[llvm] [InstCombine] Fold special cases of vector partial reduction add (PR #226714)

via llvm-commits llvm-commits at lists.llvm.org
Sat Sep 26 08:52:18 PDT 2026


llvmorg-github-actions[bot] wrote:


<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-llvm-transforms

@llvm/pr-subscribers-llvm-analysis

Author: Chennes (Chennesxu)

<details>
<summary>Changes</summary>

Extend the constant folding from #<!-- -->212112 to scalable splat constants and combine partial reductions with constant inputs and non-constant accumulators into vector additions.

For fixed-length constant inputs, reuse the existing constant folder and its chosen lane grouping. For scalable constant splats and single-use runtime splats with a constant reduction ratio, compute each lane's input contribution by multiplying the splat value by that ratio. Also fold zero inputs to the accumulator and equal-width operands to an add.

For a single-use nxv16i32 splat with an nxv4i32 accumulator, an AArch64 SVE example lowers to three instructions instead of five, excluding ret.

Nonzero mixed fixed/scalable cases remain unchanged because their reduction ratio depends on vscale.

This implements the zero/splat and constant-input cases from items 1 and 2 of #<!-- -->224350, and adds equal-width and single-use runtime-splat folds from item 3.

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


6 Files Affected:

- (modified) llvm/lib/Analysis/ConstantFolding.cpp (+26) 
- (modified) llvm/lib/Analysis/InstructionSimplify.cpp (+4) 
- (modified) llvm/lib/Transforms/InstCombine/InstCombineCalls.cpp (+34) 
- (added) llvm/test/Transforms/InstCombine/vector-partial-reduce-add.ll (+156) 
- (modified) llvm/test/Transforms/InstSimplify/ConstProp/vecreduce.ll (+50-2) 
- (added) llvm/test/Transforms/InstSimplify/vector-partial-reduce-add.ll (+40) 


``````````diff
diff --git a/llvm/lib/Analysis/ConstantFolding.cpp b/llvm/lib/Analysis/ConstantFolding.cpp
index b09f348a3b30a..ea7f3a2af7efe 100644
--- a/llvm/lib/Analysis/ConstantFolding.cpp
+++ b/llvm/lib/Analysis/ConstantFolding.cpp
@@ -2431,6 +2431,30 @@ Constant *constantFoldVectorReduce(Intrinsic::ID IID, Constant *Op) {
 static Constant *constantFoldVectorPartialReduceAdd(Constant *Acc,
                                                     Constant *Input,
                                                     const DataLayout &DL) {
+  if (Input->isNullValue())
+    return Acc;
+
+  if (auto *AccTy = dyn_cast<ScalableVectorType>(Acc->getType())) {
+    auto *InputTy = dyn_cast<ScalableVectorType>(Input->getType());
+    Constant *Splat = Input->getSplatValue();
+    if (!InputTy || !Splat)
+      return nullptr;
+
+    // The vscale factors cancel, so each result lane accumulates a constant
+    // number of input elements.
+    unsigned Ratio = InputTy->getMinNumElements() / AccTy->getMinNumElements();
+    Constant *Sum = ConstantFoldBinaryOpOperands(
+        Instruction::Mul, Splat,
+        ConstantInt::get(Splat->getType(), Ratio, /*IsSigned=*/false,
+                         /*ImplicitTrunc=*/true),
+        DL);
+    if (!Sum)
+      return nullptr;
+    return ConstantFoldBinaryOpOperands(
+        Instruction::Add, Acc,
+        ConstantVector::getSplat(AccTy->getElementCount(), Sum), DL);
+  }
+
   auto *AccTy = cast<FixedVectorType>(Acc->getType());
   // A fixed result type does not guarantee a fixed input type.
   auto *InputTy = dyn_cast<FixedVectorType>(Input->getType());
@@ -4589,6 +4613,8 @@ static Constant *ConstantFoldScalableVectorCall(
     ArrayRef<Constant *> Operands, const DataLayout &DL,
     const TargetLibraryInfo *TLI, const CallBase *Call) {
   switch (IntrinsicID) {
+  case Intrinsic::vector_partial_reduce_add:
+    return constantFoldVectorPartialReduceAdd(Operands[0], Operands[1], DL);
   case Intrinsic::aarch64_sve_convert_from_svbool: {
     Constant *Src = Operands[0];
     if (!Src->isNullValue())
diff --git a/llvm/lib/Analysis/InstructionSimplify.cpp b/llvm/lib/Analysis/InstructionSimplify.cpp
index 4e91cb44c7a18..941e254e15fa6 100644
--- a/llvm/lib/Analysis/InstructionSimplify.cpp
+++ b/llvm/lib/Analysis/InstructionSimplify.cpp
@@ -6952,6 +6952,10 @@ static Value *simplifyBinaryIntrinsic(Intrinsic::ID IID, Type *ReturnType,
                                       const SimplifyQuery &Q) {
   unsigned BitWidth = ReturnType->getScalarSizeInBits();
   switch (IID) {
+  case Intrinsic::vector_partial_reduce_add:
+    if (match(Op1, m_Zero()))
+      return Op0;
+    break;
   case Intrinsic::get_active_lane_mask: {
     if (match(Op1, m_Zero()))
       return ConstantInt::getFalse(ReturnType);
diff --git a/llvm/lib/Transforms/InstCombine/InstCombineCalls.cpp b/llvm/lib/Transforms/InstCombine/InstCombineCalls.cpp
index 7fae7b8cea710..2afeb1e869e23 100644
--- a/llvm/lib/Transforms/InstCombine/InstCombineCalls.cpp
+++ b/llvm/lib/Transforms/InstCombine/InstCombineCalls.cpp
@@ -24,6 +24,7 @@
 #include "llvm/Analysis/AliasAnalysis.h"
 #include "llvm/Analysis/AssumeBundleQueries.h"
 #include "llvm/Analysis/AssumptionCache.h"
+#include "llvm/Analysis/ConstantFolding.h"
 #include "llvm/Analysis/InstructionSimplify.h"
 #include "llvm/Analysis/Loads.h"
 #include "llvm/Analysis/MemoryBuiltins.h"
@@ -4227,6 +4228,39 @@ Instruction *InstCombinerImpl::visitCallInst(CallInst &CI) {
     }
     break;
   }
+  case Intrinsic::vector_partial_reduce_add: {
+    Value *Acc = II->getArgOperand(0);
+    Value *Input = II->getArgOperand(1);
+    if (Acc->getType() == Input->getType())
+      return BinaryOperator::CreateAdd(Acc, Input);
+
+    // Separate the accumulator so the constant input can be reduced using
+    // the same grouping as a fully constant partial reduction.
+    if (auto *C = dyn_cast<Constant>(Input)) {
+      Constant *Zero = Constant::getNullValue(Acc->getType());
+      if (Constant *Sum =
+              ConstantFoldCall(II, II->getCalledFunction(), {Zero, C}, &TLI))
+        return BinaryOperator::CreateAdd(Acc, Sum);
+    }
+
+    ElementCount AccEC = cast<VectorType>(Acc->getType())->getElementCount();
+    ElementCount InputEC =
+        cast<VectorType>(Input->getType())->getElementCount();
+    // A fixed accumulator and scalable input have a vscale-dependent ratio.
+    if (AccEC.isScalable() != InputEC.isScalable())
+      break;
+
+    // Avoid creating an additional splat when the input has other uses.
+    if (Value *Splat = Input->hasOneUse() ? getSplatValue(Input) : nullptr) {
+      unsigned Ratio = InputEC.getKnownMinValue() / AccEC.getKnownMinValue();
+      Value *Sum = Builder.CreateMul(
+          Splat, ConstantInt::get(Splat->getType(), Ratio, /*IsSigned=*/false,
+                                  /*ImplicitTrunc=*/true));
+      return BinaryOperator::CreateAdd(Acc,
+                                       Builder.CreateVectorSplat(AccEC, Sum));
+    }
+    break;
+  }
   case Intrinsic::vector_reduce_or:
   case Intrinsic::vector_reduce_and: {
     // Canonicalize logical or/and reductions:
diff --git a/llvm/test/Transforms/InstCombine/vector-partial-reduce-add.ll b/llvm/test/Transforms/InstCombine/vector-partial-reduce-add.ll
new file mode 100644
index 0000000000000..a8fa6b1f3dfb5
--- /dev/null
+++ b/llvm/test/Transforms/InstCombine/vector-partial-reduce-add.ll
@@ -0,0 +1,156 @@
+; NOTE: Assertions have been autogenerated by utils/update_test_checks.py UTC_ARGS: --function-signature
+; RUN: opt < %s -passes=instcombine -S | FileCheck %s
+; RUN: opt < %s -passes=instcombine -use-constant-int-for-fixed-length-splat -use-constant-int-for-scalable-splat -S | FileCheck %s
+
+define <2 x i8> @constant_input(<2 x i8> %acc) {
+;
+; CHECK-LABEL: define {{[^@]+}}@constant_input
+; CHECK-SAME: (<2 x i8> [[ACC:%.*]]) {
+; CHECK-NEXT:    [[R:%.*]] = add <2 x i8> [[ACC]], <i8 4, i8 6>
+; CHECK-NEXT:    ret <2 x i8> [[R]]
+;
+  %r = call <2 x i8> @llvm.vector.partial.reduce.add.v2i8.v4i8(<2 x i8> %acc, <4 x i8> <i8 1, i8 2, i8 3, i8 4>)
+  ret <2 x i8> %r
+}
+
+define <2 x i8> @constant_input_poison(<2 x i8> %acc) {
+;
+; CHECK-LABEL: define {{[^@]+}}@constant_input_poison
+; CHECK-SAME: (<2 x i8> [[ACC:%.*]]) {
+; CHECK-NEXT:    [[R:%.*]] = add <2 x i8> [[ACC]], <i8 4, i8 poison>
+; CHECK-NEXT:    ret <2 x i8> [[R]]
+;
+  %r = call <2 x i8> @llvm.vector.partial.reduce.add.v2i8.v4i8(<2 x i8> %acc, <4 x i8> <i8 1, i8 poison, i8 3, i8 4>)
+  ret <2 x i8> %r
+}
+
+define <vscale x 2 x i8> @scalable_constant_input(<vscale x 2 x i8> %acc) {
+;
+; CHECK-LABEL: define {{[^@]+}}@scalable_constant_input
+; CHECK-SAME: (<vscale x 2 x i8> [[ACC:%.*]]) {
+; CHECK-NEXT:    [[R:%.*]] = add <vscale x 2 x i8> [[ACC]], splat (i8 12)
+; CHECK-NEXT:    ret <vscale x 2 x i8> [[R]]
+;
+  %r = call <vscale x 2 x i8> @llvm.vector.partial.reduce.add.nxv2i8.nxv8i8(<vscale x 2 x i8> %acc, <vscale x 8 x i8> splat (i8 3))
+  ret <vscale x 2 x i8> %r
+}
+
+define <2 x i8> @same_width(<2 x i8> %acc, <2 x i8> %input) {
+;
+; CHECK-LABEL: define {{[^@]+}}@same_width
+; CHECK-SAME: (<2 x i8> [[ACC:%.*]], <2 x i8> [[INPUT:%.*]]) {
+; CHECK-NEXT:    [[R:%.*]] = add <2 x i8> [[ACC]], [[INPUT]]
+; CHECK-NEXT:    ret <2 x i8> [[R]]
+;
+  %r = call <2 x i8> @llvm.vector.partial.reduce.add.v2i8.v2i8(<2 x i8> %acc, <2 x i8> %input)
+  ret <2 x i8> %r
+}
+
+define <vscale x 2 x i8> @scalable_same_width(<vscale x 2 x i8> %acc, <vscale x 2 x i8> %input) {
+;
+; CHECK-LABEL: define {{[^@]+}}@scalable_same_width
+; CHECK-SAME: (<vscale x 2 x i8> [[ACC:%.*]], <vscale x 2 x i8> [[INPUT:%.*]]) {
+; CHECK-NEXT:    [[R:%.*]] = add <vscale x 2 x i8> [[ACC]], [[INPUT]]
+; CHECK-NEXT:    ret <vscale x 2 x i8> [[R]]
+;
+  %r = call <vscale x 2 x i8> @llvm.vector.partial.reduce.add.nxv2i8.nxv2i8(<vscale x 2 x i8> %acc, <vscale x 2 x i8> %input)
+  ret <vscale x 2 x i8> %r
+}
+
+define <2 x i8> @splat_input(<2 x i8> %acc, i8 %x) {
+;
+; CHECK-LABEL: define {{[^@]+}}@splat_input
+; CHECK-SAME: (<2 x i8> [[ACC:%.*]], i8 [[X:%.*]]) {
+; CHECK-NEXT:    [[TMP1:%.*]] = shl i8 [[X]], 1
+; CHECK-NEXT:    [[DOTSPLATINSERT:%.*]] = insertelement <2 x i8> poison, i8 [[TMP1]], i64 0
+; CHECK-NEXT:    [[DOTSPLAT:%.*]] = shufflevector <2 x i8> [[DOTSPLATINSERT]], <2 x i8> poison, <2 x i32> zeroinitializer
+; CHECK-NEXT:    [[R:%.*]] = add <2 x i8> [[ACC]], [[DOTSPLAT]]
+; CHECK-NEXT:    ret <2 x i8> [[R]]
+;
+  %ins = insertelement <4 x i8> poison, i8 %x, i32 0
+  %splat = shufflevector <4 x i8> %ins, <4 x i8> poison, <4 x i32> zeroinitializer
+  %r = call <2 x i8> @llvm.vector.partial.reduce.add.v2i8.v4i8(<2 x i8> %acc, <4 x i8> %splat)
+  ret <2 x i8> %r
+}
+
+define <vscale x 2 x i8> @scalable_splat_input(<vscale x 2 x i8> %acc, i8 %x) {
+;
+; CHECK-LABEL: define {{[^@]+}}@scalable_splat_input
+; CHECK-SAME: (<vscale x 2 x i8> [[ACC:%.*]], i8 [[X:%.*]]) {
+; CHECK-NEXT:    [[TMP1:%.*]] = mul i8 [[X]], 3
+; CHECK-NEXT:    [[DOTSPLATINSERT:%.*]] = insertelement <vscale x 2 x i8> poison, i8 [[TMP1]], i64 0
+; CHECK-NEXT:    [[DOTSPLAT:%.*]] = shufflevector <vscale x 2 x i8> [[DOTSPLATINSERT]], <vscale x 2 x i8> poison, <vscale x 2 x i32> zeroinitializer
+; CHECK-NEXT:    [[R:%.*]] = add <vscale x 2 x i8> [[ACC]], [[DOTSPLAT]]
+; CHECK-NEXT:    ret <vscale x 2 x i8> [[R]]
+;
+  %ins = insertelement <vscale x 6 x i8> poison, i8 %x, i32 0
+  %splat = shufflevector <vscale x 6 x i8> %ins, <vscale x 6 x i8> poison, <vscale x 6 x i32> zeroinitializer
+  %r = call <vscale x 2 x i8> @llvm.vector.partial.reduce.add.nxv2i8.nxv6i8(<vscale x 2 x i8> %acc, <vscale x 6 x i8> %splat)
+  ret <vscale x 2 x i8> %r
+}
+
+define <2 x i1> @splat_input_i1(<2 x i1> %acc, i1 %x) {
+;
+; CHECK-LABEL: define {{[^@]+}}@splat_input_i1
+; CHECK-SAME: (<2 x i1> [[ACC:%.*]], i1 [[X:%.*]]) {
+; CHECK-NEXT:    ret <2 x i1> [[ACC]]
+;
+  %ins = insertelement <4 x i1> poison, i1 %x, i32 0
+  %splat = shufflevector <4 x i1> %ins, <4 x i1> poison, <4 x i32> zeroinitializer
+  %r = call <2 x i1> @llvm.vector.partial.reduce.add.v2i1.v4i1(<2 x i1> %acc, <4 x i1> %splat)
+  ret <2 x i1> %r
+}
+
+define <2 x i8> @splat_input_multiuse(<2 x i8> %acc, i8 %x, ptr %p) {
+; CHECK-LABEL: define {{[^@]+}}@splat_input_multiuse
+; CHECK-SAME: (<2 x i8> [[ACC:%.*]], i8 [[X:%.*]], ptr [[P:%.*]]) {
+; CHECK-NEXT:    [[INS:%.*]] = insertelement <4 x i8> poison, i8 [[X]], i64 0
+; CHECK-NEXT:    [[SPLAT:%.*]] = shufflevector <4 x i8> [[INS]], <4 x i8> poison, <4 x i32> zeroinitializer
+; CHECK-NEXT:    store <4 x i8> [[SPLAT]], ptr [[P]], align 4
+; CHECK-NEXT:    [[R:%.*]] = call <2 x i8> @llvm.vector.partial.reduce.add.v2i8.v4i8(<2 x i8> [[ACC]], <4 x i8> [[SPLAT]])
+; CHECK-NEXT:    ret <2 x i8> [[R]]
+;
+  %ins = insertelement <4 x i8> poison, i8 %x, i32 0
+  %splat = shufflevector <4 x i8> %ins, <4 x i8> poison, <4 x i32> zeroinitializer
+  store <4 x i8> %splat, ptr %p
+  %r = call <2 x i8> @llvm.vector.partial.reduce.add.v2i8.v4i8(<2 x i8> %acc, <4 x i8> %splat)
+  ret <2 x i8> %r
+}
+
+define <2 x i8> @nonconstant_input(<2 x i8> %acc, <4 x i8> %input) {
+;
+; CHECK-LABEL: define {{[^@]+}}@nonconstant_input
+; CHECK-SAME: (<2 x i8> [[ACC:%.*]], <4 x i8> [[INPUT:%.*]]) {
+; CHECK-NEXT:    [[R:%.*]] = call <2 x i8> @llvm.vector.partial.reduce.add.v2i8.v4i8(<2 x i8> [[ACC]], <4 x i8> [[INPUT]])
+; CHECK-NEXT:    ret <2 x i8> [[R]]
+;
+  %r = call <2 x i8> @llvm.vector.partial.reduce.add.v2i8.v4i8(<2 x i8> %acc, <4 x i8> %input)
+  ret <2 x i8> %r
+}
+
+; The reduction ratio depends on vscale, so it cannot be used as a constant.
+define <2 x i8> @mixed_constant_input(<2 x i8> %acc) {
+;
+; CHECK-LABEL: define {{[^@]+}}@mixed_constant_input
+; CHECK-SAME: (<2 x i8> [[ACC:%.*]]) {
+; CHECK-NEXT:    [[R:%.*]] = call <2 x i8> @llvm.vector.partial.reduce.add.v2i8.nxv4i8(<2 x i8> [[ACC]], <vscale x 4 x i8> splat (i8 3))
+; CHECK-NEXT:    ret <2 x i8> [[R]]
+;
+  %r = call <2 x i8> @llvm.vector.partial.reduce.add.v2i8.nxv4i8(<2 x i8> %acc, <vscale x 4 x i8> splat (i8 3))
+  ret <2 x i8> %r
+}
+
+define <2 x i8> @mixed_splat_input(<2 x i8> %acc, i8 %x) {
+;
+; CHECK-LABEL: define {{[^@]+}}@mixed_splat_input
+; CHECK-SAME: (<2 x i8> [[ACC:%.*]], i8 [[X:%.*]]) {
+; CHECK-NEXT:    [[INS:%.*]] = insertelement <vscale x 4 x i8> poison, i8 [[X]], i64 0
+; CHECK-NEXT:    [[SPLAT:%.*]] = shufflevector <vscale x 4 x i8> [[INS]], <vscale x 4 x i8> poison, <vscale x 4 x i32> zeroinitializer
+; CHECK-NEXT:    [[R:%.*]] = call <2 x i8> @llvm.vector.partial.reduce.add.v2i8.nxv4i8(<2 x i8> [[ACC]], <vscale x 4 x i8> [[SPLAT]])
+; CHECK-NEXT:    ret <2 x i8> [[R]]
+;
+  %ins = insertelement <vscale x 4 x i8> poison, i8 %x, i32 0
+  %splat = shufflevector <vscale x 4 x i8> %ins, <vscale x 4 x i8> poison, <vscale x 4 x i32> zeroinitializer
+  %r = call <2 x i8> @llvm.vector.partial.reduce.add.v2i8.nxv4i8(<2 x i8> %acc, <vscale x 4 x i8> %splat)
+  ret <2 x i8> %r
+}
diff --git a/llvm/test/Transforms/InstSimplify/ConstProp/vecreduce.ll b/llvm/test/Transforms/InstSimplify/ConstProp/vecreduce.ll
index 6c2ee928d7c0e..353c19d4c28f1 100644
--- a/llvm/test/Transforms/InstSimplify/ConstProp/vecreduce.ll
+++ b/llvm/test/Transforms/InstSimplify/ConstProp/vecreduce.ll
@@ -969,8 +969,7 @@ define <4 x i32> @partial_reduce_add_nonconstant_input(<16 x i32> %input) {
 
 define <4 x i32> @partial_reduce_add_mixed_fixed_scalable() {
 ; CHECK-LABEL: @partial_reduce_add_mixed_fixed_scalable(
-; CHECK-NEXT:    [[X:%.*]] = call <4 x i32> @llvm.vector.partial.reduce.add.v4i32.nxv8i32(<4 x i32> zeroinitializer, <vscale x 8 x i32> zeroinitializer)
-; CHECK-NEXT:    ret <4 x i32> [[X]]
+; CHECK-NEXT:    ret <4 x i32> zeroinitializer
 ;
   %x = call <4 x i32> @llvm.vector.partial.reduce.add.v4i32.nxv8i32(
   <4 x i32> zeroinitializer, <vscale x 8 x i32> zeroinitializer)
@@ -1029,3 +1028,52 @@ define <2 x i8> @partial_reduce_add_wrap() {
   <4 x i8> <i8 1, i8 1, i8 2, i8 4>)
   ret <2 x i8> %x
 }
+
+define <vscale x 2 x i8> @partial_reduce_add_scalable_zero() {
+; CHECK-LABEL: @partial_reduce_add_scalable_zero(
+; CHECK-NEXT:    ret <vscale x 2 x i8> splat (i8 3)
+;
+  %r = call <vscale x 2 x i8> @llvm.vector.partial.reduce.add.nxv2i8.nxv4i8(<vscale x 2 x i8> splat (i8 3), <vscale x 4 x i8> zeroinitializer)
+  ret <vscale x 2 x i8> %r
+}
+
+define <vscale x 2 x i8> @partial_reduce_add_scalable_splat() {
+; CHECK-LABEL: @partial_reduce_add_scalable_splat(
+; CHECK-NEXT:    ret <vscale x 2 x i8> splat (i8 18)
+;
+  %r = call <vscale x 2 x i8> @llvm.vector.partial.reduce.add.nxv2i8.nxv6i8(<vscale x 2 x i8> splat (i8 3), <vscale x 6 x i8> splat (i8 5))
+  ret <vscale x 2 x i8> %r
+}
+
+define <vscale x 2 x i8> @partial_reduce_add_scalable_wrap() {
+; CHECK-LABEL: @partial_reduce_add_scalable_wrap(
+; CHECK-NEXT:    ret <vscale x 2 x i8> splat (i8 123)
+;
+  %r = call <vscale x 2 x i8> @llvm.vector.partial.reduce.add.nxv2i8.nxv8i8(<vscale x 2 x i8> splat (i8 127), <vscale x 8 x i8> splat (i8 127))
+  ret <vscale x 2 x i8> %r
+}
+
+define <vscale x 2 x i1> @partial_reduce_add_scalable_i1() {
+; CHECK-LABEL: @partial_reduce_add_scalable_i1(
+; CHECK-NEXT:    ret <vscale x 2 x i1> splat (i1 true)
+;
+  %r = call <vscale x 2 x i1> @llvm.vector.partial.reduce.add.nxv2i1.nxv4i1(<vscale x 2 x i1> splat (i1 true), <vscale x 4 x i1> splat (i1 true))
+  ret <vscale x 2 x i1> %r
+}
+
+define <vscale x 2 x i8> @partial_reduce_add_scalable_poison() {
+; CHECK-LABEL: @partial_reduce_add_scalable_poison(
+; CHECK-NEXT:    ret <vscale x 2 x i8> poison
+;
+  %r = call <vscale x 2 x i8> @llvm.vector.partial.reduce.add.nxv2i8.nxv4i8(<vscale x 2 x i8> splat (i8 3), <vscale x 4 x i8> poison)
+  ret <vscale x 2 x i8> %r
+}
+
+define <2 x i8> @partial_reduce_add_mixed_nonzero() {
+; CHECK-LABEL: @partial_reduce_add_mixed_nonzero(
+; CHECK-NEXT:    [[R:%.*]] = call <2 x i8> @llvm.vector.partial.reduce.add.v2i8.nxv4i8(<2 x i8> <i8 1, i8 2>, <vscale x 4 x i8> splat (i8 3))
+; CHECK-NEXT:    ret <2 x i8> [[R]]
+;
+  %r = call <2 x i8> @llvm.vector.partial.reduce.add.v2i8.nxv4i8(<2 x i8> <i8 1, i8 2>, <vscale x 4 x i8> splat (i8 3))
+  ret <2 x i8> %r
+}
diff --git a/llvm/test/Transforms/InstSimplify/vector-partial-reduce-add.ll b/llvm/test/Transforms/InstSimplify/vector-partial-reduce-add.ll
new file mode 100644
index 0000000000000..60db23a26df9a
--- /dev/null
+++ b/llvm/test/Transforms/InstSimplify/vector-partial-reduce-add.ll
@@ -0,0 +1,40 @@
+; NOTE: Assertions have been autogenerated by utils/update_test_checks.py UTC_ARGS: --function-signature
+; RUN: opt < %s -passes=instsimplify -S | FileCheck %s
+; RUN: opt < %s -passes=instsimplify -use-constant-int-for-fixed-length-splat -use-constant-int-for-scalable-splat -S | FileCheck %s
+
+define <2 x i8> @zero_input(<2 x i8> %acc) {
+; CHECK-LABEL: define {{[^@]+}}@zero_input
+; CHECK-SAME: (<2 x i8> [[ACC:%.*]]) {
+; CHECK-NEXT:    ret <2 x i8> [[ACC]]
+;
+  %r = call <2 x i8> @llvm.vector.partial.reduce.add.v2i8.v4i8(<2 x i8> %acc, <4 x i8> zeroinitializer)
+  ret <2 x i8> %r
+}
+
+define <vscale x 2 x i8> @scalable_zero_input(<vscale x 2 x i8> %acc) {
+; CHECK-LABEL: define {{[^@]+}}@scalable_zero_input
+; CHECK-SAME: (<vscale x 2 x i8> [[ACC:%.*]]) {
+; CHECK-NEXT:    ret <vscale x 2 x i8> [[ACC]]
+;
+  %r = call <vscale x 2 x i8> @llvm.vector.partial.reduce.add.nxv2i8.nxv4i8(<vscale x 2 x i8> %acc, <vscale x 4 x i8> zeroinitializer)
+  ret <vscale x 2 x i8> %r
+}
+
+define <2 x i8> @mixed_zero_input(<2 x i8> %acc) {
+; CHECK-LABEL: define {{[^@]+}}@mixed_zero_input
+; CHECK-SAME: (<2 x i8> [[ACC:%.*]]) {
+; CHECK-NEXT:    ret <2 x i8> [[ACC]]
+;
+  %r = call <2 x i8> @llvm.vector.partial.reduce.add.v2i8.nxv4i8(<2 x i8> %acc, <vscale x 4 x i8> zeroinitializer)
+  ret <2 x i8> %r
+}
+
+define <vscale x 2 x i8> @scalable_nonzero_input(<vscale x 2 x i8> %acc) {
+; CHECK-LABEL: define {{[^@]+}}@scalable_nonzero_input
+; CHECK-SAME: (<vscale x 2 x i8> [[ACC:%.*]]) {
+; CHECK-NEXT:    [[R:%.*]] = call <vscale x 2 x i8> @llvm.vector.partial.reduce.add.nxv2i8.nxv4i8(<vscale x 2 x i8> [[ACC]], <vscale x 4 x i8> splat (i8 1))
+; CHECK-NEXT:    ret <vscale x 2 x i8> [[R]]
+;
+  %r = call <vscale x 2 x i8> @llvm.vector.partial.reduce.add.nxv2i8.nxv4i8(<vscale x 2 x i8> %acc, <vscale x 4 x i8> splat (i8 1))
+  ret <vscale x 2 x i8> %r
+}

``````````

</details>


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


More information about the llvm-commits mailing list