[llvm] [SLP]Drop unprofitable splat gather subtrees when nothing was trimmed (PR #221717)

via llvm-commits llvm-commits at lists.llvm.org
Mon Sep 7 05:41:16 PDT 2026


llvmorg-github-actions[bot] wrote:


<!--LLVM PR SUMMARY COMMENT-->

@llvm/pr-subscribers-vectorizers

Author: Alexey Bataev (alexey-bataev)

<details>
<summary>Changes</summary>

The splat subtree keep/drop check ran only when the main tree trimming
changed something; otherwise the subtree was kept unconditionally, and
its cost plus the extracts it forces for the remaining scalar uses
rejected otherwise profitable trees. Run the check in both paths and
delete gathered loads subtrees left without surviving gather users.

Fixes the perf regression from #<!-- -->220250, reported in https://github.com/llvm/llvm-project/pull/220250?email_source=notifications&email_token=ABI45DSHUYRFN4NA4DILB5T5NZTZTA5CNFSNUABFM5UWIORPF5TWS5BNNB2WEL2JONZXKZKDN5WW2ZLOOQXTKNJWG4ZTSMBZGUYKM4TFMFZW63VMON2GC5DFL5RWQYLOM5S2KZLWMVXHJLDGN5XXIZLSL5RWY2LDNM#issuecomment-5567390950.


---
Full diff: https://github.com/llvm/llvm-project/pull/221717.diff


2 Files Affected:

- (modified) llvm/lib/Transforms/Vectorize/SLPVectorizer.cpp (+62-38) 
- (modified) llvm/test/Transforms/SLPVectorizer/AArch64/splat-gather-subtree-store-chain.ll (+21-36) 


``````````diff
diff --git a/llvm/lib/Transforms/Vectorize/SLPVectorizer.cpp b/llvm/lib/Transforms/Vectorize/SLPVectorizer.cpp
index 2b6eed20572fa..f2729b44df67e 100644
--- a/llvm/lib/Transforms/Vectorize/SLPVectorizer.cpp
+++ b/llvm/lib/Transforms/Vectorize/SLPVectorizer.cpp
@@ -19563,18 +19563,6 @@ BoUpSLP::calculateTreeCostAndTrimNonProfitable(ArrayRef<Value *> VectorizedVals,
     }
     Worklist.pop();
   }
-  if (!Changed) {
-    // The splat subtrees are not linked to the tree root, so their cost is
-    // not included in the root's subtree cost; add it explicitly.
-    InstructionCost TotalCost = std::get<1>(SubtreeCosts.front());
-    for (const TreeEntry *TE : SplatGatheredScalarsRoots)
-      TotalCost += std::get<1>(SubtreeCosts[TE->Idx]);
-    return TotalCost;
-  }
-
-  SmallPtrSet<TreeEntry *, 4> SubtreesToDelete;
-  SmallPtrSet<TreeEntry *, 4> DroppedSplatSubtrees;
-  InstructionCost LoadsExtractsCost = 0;
   using ValuesToInsertTy =
       SmallDenseMap<const TreeEntry *, SmallVector<Value *>>;
   auto GetScalarTy = [&](const TreeEntry *TE) {
@@ -19621,6 +19609,67 @@ BoUpSLP::calculateTreeCostAndTrimNonProfitable(ArrayRef<Value *> VectorizedVals,
     }
     return BVCost;
   };
