[llvm] [AArch64][TTI] Fix extract cost for scalar fmul users (PR #212739)

Anutosh Bhat via llvm-commits llvm-commits at lists.llvm.org
Wed Jul 29 04:19:58 PDT 2026


https://github.com/anutosh491 updated https://github.com/llvm/llvm-project/pull/212739

>From ae7eeea24c20ba0d0cf9357af91e531c5c2c1bb8 Mon Sep 17 00:00:00 2001
From: anutosh491 <andersonbhat491 at gmail.com>
Date: Wed, 29 Jul 2026 16:34:16 +0530
Subject: [PATCH] [AArch64][TTI] Fix extract cost for scalar fmul users

---
 .../AArch64/AArch64TargetTransformInfo.cpp    | 24 ++++----
 .../AArch64/extract-fmul-cost.ll              | 59 +++++++++++++++++++
 2 files changed, 72 insertions(+), 11 deletions(-)
 create mode 100644 llvm/test/Transforms/SLPVectorizer/AArch64/extract-fmul-cost.ll

diff --git a/llvm/lib/Target/AArch64/AArch64TargetTransformInfo.cpp b/llvm/lib/Target/AArch64/AArch64TargetTransformInfo.cpp
index 341ac1008d7e7..fbfabd20ab28f 100644
--- a/llvm/lib/Target/AArch64/AArch64TargetTransformInfo.cpp
+++ b/llvm/lib/Target/AArch64/AArch64TargetTransformInfo.cpp
@@ -4433,7 +4433,9 @@ InstructionCost AArch64TTIImpl::getVectorInstrCostHelper(
 
     if (Scalar) {
       DenseMap<User *, unsigned> UserToExtractIdx;
-      for (auto *U : Scalar->users()) {
+      for (auto &[S, U, L] : ScalarUserAndIdx) {
+        if (S != Scalar)
+          continue;
         if (!IsUserFMulScalarTy(U))
           return false;
         // Recording entry for the user is important. Index value is not
@@ -4442,19 +4444,19 @@ InstructionCost AArch64TTIImpl::getVectorInstrCostHelper(
       }
       if (UserToExtractIdx.empty())
         return false;
-      for (auto &[S, U, L] : ScalarUserAndIdx) {
-        for (auto *U : S->users()) {
-          if (UserToExtractIdx.contains(U)) {
-            auto *FMul = cast<BinaryOperator>(U);
-            auto *Op0 = FMul->getOperand(0);
-            auto *Op1 = FMul->getOperand(1);
-            if ((Op0 == S && Op1 == S) || Op0 != S || Op1 != S) {
-              UserToExtractIdx[U] = L;
-              break;
-            }
+
+      for (auto &[U, L] : UserToExtractIdx) {
+        auto *FMul = cast<BinaryOperator>(U);
+        Value *OtherOp = FMul->getOperand(0) == Scalar ? FMul->getOperand(1)
+                                                       : FMul->getOperand(0);
+        for (auto &[S, ScalarUser, Lane] : ScalarUserAndIdx) {
+          if (S == OtherOp && ScalarUser == U) {
+            L = Lane;
+            break;
           }
         }
       }
+
       for (auto &[U, L] : UserToExtractIdx) {
         if (!IsExtractLaneEquivalentToZero(Index, Val->getScalarSizeInBits()) &&
             !IsExtractLaneEquivalentToZero(L, Val->getScalarSizeInBits()))
diff --git a/llvm/test/Transforms/SLPVectorizer/AArch64/extract-fmul-cost.ll b/llvm/test/Transforms/SLPVectorizer/AArch64/extract-fmul-cost.ll
new file mode 100644
index 0000000000000..a2d1b1f78c9d2
--- /dev/null
+++ b/llvm/test/Transforms/SLPVectorizer/AArch64/extract-fmul-cost.ll
@@ -0,0 +1,59 @@
+; NOTE: Assertions have been autogenerated by utils/update_test_checks.py UTC_ARGS: --version 6
+; RUN: opt -passes=slp-vectorizer -slp-threshold=3 \
+; RUN:   -mtriple=aarch64-unknown-linux -S < %s | FileCheck %s
+
+define double @lane1_times_lane2(ptr %a, ptr %b, ptr %out) {
+; CHECK-LABEL: define double @lane1_times_lane2(
+; CHECK-SAME: ptr [[A:%.*]], ptr [[B:%.*]], ptr [[OUT:%.*]]) {
+; CHECK-NEXT:  [[ENTRY:.*:]]
+; CHECK-NEXT:    [[A0P:%.*]] = getelementptr double, ptr [[A]], i64 0
+; CHECK-NEXT:    [[A2P:%.*]] = getelementptr double, ptr [[A]], i64 2
+; CHECK-NEXT:    [[B0P:%.*]] = getelementptr double, ptr [[B]], i64 0
+; CHECK-NEXT:    [[B2P:%.*]] = getelementptr double, ptr [[B]], i64 2
+; CHECK-NEXT:    [[O0:%.*]] = getelementptr double, ptr [[OUT]], i64 0
+; CHECK-NEXT:    [[O2:%.*]] = getelementptr double, ptr [[OUT]], i64 2
+; CHECK-NEXT:    [[TMP0:%.*]] = load <2 x double>, ptr [[A0P]], align 8
+; CHECK-NEXT:    [[TMP1:%.*]] = load <2 x double>, ptr [[B0P]], align 8
+; CHECK-NEXT:    [[TMP2:%.*]] = fadd <2 x double> [[TMP0]], [[TMP1]]
+; CHECK-NEXT:    [[TMP3:%.*]] = load <2 x double>, ptr [[A2P]], align 8
+; CHECK-NEXT:    [[TMP4:%.*]] = load <2 x double>, ptr [[B2P]], align 8
+; CHECK-NEXT:    [[TMP5:%.*]] = fadd <2 x double> [[TMP3]], [[TMP4]]
+; CHECK-NEXT:    store <2 x double> [[TMP2]], ptr [[O0]], align 8
+; CHECK-NEXT:    store <2 x double> [[TMP5]], ptr [[O2]], align 8
+; CHECK-NEXT:    [[TMP6:%.*]] = extractelement <2 x double> [[TMP2]], i64 1
+; CHECK-NEXT:    [[TMP7:%.*]] = extractelement <2 x double> [[TMP5]], i64 0
+; CHECK-NEXT:    [[MUL:%.*]] = fmul double [[TMP6]], [[TMP7]]
+; CHECK-NEXT:    ret double [[MUL]]
+;
+entry:
+  %a0p = getelementptr double, ptr %a, i64 0
+  %a1p = getelementptr double, ptr %a, i64 1
+  %a2p = getelementptr double, ptr %a, i64 2
+  %a3p = getelementptr double, ptr %a, i64 3
+  %b0p = getelementptr double, ptr %b, i64 0
+  %b1p = getelementptr double, ptr %b, i64 1
+  %b2p = getelementptr double, ptr %b, i64 2
+  %b3p = getelementptr double, ptr %b, i64 3
+  %a0 = load double, ptr %a0p, align 8
+  %a1 = load double, ptr %a1p, align 8
+  %a2 = load double, ptr %a2p, align 8
+  %a3 = load double, ptr %a3p, align 8
+  %b0 = load double, ptr %b0p, align 8
+  %b1 = load double, ptr %b1p, align 8
+  %b2 = load double, ptr %b2p, align 8
+  %b3 = load double, ptr %b3p, align 8
+  %x0 = fadd double %a0, %b0
+  %x1 = fadd double %a1, %b1
+  %x2 = fadd double %a2, %b2
+  %x3 = fadd double %a3, %b3
+  %o0 = getelementptr double, ptr %out, i64 0
+  %o1 = getelementptr double, ptr %out, i64 1
+  %o2 = getelementptr double, ptr %out, i64 2
+  %o3 = getelementptr double, ptr %out, i64 3
+  store double %x0, ptr %o0, align 8
+  store double %x1, ptr %o1, align 8
+  store double %x2, ptr %o2, align 8
+  store double %x3, ptr %o3, align 8
+  %mul = fmul double %x1, %x2
+  ret double %mul
+}



More information about the llvm-commits mailing list