[llvm] [ConstantFolding] Fold vector.partial.reduce.add constants (PR #212112)

Jeong Jihyeon via llvm-commits llvm-commits at lists.llvm.org
Wed Sep 9 08:47:21 PDT 2026


https://github.com/JihyeonJeong129 updated https://github.com/llvm/llvm-project/pull/212112

>From f591178bd75b2491efaa2ab091d284f65e718387 Mon Sep 17 00:00:00 2001
From: Jihyeon Jeong <jh.jeong129 at gmail.com>
Date: Tue, 8 Sep 2026 13:21:36 +0000
Subject: [PATCH 1/4] [ConstantFolding] Add baseline tests for
 vector.partial.reduce.add

---
 .../InstSimplify/ConstProp/vecreduce.ll       | 97 +++++++++++++++++++
 1 file changed, 97 insertions(+)

diff --git a/llvm/test/Transforms/InstSimplify/ConstProp/vecreduce.ll b/llvm/test/Transforms/InstSimplify/ConstProp/vecreduce.ll
index 479b3f8ea4128..520b8906c9c2b 100644
--- a/llvm/test/Transforms/InstSimplify/ConstProp/vecreduce.ll
+++ b/llvm/test/Transforms/InstSimplify/ConstProp/vecreduce.ll
@@ -928,3 +928,100 @@ define i32 @umax_poison_elt() {
   %x = call i32 @llvm.vector.reduce.umax.v8i32(<8 x i32> <i32 1, i32 1, i32 poison, i32 1, i32 1, i32 poison, i32 1, i32 1>)
   ret i32 %x
 }
+
+define <4 x i32> @partial_reduce_add_constants() {
+; CHECK-LABEL: @partial_reduce_add_constants(
+; CHECK-NEXT:    [[X:%.*]] = call <4 x i32> @llvm.vector.partial.reduce.add.v4i32.v16i32(<4 x i32> <i32 100, i32 200, i32 300, i32 400>, <16 x i32> <i32 1, i32 2, i32 3, i32 4, i32 5, i32 6, i32 7, i32 8, i32 9, i32 10, i32 11, i32 12, i32 13, i32 14, i32 15, i32 16>)
+; CHECK-NEXT:    ret <4 x i32> [[X]]
+;
+  %x = call <4 x i32> @llvm.vector.partial.reduce.add.v4i32.v16i32(
+  <4 x i32> <i32 100, i32 200, i32 300, i32 400>,
+  <16 x i32> <i32 1, i32 2, i32 3, i32 4,
+  i32 5, i32 6, i32 7, i32 8,
+  i32 9, i32 10, i32 11, i32 12,
+  i32 13, i32 14, i32 15, i32 16>)
+  ret <4 x i32> %x
+}
+
+define <4 x i32> @partial_reduce_add_nonconstant_acc(<4 x i32> %acc) {
+; CHECK-LABEL: @partial_reduce_add_nonconstant_acc(
+; CHECK-NEXT:    [[X:%.*]] = call <4 x i32> @llvm.vector.partial.reduce.add.v4i32.v16i32(<4 x i32> [[ACC:%.*]], <16 x i32> <i32 1, i32 2, i32 3, i32 4, i32 5, i32 6, i32 7, i32 8, i32 9, i32 10, i32 11, i32 12, i32 13, i32 14, i32 15, i32 16>)
+; CHECK-NEXT:    ret <4 x i32> [[X]]
+;
+  %x = call <4 x i32> @llvm.vector.partial.reduce.add.v4i32.v16i32(
+  <4 x i32> %acc,
+  <16 x i32> <i32 1, i32 2, i32 3, i32 4,
+  i32 5, i32 6, i32 7, i32 8,
+  i32 9, i32 10, i32 11, i32 12,
+  i32 13, i32 14, i32 15, i32 16>)
+  ret <4 x i32> %x
+}
+
+define <4 x i32> @partial_reduce_add_nonconstant_input(<16 x i32> %input) {
+; CHECK-LABEL: @partial_reduce_add_nonconstant_input(
+; CHECK-NEXT:    [[X:%.*]] = call <4 x i32> @llvm.vector.partial.reduce.add.v4i32.v16i32(<4 x i32> <i32 100, i32 200, i32 300, i32 400>, <16 x i32> [[INPUT:%.*]])
+; CHECK-NEXT:    ret <4 x i32> [[X]]
+;
+  %x = call <4 x i32> @llvm.vector.partial.reduce.add.v4i32.v16i32(
+  <4 x i32> <i32 100, i32 200, i32 300, i32 400>,
+  <16 x i32> %input)
+  ret <4 x i32> %x
+}
+
+define <4 x i32> @partial_reduce_add_poison_element() {
+; CHECK-LABEL: @partial_reduce_add_poison_element(
+; CHECK-NEXT:    [[X:%.*]] = call <4 x i32> @llvm.vector.partial.reduce.add.v4i32.v16i32(<4 x i32> <i32 100, i32 200, i32 300, i32 400>, <16 x i32> <i32 1, i32 2, i32 3, i32 4, i32 5, i32 poison, i32 7, i32 8, i32 9, i32 10, i32 11, i32 12, i32 13, i32 14, i32 15, i32 16>)
+; CHECK-NEXT:    ret <4 x i32> [[X]]
+;
+  %x = call <4 x i32> @llvm.vector.partial.reduce.add.v4i32.v16i32(
+  <4 x i32> <i32 100, i32 200, i32 300, i32 400>,
+  <16 x i32> <i32 1, i32 2, i32 3, i32 4,
+  i32 5, i32 poison, i32 7, i32 8,
+  i32 9, i32 10, i32 11, i32 12,
+  i32 13, i32 14, i32 15, i32 16>)
+  ret <4 x i32> %x
+}
+
+define <4 x i32> @partial_reduce_add_ratio_one() {
+; CHECK-LABEL: @partial_reduce_add_ratio_one(
+; CHECK-NEXT:    [[X:%.*]] = call <4 x i32> @llvm.vector.partial.reduce.add.v4i32.v4i32(<4 x i32> <i32 100, i32 200, i32 300, i32 400>, <4 x i32> <i32 1, i32 2, i32 3, i32 4>)
+; CHECK-NEXT:    ret <4 x i32> [[X]]
+;
+  %x = call <4 x i32> @llvm.vector.partial.reduce.add.v4i32.v4i32(
+  <4 x i32> <i32 100, i32 200, i32 300, i32 400>,
+  <4 x i32> <i32 1, i32 2, i32 3, i32 4>)
+  ret <4 x i32> %x
+}
+
+define <2 x i32> @partial_reduce_add_ratio_two() {
+; CHECK-LABEL: @partial_reduce_add_ratio_two(
+; CHECK-NEXT:    [[X:%.*]] = call <2 x i32> @llvm.vector.partial.reduce.add.v2i32.v4i32(<2 x i32> <i32 100, i32 200>, <4 x i32> <i32 1, i32 2, i32 3, i32 4>)
+; CHECK-NEXT:    ret <2 x i32> [[X]]
+;
+  %x = call <2 x i32> @llvm.vector.partial.reduce.add.v2i32.v4i32(
+  <2 x i32> <i32 100, i32 200>,
+  <4 x i32> <i32 1, i32 2, i32 3, i32 4>)
+  ret <2 x i32> %x
+}
+
+define <2 x i32> @partial_reduce_add_negative() {
+; CHECK-LABEL: @partial_reduce_add_negative(
+; CHECK-NEXT:    [[X:%.*]] = call <2 x i32> @llvm.vector.partial.reduce.add.v2i32.v4i32(<2 x i32> <i32 -100, i32 -200>, <4 x i32> <i32 1, i32 2, i32 3, i32 4>)
+; CHECK-NEXT:    ret <2 x i32> [[X]]
+;
+  %x = call <2 x i32> @llvm.vector.partial.reduce.add.v2i32.v4i32(
+  <2 x i32> <i32 -100, i32 -200>,
+  <4 x i32> <i32 1, i32 2, i32 3, i32 4>)
+  ret <2 x i32> %x
+}
+
+define <2 x i8> @partial_reduce_add_wrap() {
+; CHECK-LABEL: @partial_reduce_add_wrap(
+; CHECK-NEXT:    [[X:%.*]] = call <2 x i8> @llvm.vector.partial.reduce.add.v2i8.v4i8(<2 x i8> <i8 127, i8 126>, <4 x i8> <i8 1, i8 1, i8 2, i8 4>)
+; CHECK-NEXT:    ret <2 x i8> [[X]]
+;
+  %x = call <2 x i8> @llvm.vector.partial.reduce.add.v2i8.v4i8(
+  <2 x i8> <i8 127, i8 126>,
+  <4 x i8> <i8 1, i8 1, i8 2, i8 4>)
+  ret <2 x i8> %x
+}

>From f20c89d9eb0b594e805104d7d8912fe4c2269c24 Mon Sep 17 00:00:00 2001
From: Jihyeon Jeong <jh.jeong129 at gmail.com>
Date: Tue, 8 Sep 2026 13:22:53 +0000
Subject: [PATCH 2/4] [ConstantFolding] Fold vector.partial.reduce.add
 constants

---
 llvm/lib/Analysis/ConstantFolding.cpp         | 47 +++++++++++++++++++
 .../InstSimplify/ConstProp/vecreduce.ll       | 18 +++----
 2 files changed, 53 insertions(+), 12 deletions(-)

diff --git a/llvm/lib/Analysis/ConstantFolding.cpp b/llvm/lib/Analysis/ConstantFolding.cpp
index 7eb44a398ebca..569f0a2459ead 100644
--- a/llvm/lib/Analysis/ConstantFolding.cpp
+++ b/llvm/lib/Analysis/ConstantFolding.cpp
@@ -1777,6 +1777,7 @@ static bool canConstantFoldIntrinsic(Intrinsic::ID ID, bool IsStrictFP) {
   case Intrinsic::vector_reduce_smax:
   case Intrinsic::vector_reduce_umin:
   case Intrinsic::vector_reduce_umax:
+  case Intrinsic::vector_partial_reduce_add:
   case Intrinsic::vector_extract:
   case Intrinsic::vector_insert:
   case Intrinsic::vector_interleave2:
@@ -2395,6 +2396,50 @@ Constant *constantFoldVectorReduce(Intrinsic::ID IID, Constant *Op) {
   return ConstantInt::get(Op->getContext(), Acc);
 }
 
+/// Fold a vector partial reduction add using the deterministic grouping
+/// chosen by TargetLowering::expandPartialReduceMLA. Although the
+/// LangRef leaves the grouping unspecified, input element I is accumulated
+/// into result lane I % NumAccElts, with each accumulator element seeding
+/// its corresponding result lane. Returns nullptr if any element cannot be
+/// folded.
+static Constant *constantFoldVectorPartialReduceAdd(Constant *Acc,
+                                                    Constant *Input,
+                                                    const DataLayout &DL) {
+  auto *AccTy = cast<FixedVectorType>(Acc->getType());
+  auto *InputTy = cast<FixedVectorType>(Input->getType());
+
+  unsigned NumAccElts = AccTy->getNumElements();
+  unsigned NumInputElts = InputTy->getNumElements();
+
+  SmallVector<Constant *> ResultElts;
+  ResultElts.reserve(NumAccElts);
+
+  for (unsigned I = 0; I < NumAccElts; ++I) {
+    Constant *AccElt = Acc->getAggregateElement(I);
+
+    if (!AccElt)
+      return nullptr;
+
+    ResultElts.push_back(AccElt);
+  }
+
+  for (unsigned I = 0; I < NumInputElts; ++I) {
+    Constant *InputElt = Input->getAggregateElement(I);
+    if (!InputElt)
+      return nullptr;
+
+    unsigned ResultIdx = I % NumAccElts;
+    Constant *Folded = ConstantFoldBinaryOpOperands(
+        Instruction::Add, ResultElts[ResultIdx], InputElt, DL);
+    if (!Folded)
+      return nullptr;
+
+    ResultElts[ResultIdx] = Folded;
+  }
+
+  return ConstantVector::get(ResultElts);
+}
+
 /// Attempt to fold an SSE floating point to integer conversion of a constant
 /// floating point. If roundTowardZero is false, the default IEEE rounding is
 /// used (toward nearest, ties to even). This matches the behavior of the
@@ -4435,6 +4480,8 @@ static Constant *ConstantFoldFixedVectorCall(
     }
     return ConstantVector::get(Result);
   }
+  case Intrinsic::vector_partial_reduce_add:
+    return constantFoldVectorPartialReduceAdd(Operands[0], Operands[1], DL);
   case Intrinsic::wasm_dot: {
     unsigned NumElements =
         cast<FixedVectorType>(Operands[0]->getType())->getNumElements();
diff --git a/llvm/test/Transforms/InstSimplify/ConstProp/vecreduce.ll b/llvm/test/Transforms/InstSimplify/ConstProp/vecreduce.ll
index 520b8906c9c2b..314d7165dce0f 100644
--- a/llvm/test/Transforms/InstSimplify/ConstProp/vecreduce.ll
+++ b/llvm/test/Transforms/InstSimplify/ConstProp/vecreduce.ll
@@ -931,8 +931,7 @@ define i32 @umax_poison_elt() {
 
 define <4 x i32> @partial_reduce_add_constants() {
 ; CHECK-LABEL: @partial_reduce_add_constants(
-; CHECK-NEXT:    [[X:%.*]] = call <4 x i32> @llvm.vector.partial.reduce.add.v4i32.v16i32(<4 x i32> <i32 100, i32 200, i32 300, i32 400>, <16 x i32> <i32 1, i32 2, i32 3, i32 4, i32 5, i32 6, i32 7, i32 8, i32 9, i32 10, i32 11, i32 12, i32 13, i32 14, i32 15, i32 16>)
-; CHECK-NEXT:    ret <4 x i32> [[X]]
+; CHECK-NEXT:    ret <4 x i32> <i32 128, i32 232, i32 336, i32 440>
 ;
   %x = call <4 x i32> @llvm.vector.partial.reduce.add.v4i32.v16i32(
   <4 x i32> <i32 100, i32 200, i32 300, i32 400>,
@@ -970,8 +969,7 @@ define <4 x i32> @partial_reduce_add_nonconstant_input(<16 x i32> %input) {
 
 define <4 x i32> @partial_reduce_add_poison_element() {
 ; CHECK-LABEL: @partial_reduce_add_poison_element(
-; CHECK-NEXT:    [[X:%.*]] = call <4 x i32> @llvm.vector.partial.reduce.add.v4i32.v16i32(<4 x i32> <i32 100, i32 200, i32 300, i32 400>, <16 x i32> <i32 1, i32 2, i32 3, i32 4, i32 5, i32 poison, i32 7, i32 8, i32 9, i32 10, i32 11, i32 12, i32 13, i32 14, i32 15, i32 16>)
-; CHECK-NEXT:    ret <4 x i32> [[X]]
+; CHECK-NEXT:    ret <4 x i32> <i32 128, i32 poison, i32 336, i32 440>
 ;
   %x = call <4 x i32> @llvm.vector.partial.reduce.add.v4i32.v16i32(
   <4 x i32> <i32 100, i32 200, i32 300, i32 400>,
@@ -984,8 +982,7 @@ define <4 x i32> @partial_reduce_add_poison_element() {
 
 define <4 x i32> @partial_reduce_add_ratio_one() {
 ; CHECK-LABEL: @partial_reduce_add_ratio_one(
-; CHECK-NEXT:    [[X:%.*]] = call <4 x i32> @llvm.vector.partial.reduce.add.v4i32.v4i32(<4 x i32> <i32 100, i32 200, i32 300, i32 400>, <4 x i32> <i32 1, i32 2, i32 3, i32 4>)
-; CHECK-NEXT:    ret <4 x i32> [[X]]
+; CHECK-NEXT:    ret <4 x i32> <i32 101, i32 202, i32 303, i32 404>
 ;
   %x = call <4 x i32> @llvm.vector.partial.reduce.add.v4i32.v4i32(
   <4 x i32> <i32 100, i32 200, i32 300, i32 400>,
@@ -995,8 +992,7 @@ define <4 x i32> @partial_reduce_add_ratio_one() {
 
 define <2 x i32> @partial_reduce_add_ratio_two() {
 ; CHECK-LABEL: @partial_reduce_add_ratio_two(
-; CHECK-NEXT:    [[X:%.*]] = call <2 x i32> @llvm.vector.partial.reduce.add.v2i32.v4i32(<2 x i32> <i32 100, i32 200>, <4 x i32> <i32 1, i32 2, i32 3, i32 4>)
-; CHECK-NEXT:    ret <2 x i32> [[X]]
+; CHECK-NEXT:    ret <2 x i32> <i32 104, i32 206>
 ;
   %x = call <2 x i32> @llvm.vector.partial.reduce.add.v2i32.v4i32(
   <2 x i32> <i32 100, i32 200>,
@@ -1006,8 +1002,7 @@ define <2 x i32> @partial_reduce_add_ratio_two() {
 
 define <2 x i32> @partial_reduce_add_negative() {
 ; CHECK-LABEL: @partial_reduce_add_negative(
-; CHECK-NEXT:    [[X:%.*]] = call <2 x i32> @llvm.vector.partial.reduce.add.v2i32.v4i32(<2 x i32> <i32 -100, i32 -200>, <4 x i32> <i32 1, i32 2, i32 3, i32 4>)
-; CHECK-NEXT:    ret <2 x i32> [[X]]
+; CHECK-NEXT:    ret <2 x i32> <i32 -96, i32 -194>
 ;
   %x = call <2 x i32> @llvm.vector.partial.reduce.add.v2i32.v4i32(
   <2 x i32> <i32 -100, i32 -200>,
@@ -1017,8 +1012,7 @@ define <2 x i32> @partial_reduce_add_negative() {
 
 define <2 x i8> @partial_reduce_add_wrap() {
 ; CHECK-LABEL: @partial_reduce_add_wrap(
-; CHECK-NEXT:    [[X:%.*]] = call <2 x i8> @llvm.vector.partial.reduce.add.v2i8.v4i8(<2 x i8> <i8 127, i8 126>, <4 x i8> <i8 1, i8 1, i8 2, i8 4>)
-; CHECK-NEXT:    ret <2 x i8> [[X]]
+; CHECK-NEXT:    ret <2 x i8> <i8 -126, i8 -125>
 ;
   %x = call <2 x i8> @llvm.vector.partial.reduce.add.v2i8.v4i8(
   <2 x i8> <i8 127, i8 126>,

>From 20fe19a532e05a1c8afe74c3bb7976d7127c13b8 Mon Sep 17 00:00:00 2001
From: Jihyeon Jeong <jh.jeong129 at gmail.com>
Date: Wed, 9 Sep 2026 15:46:16 +0000
Subject: [PATCH 3/4] [ConstantFolding] Add regression test for mixed
 fixed/scalable partial reduction

---
 .../Transforms/InstSimplify/ConstProp/vecreduce.ll     | 10 ++++++++++
 1 file changed, 10 insertions(+)

diff --git a/llvm/test/Transforms/InstSimplify/ConstProp/vecreduce.ll b/llvm/test/Transforms/InstSimplify/ConstProp/vecreduce.ll
index 314d7165dce0f..6c2ee928d7c0e 100644
--- a/llvm/test/Transforms/InstSimplify/ConstProp/vecreduce.ll
+++ b/llvm/test/Transforms/InstSimplify/ConstProp/vecreduce.ll
@@ -967,6 +967,16 @@ define <4 x i32> @partial_reduce_add_nonconstant_input(<16 x i32> %input) {
   ret <4 x i32> %x
 }
 
+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]]
+;
+  %x = call <4 x i32> @llvm.vector.partial.reduce.add.v4i32.nxv8i32(
+  <4 x i32> zeroinitializer, <vscale x 8 x i32> zeroinitializer)
+  ret <4 x i32> %x
+}
+
 define <4 x i32> @partial_reduce_add_poison_element() {
 ; CHECK-LABEL: @partial_reduce_add_poison_element(
 ; CHECK-NEXT:    ret <4 x i32> <i32 128, i32 poison, i32 336, i32 440>

>From d800f3aab150aa8aa7edc9484a4e3dfb97c34f75 Mon Sep 17 00:00:00 2001
From: Jihyeon Jeong <jh.jeong129 at gmail.com>
Date: Wed, 9 Sep 2026 15:47:00 +0000
Subject: [PATCH 4/4] [ConstantFolding] Avoid crash on scalable partial
 reduction input

---
 llvm/lib/Analysis/ConstantFolding.cpp | 5 ++++-
 1 file changed, 4 insertions(+), 1 deletion(-)

diff --git a/llvm/lib/Analysis/ConstantFolding.cpp b/llvm/lib/Analysis/ConstantFolding.cpp
index 569f0a2459ead..43d2bbc2002a6 100644
--- a/llvm/lib/Analysis/ConstantFolding.cpp
+++ b/llvm/lib/Analysis/ConstantFolding.cpp
@@ -2406,7 +2406,10 @@ static Constant *constantFoldVectorPartialReduceAdd(Constant *Acc,
                                                     Constant *Input,
                                                     const DataLayout &DL) {
   auto *AccTy = cast<FixedVectorType>(Acc->getType());
-  auto *InputTy = cast<FixedVectorType>(Input->getType());
+  // A fixed result type does not guarantee a fixed input type.
+  auto *InputTy = dyn_cast<FixedVectorType>(Input->getType());
+  if (!InputTy)
+    return nullptr;
 
   unsigned NumAccElts = AccTy->getNumElements();
   unsigned NumInputElts = InputTy->getNumElements();



More information about the llvm-commits mailing list