+  // A splat subtree pays off only if its full price plus the extracts of its
+  // scalars used by the remaining scalar code is not worse than materializing
+  // the splatted scalars directly in the surviving gathers.
+  auto IsSplatSubtreeProfitable = [&](TreeEntry *TE, Type *ScalarTy,
+                                      const ValuesToInsertTy &ValuesToInsert) {
+    APInt ExtractElts = APInt::getZero(TE->getVectorFactor());
+    for (Value *V : TE->Scalars) {
+      if (!isa<Instruction>(V) || TE->isCopyableElement(V))
+        continue;
+      // Too many users - the scalar is extracted anyway.
+      if (V->hasNUsesOrMore(UsesLimit) || any_of(V->users(), [&](User *U) {
+            return none_of(getTreeEntries(U), [&](const TreeEntry *UseTE) {
+              return !DeletedNodes.contains(UseTE) &&
+                     !TransformedToGatherNodes.contains(UseTE);
+            });
+          }))
+        ExtractElts.setBit(TE->findLaneForValue(V));
+    }
+    InstructionCost KeepCost = getScalarizationOverhead(
+        *TTI, SLPReVec, ScalarTy,
+        cast<VectorType>(getWidenedType(ScalarTy, TE->getVectorFactor())),
+        ExtractElts, /*Insert=*/false, /*Extract=*/true, CostKind);
+    // Add the cost of the subtree itself, computed before any trimming:
+    // trimming of the subtree's own nodes would otherwise make it look
+    // artificially cheap.
+    KeepCost += std::get<1>(SubtreeCosts[TE->Idx]);
+    return KeepCost <= GetGatherInsertCost(ScalarTy, ValuesToInsert);
+  };
+  if (!Changed) {
+    // The splat subtrees are not linked to the tree root, so their cost is
+    // not included in the root's subtree cost; add it explicitly. Drop the
+    // unprofitable ones instead of letting them reject the whole tree.
+    InstructionCost TotalCost = std::get<1>(SubtreeCosts.front());
+    for (TreeEntry *TE : SplatGatheredScalarsRoots) {
+      ValuesToInsertTy ValuesToInsert;
+      if (!FindDemandedElts(TE, ValuesToInsert).isZero() &&
+          IsSplatSubtreeProfitable(TE, GetScalarTy(TE), ValuesToInsert)) {
+        TotalCost += std::get<1>(SubtreeCosts[TE->Idx]);
+        continue;
+      }
+      DeletedNodes.insert(TE);
+      for (unsigned Idx : std::get<2>(SubtreeCosts[TE->Idx]))
+        DeletedNodes.insert(VectorizableTree[Idx].get());
+    }
+    // Gathered loads subtrees left without surviving gather users are dead.
+    for (TreeEntry *TE : GatheredLoadsNodes) {
+      if (DeletedNodes.contains(TE))
+        continue;
+      ValuesToInsertTy ValuesToInsert;
+      if (!FindDemandedElts(TE, ValuesToInsert).isZero())
+        continue;
+      DeletedNodes.insert(TE);
+      for (unsigned Idx : std::get<2>(SubtreeCosts[TE->Idx]))
+        DeletedNodes.insert(VectorizableTree[Idx].get());
+    }
+    return TotalCost;
+  }
+
+  SmallPtrSet<TreeEntry *, 4> SubtreesToDelete;
+  SmallPtrSet<TreeEntry *, 4> DroppedSplatSubtrees;
+  InstructionCost LoadsExtractsCost = 0;
   // Check if all loads of gathered loads nodes are marked for deletion. In this
   // case the whole gathered loads subtree must be deleted.
   // Also, try to account for extracts, which might be required, if only part of
@@ -19662,32 +19711,7 @@ BoUpSLP::calculateTreeCostAndTrimNonProfitable(ArrayRef<Value *> VectorizedVals,
     ValuesToInsertTy ValuesToInsert;
     APInt DemandedElts = FindDemandedElts(TE, ValuesToInsert);
     if (!DemandedElts.isZero()) {
-      Type *ScalarTy = GetScalarTy(TE);
-      // Lanes of the subtree scalars still used by the remaining scalar code
-      // must be extracted if the subtree is kept.
-      APInt ExtractElts = APInt::getZero(TE->getVectorFactor());
-      for (Value *V : TE->Scalars) {
-        if (!isa<Instruction>(V) || TE->isCopyableElement(V))
-          continue;
-        // Too many users - the scalar is extracted anyway.
-        if (V->hasNUsesOrMore(UsesLimit) || any_of(V->users(), [&](User *U) {
-              return none_of(getTreeEntries(U), [&](const TreeEntry *UseTE) {
-                return !DeletedNodes.contains(UseTE) &&
-                       !TransformedToGatherNodes.contains(UseTE);
-              });
-            }))
-          ExtractElts.setBit(TE->findLaneForValue(V));
-      }
-      InstructionCost KeepCost = getScalarizationOverhead(
-          *TTI, SLPReVec, ScalarTy,
-          cast<VectorType>(getWidenedType(ScalarTy, TE->getVectorFactor())),
-          ExtractElts, /*Insert=*/false, /*Extract=*/true, CostKind);
-      // Add the cost of the subtree itself, computed before any trimming:
-      // trimming of the subtree's own nodes would otherwise make it look
-      // artificially cheap.
-      KeepCost += std::get<1>(SubtreeCosts[TE->Idx]);
-      InstructionCost DropCost = GetGatherInsertCost(ScalarTy, ValuesToInsert);
-      if (KeepCost <= DropCost)
+      if (IsSplatSubtreeProfitable(TE, GetScalarTy(TE), ValuesToInsert))
         continue;
       // Dropped as unprofitable: exclude its cost from the reference cost, so
       // the trimming of the remaining tree is not reverted because of it, and
diff --git a/llvm/test/Transforms/SLPVectorizer/AArch64/splat-gather-subtree-store-chain.ll b/llvm/test/Transforms/SLPVectorizer/AArch64/splat-gather-subtree-store-chain.ll
index 5d7998dd7311d..678512bb23d25 100644
--- a/llvm/test/Transforms/SLPVectorizer/AArch64/splat-gather-subtree-store-chain.ll
+++ b/llvm/test/Transforms/SLPVectorizer/AArch64/splat-gather-subtree-store-chain.ll
@@ -69,53 +69,38 @@ define void @splat_subtree_with_scalar_uses(ptr noalias %out, ptr noalias %in) {
 ; CHECK-LABEL: define void @splat_subtree_with_scalar_uses(
 ; CHECK-SAME: ptr noalias [[OUT:%.*]], ptr noalias [[IN:%.*]]) {
 ; CHECK-NEXT:  [[ENTRY_RTVEC:.*:]]
-; CHECK-NEXT:    [[TMP0:%.*]] = load i32, ptr [[IN]], align 4
 ; CHECK-NEXT:    [[ARRAYIDX1:%.*]] = getelementptr inbounds nuw i8, ptr [[IN]], i64 4
-; CHECK-NEXT:    [[TMP1:%.*]] = load i32, ptr [[ARRAYIDX1]], align 4
-; CHECK-NEXT:    [[TMP5:%.*]] = add i32 [[TMP1]], [[TMP0]]
 ; CHECK-NEXT:    [[ARRAYIDX2:%.*]] = getelementptr inbounds nuw i8, ptr [[IN]], i64 8
-; CHECK-NEXT:    [[TMP7:%.*]] = load i32, ptr [[ARRAYIDX2]], align 4
 ; CHECK-NEXT:    [[ARRAYIDX3:%.*]] = getelementptr inbounds nuw i8, ptr [[IN]], i64 12
-; CHECK-NEXT:    [[TMP9:%.*]] = load i32, ptr [[ARRAYIDX3]], align 4
-; CHECK-NEXT:    [[ADD5:%.*]] = add i32 [[TMP9]], [[TMP7]]
 ; CHECK-NEXT:    [[ARRAYIDX5:%.*]] = getelementptr inbounds nuw i8, ptr [[IN]], i64 16
-; CHECK-NEXT:    [[TMP10:%.*]] = load i32, ptr [[ARRAYIDX5]], align 4
 ; CHECK-NEXT:    [[ARRAYIDX7:%.*]] = getelementptr inbounds nuw i8, ptr [[IN]], i64 20
-; CHECK-NEXT:    [[TMP11:%.*]] = load i32, ptr [[ARRAYIDX7]], align 4
-; CHECK-NEXT:    [[TMP3:%.*]] = add i32 [[TMP11]], [[TMP10]]
 ; CHECK-NEXT:    [[ARRAYIDX8:%.*]] = getelementptr inbounds nuw i8, ptr [[IN]], i64 24
-; CHECK-NEXT:    [[TMP12:%.*]] = load i32, ptr [[ARRAYIDX8]], align 4
-; CHECK-NEXT:    [[ADD9:%.*]] = add i32 [[TMP12]], [[TMP5]]
-; CHECK-NEXT:    [[XOR:%.*]] = xor i32 [[ADD9]], [[ADD5]]
-; CHECK-NEXT:    [[ADD10:%.*]] = add i32 [[XOR]], [[TMP3]]
-; CHECK-NEXT:    store i32 [[ADD10]], ptr [[OUT]], align 4
-; CHECK-NEXT:    [[ARRAYIDX6:%.*]] = getelementptr inbounds nuw i8, ptr [[IN]], i64 28
-; CHECK-NEXT:    [[TMP6:%.*]] = load i32, ptr [[ARRAYIDX6]], align 4
-; CHECK-NEXT:    [[ADD13:%.*]] = add i32 [[TMP6]], [[TMP5]]
-; CHECK-NEXT:    [[XOR14:%.*]] = xor i32 [[ADD13]], [[ADD5]]
-; CHECK-NEXT:    [[ADD15:%.*]] = add i32 [[XOR14]], [[TMP3]]
-; CHECK-NEXT:    [[ARRAYIDX16:%.*]] = getelementptr inbounds nuw i8, ptr [[OUT]], i64 4
-; CHECK-NEXT:    store i32 [[ADD15]], ptr [[ARRAYIDX16]], align 4
-; CHECK-NEXT:    [[ARRAYIDX17:%.*]] = getelementptr inbounds nuw i8, ptr [[IN]], i64 32
-; CHECK-NEXT:    [[TMP8:%.*]] = load i32, ptr [[ARRAYIDX17]], align 4
-; CHECK-NEXT:    [[ADD18:%.*]] = add i32 [[TMP8]], [[TMP5]]
-; CHECK-NEXT:    [[TMP2:%.*]] = xor i32 [[ADD18]], [[ADD5]]
-; CHECK-NEXT:    [[ADD4:%.*]] = add i32 [[TMP2]], [[TMP3]]
-; CHECK-NEXT:    [[ARRAYIDX21:%.*]] = getelementptr inbounds nuw i8, ptr [[OUT]], i64 8
-; CHECK-NEXT:    store i32 [[ADD4]], ptr [[ARRAYIDX21]], align 4
-; CHECK-NEXT:    [[ARRAYIDX22:%.*]] = getelementptr inbounds nuw i8, ptr [[IN]], i64 36
-; CHECK-NEXT:    [[TMP4:%.*]] = load i32, ptr [[ARRAYIDX22]], align 4
+; CHECK-NEXT:    [[XOR24:%.*]] = load i32, ptr [[ARRAYIDX3]], align 4
+; CHECK-NEXT:    [[TMP3:%.*]] = load i32, ptr [[ARRAYIDX2]], align 4
+; CHECK-NEXT:    [[TMP2:%.*]] = load i32, ptr [[ARRAYIDX1]], align 4
+; CHECK-NEXT:    [[TMP16:%.*]] = load i32, ptr [[IN]], align 4
+; CHECK-NEXT:    [[TMP4:%.*]] = load i32, ptr [[ARRAYIDX7]], align 4
+; CHECK-NEXT:    [[TMP5:%.*]] = load i32, ptr [[ARRAYIDX5]], align 4
 ; CHECK-NEXT:    [[ADD:%.*]] = add i32 [[TMP4]], [[TMP5]]
-; CHECK-NEXT:    [[XOR24:%.*]] = xor i32 [[ADD]], [[ADD5]]
 ; CHECK-NEXT:    [[ADD25:%.*]] = add i32 [[XOR24]], [[TMP3]]
-; CHECK-NEXT:    [[ARRAYIDX26:%.*]] = getelementptr inbounds nuw i8, ptr [[OUT]], i64 12
-; CHECK-NEXT:    store i32 [[ADD25]], ptr [[ARRAYIDX26]], align 4
+; CHECK-NEXT:    [[ADD1:%.*]] = add i32 [[TMP2]], [[TMP16]]
+; CHECK-NEXT:    [[TMP6:%.*]] = load <4 x i32>, ptr [[ARRAYIDX8]], align 4
+; CHECK-NEXT:    [[TMP7:%.*]] = insertelement <4 x i32> poison, i32 [[ADD1]], i64 0
+; CHECK-NEXT:    [[TMP8:%.*]] = shufflevector <4 x i32> [[TMP7]], <4 x i32> poison, <4 x i32> zeroinitializer
+; CHECK-NEXT:    [[TMP9:%.*]] = add <4 x i32> [[TMP6]], [[TMP8]]
+; CHECK-NEXT:    [[TMP10:%.*]] = insertelement <4 x i32> poison, i32 [[ADD25]], i64 0
+; CHECK-NEXT:    [[TMP11:%.*]] = shufflevector <4 x i32> [[TMP10]], <4 x i32> poison, <4 x i32> zeroinitializer
+; CHECK-NEXT:    [[TMP12:%.*]] = xor <4 x i32> [[TMP9]], [[TMP11]]
+; CHECK-NEXT:    [[TMP13:%.*]] = insertelement <4 x i32> poison, i32 [[ADD]], i64 0
+; CHECK-NEXT:    [[TMP14:%.*]] = shufflevector <4 x i32> [[TMP13]], <4 x i32> poison, <4 x i32> zeroinitializer
+; CHECK-NEXT:    [[TMP15:%.*]] = add <4 x i32> [[TMP12]], [[TMP14]]
+; CHECK-NEXT:    store <4 x i32> [[TMP15]], ptr [[OUT]], align 4
 ; CHECK-NEXT:    [[ARRAYIDX27_SCALAR:%.*]] = getelementptr inbounds nuw i8, ptr [[OUT]], i64 16
-; CHECK-NEXT:    store i32 [[TMP5]], ptr [[ARRAYIDX27_SCALAR]], align 4
+; CHECK-NEXT:    store i32 [[ADD1]], ptr [[ARRAYIDX27_SCALAR]], align 4
 ; CHECK-NEXT:    [[ARRAYIDX28_SCALAR:%.*]] = getelementptr inbounds nuw i8, ptr [[OUT]], i64 24
-; CHECK-NEXT:    store i32 [[ADD5]], ptr [[ARRAYIDX28_SCALAR]], align 4
+; CHECK-NEXT:    store i32 [[ADD25]], ptr [[ARRAYIDX28_SCALAR]], align 4
 ; CHECK-NEXT:    [[ARRAYIDX29_SCALAR:%.*]] = getelementptr inbounds nuw i8, ptr [[OUT]], i64 32
-; CHECK-NEXT:    store i32 [[TMP3]], ptr [[ARRAYIDX29_SCALAR]], align 4
+; CHECK-NEXT:    store i32 [[ADD]], ptr [[ARRAYIDX29_SCALAR]], align 4
 ; CHECK-NEXT:    ret void
 ;
 entry:

``````````

</details>


https://github.com/llvm/llvm-project/pull/221717


More information about the llvm-commits mailing list