[llvm] [SelectionDAG] Fold constant PARTIAL_REDUCE_SMLA/UMLA/SUMLA nodes (PR #210351)

via llvm-commits llvm-commits at lists.llvm.org
Sat Jul 18 20:41:09 PDT 2026


https://github.com/Chennesxu updated https://github.com/llvm/llvm-project/pull/210351

>From b7078d6a1323258d76c12bdd9088877aff524d4a Mon Sep 17 00:00:00 2001
From: Chennes Xu <xuchen359 at gmail.com>
Date: Fri, 17 Jul 2026 22:19:07 +0800
Subject: [PATCH 1/3] [RISCV] Precommit tests for partial reduction constant
 folding

Add baseline RISC-V coverage for fixed-length constant operands and scalable constant splats in llvm.vector.partial.reduce.add.

The fixed-length case exposes the codegen change from the follow-up SelectionDAG fold. The scalable case ensures unsupported splats continue through normal lowering.
---
 .../RISCV/rvv/partial-reduction-add.ll        | 39 +++++++++++++++++++
 1 file changed, 39 insertions(+)

diff --git a/llvm/test/CodeGen/RISCV/rvv/partial-reduction-add.ll b/llvm/test/CodeGen/RISCV/rvv/partial-reduction-add.ll
index 209d0b5149fd4..9f7ce0f5e9f10 100644
--- a/llvm/test/CodeGen/RISCV/rvv/partial-reduction-add.ll
+++ b/llvm/test/CodeGen/RISCV/rvv/partial-reduction-add.ll
@@ -76,3 +76,42 @@ entry:
   ret <vscale x 8 x i32> %partial.reduce
 }
 
+; FIXME: Fold constant partial reductions.
+define <4 x i32> @partial_reduce_add_constants() {
+; CHECK-LABEL: partial_reduce_add_constants:
+; CHECK:       # %bb.0:
+; CHECK-NEXT:    vsetivli zero, 4, e32, m1, ta, ma
+; CHECK-NEXT:    vid.v v8
+; CHECK-NEXT:    li a0, 100
+; CHECK-NEXT:    vmv.v.x v9, a0
+; CHECK-NEXT:    vadd.vi v10, v8, 1
+; CHECK-NEXT:    vadd.vi v11, v8, 13
+; CHECK-NEXT:    vmacc.vx v9, a0, v8
+; CHECK-NEXT:    vadd.vv v10, v11, v10
+; CHECK-NEXT:    vadd.vi v11, v8, 9
+; CHECK-NEXT:    vadd.vi v8, v8, 5
+; CHECK-NEXT:    vadd.vv v9, v10, v9
+; CHECK-NEXT:    vadd.vv v8, v8, v11
+; CHECK-NEXT:    vadd.vv v8, v8, v9
+; CHECK-NEXT:    ret
+  %partial.reduce = call <4 x i32> @llvm.vector.partial.reduce.add(
+      <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> %partial.reduce
+}
+
+; Ensure scalable splats do not fall through to generic scalar folding.
+define <vscale x 4 x i32> @partial_reduce_add_nxv4i32_splat_constants() {
+; CHECK-LABEL: partial_reduce_add_nxv4i32_splat_constants:
+; CHECK:       # %bb.0:
+; CHECK-NEXT:    vsetvli a0, zero, e32, m2, ta, ma
+; CHECK-NEXT:    vmv.v.i v8, 3
+; CHECK-NEXT:    ret
+  %partial.reduce = call <vscale x 4 x i32> @llvm.vector.partial.reduce.add(
+      <vscale x 4 x i32> splat (i32 1),
+      <vscale x 4 x i32> splat (i32 2))
+  ret <vscale x 4 x i32> %partial.reduce
+}

>From 48c3fe581627fd154081a13c6f3d3a87956517c6 Mon Sep 17 00:00:00 2001
From: Chennes Xu <xuchen359 at gmail.com>
Date: Fri, 17 Jul 2026 22:28:22 +0800
Subject: [PATCH 2/3] [SelectionDAG] Fold constant
 PARTIAL_REDUCE_SMLA/UMLA/SUMLA nodes

Fold integer partial-reduce MLA nodes with BUILD_VECTOR operands containing integer constants or poison.

Truncate input constants to their logical element width before applying the signedness required by each opcode. Accumulate input lane I into result lane I % NumAccElts, matching TargetLowering::expandPartialReduceMLA, and perform the arithmetic at the accumulator element width.

Propagate poison only to affected result lanes, while leaving undef, opaque, and unsupported operand forms unfolded. Return early for unsupported forms because partial-reduce nodes have no scalar form and must not fall through to generic per-lane folding.

Add direct SelectionDAG tests for signedness, narrow vector elements represented by wider scalar constants, fixed-width wraparound, poison propagation, and no-fold cases.

Fixes #209191
---
 .../lib/CodeGen/SelectionDAG/SelectionDAG.cpp |  71 +++++++++
 .../RISCV/rvv/partial-reduction-add.ll        |  16 +-
 .../SelectionDAGNodeConstructionTest.cpp      | 148 +++++++++++++++++-
 3 files changed, 222 insertions(+), 13 deletions(-)

diff --git a/llvm/lib/CodeGen/SelectionDAG/SelectionDAG.cpp b/llvm/lib/CodeGen/SelectionDAG/SelectionDAG.cpp
index a8ae7927726ce..86fa1593f297b 100644
--- a/llvm/lib/CodeGen/SelectionDAG/SelectionDAG.cpp
+++ b/llvm/lib/CodeGen/SelectionDAG/SelectionDAG.cpp
@@ -7931,6 +7931,74 @@ SDValue SelectionDAG::FoldConstantArithmetic(unsigned Opcode, const SDLoc &DL,
     }
   }
 
+  // Constant fold integer partial reductions with constant BUILD_VECTOR
+  // operands. The reduction order is deliberately unspecified. Use the same
+  // subvector layout as TargetLowering::expandPartialReduceMLA(), where input
+  // lane I contributes to accumulator lane I % NumAccElts.
+  if (Opcode == ISD::PARTIAL_REDUCE_SMLA ||
+      Opcode == ISD::PARTIAL_REDUCE_UMLA ||
+      Opcode == ISD::PARTIAL_REDUCE_SUMLA) {
+    // These nodes have no scalar form, so unsupported cases must not fall
+    // through to generic per-lane vector folding.
+    if (!llvm::all_of(Ops, [](SDValue Op) {
+          return ISD::isBuildVectorOfConstantSDNodes(Op.getNode());
+        }))
+      return SDValue();
+
+    const unsigned AccEltBits = VT.getScalarSizeInBits();
+    const unsigned InputEltBits = Ops[1].getScalarValueSizeInBits();
+    const unsigned NumAccElts = VT.getVectorNumElements();
+    const unsigned NumInputElts = Ops[1].getValueType().getVectorNumElements();
+    SmallVector<APInt, 8> Results(NumAccElts, APInt::getZero(AccEltBits));
+    BitVector PoisonElts(NumAccElts);
+
+    for (unsigned I = 0; I != NumAccElts; ++I) {
+      SDValue Elt = Ops[0].getOperand(I);
+      if (Elt.getOpcode() == ISD::POISON) {
+        PoisonElts.set(I);
+        continue;
+      }
+      auto *C = dyn_cast<ConstantSDNode>(Elt);
+      if (!C || C->isOpaque())
+        return SDValue();
+      Results[I] = C->getAPIntValue().trunc(AccEltBits);
+    }
+
+    const bool IsLHSSigned = Opcode != ISD::PARTIAL_REDUCE_UMLA;
+    const bool IsRHSSigned = Opcode == ISD::PARTIAL_REDUCE_SMLA;
+    for (unsigned I = 0; I != NumInputElts; ++I) {
+      const unsigned AccIdx = I % NumAccElts;
+      SDValue LHSElt = Ops[1].getOperand(I);
+      SDValue RHSElt = Ops[2].getOperand(I);
+      if (LHSElt.getOpcode() == ISD::POISON ||
+          RHSElt.getOpcode() == ISD::POISON) {
+        PoisonElts.set(AccIdx);
+        continue;
+      }
+
+      auto *LHS = dyn_cast<ConstantSDNode>(LHSElt);
+      auto *RHS = dyn_cast<ConstantSDNode>(RHSElt);
+      if (!LHS || !RHS || LHS->isOpaque() || RHS->isOpaque())
+        return SDValue();
+
+      APInt LHSVal = LHS->getAPIntValue().trunc(InputEltBits);
+      APInt RHSVal = RHS->getAPIntValue().trunc(InputEltBits);
+      LHSVal = IsLHSSigned ? LHSVal.sextOrTrunc(AccEltBits)
+                           : LHSVal.zextOrTrunc(AccEltBits);
+      RHSVal = IsRHSSigned ? RHSVal.sextOrTrunc(AccEltBits)
+                           : RHSVal.zextOrTrunc(AccEltBits);
+      Results[AccIdx] += LHSVal * RHSVal;
+    }
+
+    EVT AccEltVT = VT.getVectorElementType();
+    SmallVector<SDValue, 8> ResultOps;
+    for (unsigned I = 0; I != NumAccElts; ++I)
+      ResultOps.push_back(PoisonElts[I]
+                              ? getPOISON(AccEltVT)
+                              : getConstant(Results[I], DL, AccEltVT));
+    return getBuildVector(VT, DL, ResultOps);
+  }
+
   // This is for vector folding only from here on.
   if (!VT.isVector())
     return SDValue();
@@ -9131,6 +9199,9 @@ SDValue SelectionDAG::getNode(unsigned Opcode, const SDLoc &DL, EVT VT,
 
   // Perform trivial constant folding for arithmetic operators.
   switch (Opcode) {
+  case ISD::PARTIAL_REDUCE_SMLA:
+  case ISD::PARTIAL_REDUCE_UMLA:
+  case ISD::PARTIAL_REDUCE_SUMLA:
   case ISD::FMA:
   case ISD::FMAD:
   case ISD::SETCC:
diff --git a/llvm/test/CodeGen/RISCV/rvv/partial-reduction-add.ll b/llvm/test/CodeGen/RISCV/rvv/partial-reduction-add.ll
index 9f7ce0f5e9f10..c98aad2b83ae8 100644
--- a/llvm/test/CodeGen/RISCV/rvv/partial-reduction-add.ll
+++ b/llvm/test/CodeGen/RISCV/rvv/partial-reduction-add.ll
@@ -76,23 +76,15 @@ entry:
   ret <vscale x 8 x i32> %partial.reduce
 }
 
-; FIXME: Fold constant partial reductions.
 define <4 x i32> @partial_reduce_add_constants() {
 ; CHECK-LABEL: partial_reduce_add_constants:
 ; CHECK:       # %bb.0:
+; CHECK-NEXT:    li a0, 128
 ; CHECK-NEXT:    vsetivli zero, 4, e32, m1, ta, ma
-; CHECK-NEXT:    vid.v v8
-; CHECK-NEXT:    li a0, 100
 ; CHECK-NEXT:    vmv.v.x v9, a0
-; CHECK-NEXT:    vadd.vi v10, v8, 1
-; CHECK-NEXT:    vadd.vi v11, v8, 13
-; CHECK-NEXT:    vmacc.vx v9, a0, v8
-; CHECK-NEXT:    vadd.vv v10, v11, v10
-; CHECK-NEXT:    vadd.vi v11, v8, 9
-; CHECK-NEXT:    vadd.vi v8, v8, 5
-; CHECK-NEXT:    vadd.vv v9, v10, v9
-; CHECK-NEXT:    vadd.vv v8, v8, v11
-; CHECK-NEXT:    vadd.vv v8, v8, v9
+; CHECK-NEXT:    vid.v v8
+; CHECK-NEXT:    li a0, 104
+; CHECK-NEXT:    vmadd.vx v8, a0, v9
 ; CHECK-NEXT:    ret
   %partial.reduce = call <4 x i32> @llvm.vector.partial.reduce.add(
       <4 x i32> <i32 100, i32 200, i32 300, i32 400>,
diff --git a/llvm/unittests/CodeGen/SelectionDAGNodeConstructionTest.cpp b/llvm/unittests/CodeGen/SelectionDAGNodeConstructionTest.cpp
index 0899b04bfddb8..7ae37341cac81 100644
--- a/llvm/unittests/CodeGen/SelectionDAGNodeConstructionTest.cpp
+++ b/llvm/unittests/CodeGen/SelectionDAGNodeConstructionTest.cpp
@@ -7,10 +7,40 @@
 //===----------------------------------------------------------------------===//
 
 #include "SelectionDAGTestBase.h"
+#include "llvm/ADT/ArrayRef.h"
+#include "llvm/ADT/SmallVector.h"
 
 using namespace llvm;
 
-class SelectionDAGNodeConstructionTest : public SelectionDAGTestBase {};
+class SelectionDAGNodeConstructionTest : public SelectionDAGTestBase {
+protected:
+  SDValue buildVector(EVT VT, EVT ScalarVT, const SDLoc &DL,
+                      ArrayRef<int64_t> Values) {
+    SmallVector<SDValue, 8> Elts;
+    for (int64_t Value : Values)
+      Elts.push_back(DAG->getConstant(
+          APInt(ScalarVT.getSizeInBits(), Value, /*isSigned=*/true), DL,
+          ScalarVT));
+    return DAG->getBuildVector(VT, DL, Elts);
+  }
+
+  SDValue buildVector(EVT VT, const SDLoc &DL, ArrayRef<int64_t> Values) {
+    return buildVector(VT, VT.getVectorElementType(), DL, Values);
+  }
+
+  void checkConstant(SDValue Value, int64_t Expected) {
+    auto *C = dyn_cast<ConstantSDNode>(Value);
+    ASSERT_NE(C, nullptr);
+    EXPECT_EQ(C->getSExtValue(), Expected);
+  }
+
+  void checkBuildVector(SDValue Result, ArrayRef<int64_t> Expected) {
+    ASSERT_EQ(Result.getOpcode(), ISD::BUILD_VECTOR);
+    ASSERT_EQ(Result.getNumOperands(), Expected.size());
+    for (unsigned I = 0; I != Expected.size(); ++I)
+      checkConstant(Result.getOperand(I), Expected[I]);
+  }
+};
 
 TEST_F(SelectionDAGNodeConstructionTest, ADD) {
   SDLoc DL;
@@ -357,3 +387,119 @@ TEST_F(SelectionDAGNodeConstructionTest, CTLS) {
   SDValue Ctlsi1 = DAG->getNode(ISD::CTLS, DL, MVT::i32, i1Op);
   EXPECT_TRUE(isNullConstant(Ctlsi1));
 }
+
+TEST_F(SelectionDAGNodeConstructionTest,
+       FoldConstantPartialReduceMLASignedness) {
+  SDLoc DL;
+  SDValue Acc = buildVector(MVT::v2i32, DL, {100, 200});
+  SDValue LHS = buildVector(MVT::v8i8, DL, {-1, 2, -3, 4, -5, 6, -7, 8});
+  SDValue RHS = buildVector(MVT::v8i8, DL, {1, -2, 3, -4, 5, -6, 7, -8});
+
+  checkBuildVector(
+      DAG->getNode(ISD::PARTIAL_REDUCE_SMLA, DL, MVT::v2i32, Acc, LHS, RHS),
+      {16, 80});
+  checkBuildVector(
+      DAG->getNode(ISD::PARTIAL_REDUCE_UMLA, DL, MVT::v2i32, Acc, LHS, RHS),
+      {4112, 5200});
+  checkBuildVector(
+      DAG->getNode(ISD::PARTIAL_REDUCE_SUMLA, DL, MVT::v2i32, Acc, LHS, RHS),
+      {16, 5200});
+}
+
+TEST_F(SelectionDAGNodeConstructionTest,
+       FoldConstantPartialReduceMLAWidthAndOverflow) {
+  SDLoc DL;
+  SDValue PromotedLHS = buildVector(MVT::v4i8, MVT::i32, DL, {255, 128, 0, 0});
+  SDValue PromotedRHS = buildVector(MVT::v4i8, MVT::i32, DL, {255, 255, 0, 0});
+  SDValue ZeroAcc = buildVector(MVT::v2i32, DL, {0, 0});
+
+  checkBuildVector(DAG->getNode(ISD::PARTIAL_REDUCE_SMLA, DL, MVT::v2i32,
+                                ZeroAcc, PromotedLHS, PromotedRHS),
+                   {1, 128});
+  checkBuildVector(DAG->getNode(ISD::PARTIAL_REDUCE_UMLA, DL, MVT::v2i32,
+                                ZeroAcc, PromotedLHS, PromotedRHS),
+                   {65025, 32640});
+  checkBuildVector(DAG->getNode(ISD::PARTIAL_REDUCE_SUMLA, DL, MVT::v2i32,
+                                ZeroAcc, PromotedLHS, PromotedRHS),
+                   {-255, -32640});
+
+  SDValue WrapAcc = buildVector(MVT::v2i32, DL, {INT32_MAX, INT32_MIN});
+  SDValue Ones = buildVector(MVT::v4i8, DL, {1, 1, 1, 1});
+  checkBuildVector(DAG->getNode(ISD::PARTIAL_REDUCE_SMLA, DL, MVT::v2i32,
+                                WrapAcc, Ones, Ones),
+                   {-2147483647, -2147483646});
+}
+
+TEST_F(SelectionDAGNodeConstructionTest, FoldConstantPartialReduceMLAPoison) {
+  SDLoc DL;
+  SDValue Acc = buildVector(MVT::v2i32, DL, {100, 200});
+  SDValue LHS = buildVector(MVT::v8i8, DL, {-1, 2, -3, 4, -5, 6, -7, 8});
+  SDValue RHS = buildVector(MVT::v8i8, DL, {1, -2, 3, -4, 5, -6, 7, -8});
+
+  SmallVector<SDValue, 8> PoisonLHS;
+  for (SDValue Elt : LHS->op_values())
+    PoisonLHS.push_back(Elt);
+  PoisonLHS[2] = DAG->getPOISON(MVT::i8);
+  SDValue PoisonResult =
+      DAG->getNode(ISD::PARTIAL_REDUCE_SMLA, DL, MVT::v2i32, Acc,
+                   DAG->getBuildVector(MVT::v8i8, DL, PoisonLHS), RHS);
+  ASSERT_EQ(PoisonResult.getOpcode(), ISD::BUILD_VECTOR);
+  EXPECT_EQ(PoisonResult.getOperand(0).getOpcode(), ISD::POISON);
+  checkConstant(PoisonResult.getOperand(1), 80);
+
+  SmallVector<SDValue, 8> PoisonRHS;
+  for (SDValue Elt : RHS->op_values())
+    PoisonRHS.push_back(Elt);
+  PoisonRHS[3] = DAG->getPOISON(MVT::i8);
+  PoisonResult =
+      DAG->getNode(ISD::PARTIAL_REDUCE_SMLA, DL, MVT::v2i32, Acc, LHS,
+                   DAG->getBuildVector(MVT::v8i8, DL, PoisonRHS));
+  ASSERT_EQ(PoisonResult.getOpcode(), ISD::BUILD_VECTOR);
+  checkConstant(PoisonResult.getOperand(0), 16);
+  EXPECT_EQ(PoisonResult.getOperand(1).getOpcode(), ISD::POISON);
+
+  SmallVector<SDValue, 2> PoisonAcc;
+  for (SDValue Elt : Acc->op_values())
+    PoisonAcc.push_back(Elt);
+  PoisonAcc[0] = DAG->getPOISON(MVT::i32);
+  PoisonResult =
+      DAG->getNode(ISD::PARTIAL_REDUCE_SMLA, DL, MVT::v2i32,
+                   DAG->getBuildVector(MVT::v2i32, DL, PoisonAcc), LHS, RHS);
+  ASSERT_EQ(PoisonResult.getOpcode(), ISD::BUILD_VECTOR);
+  EXPECT_EQ(PoisonResult.getOperand(0).getOpcode(), ISD::POISON);
+  checkConstant(PoisonResult.getOperand(1), 80);
+}
+
+TEST_F(SelectionDAGNodeConstructionTest, DontFoldPartialReduceMLA) {
+  SDLoc DL;
+  SDValue Acc = buildVector(MVT::v2i32, DL, {100, 200});
+  SDValue LHS = buildVector(MVT::v8i8, DL, {-1, 2, -3, 4, -5, 6, -7, 8});
+  SDValue RHS = buildVector(MVT::v8i8, DL, {1, -2, 3, -4, 5, -6, 7, -8});
+
+  SmallVector<SDValue, 8> SpecialLHS;
+  for (SDValue Elt : LHS->op_values())
+    SpecialLHS.push_back(Elt);
+  SpecialLHS[2] = DAG->getConstant(APInt(8, 1), DL, MVT::i8,
+                                   /*isTarget=*/false, /*isOpaque=*/true);
+  SDValue OpaqueResult =
+      DAG->getNode(ISD::PARTIAL_REDUCE_SMLA, DL, MVT::v2i32, Acc,
+                   DAG->getBuildVector(MVT::v8i8, DL, SpecialLHS), RHS);
+  EXPECT_EQ(OpaqueResult.getOpcode(), ISD::PARTIAL_REDUCE_SMLA);
+
+  SpecialLHS[2] = DAG->getUNDEF(MVT::i8);
+  SDValue UndefResult =
+      DAG->getNode(ISD::PARTIAL_REDUCE_SMLA, DL, MVT::v2i32, Acc,
+                   DAG->getBuildVector(MVT::v8i8, DL, SpecialLHS), RHS);
+  EXPECT_EQ(UndefResult.getOpcode(), ISD::PARTIAL_REDUCE_SMLA);
+
+  SDValue Variable = DAG->getCopyFromReg(DAG->getEntryNode(), DL,
+                                         Register::index2VirtReg(2), MVT::i32);
+  SmallVector<SDValue, 2> MixedLHS = {Variable,
+                                      DAG->getConstant(2, DL, MVT::i32)};
+  SDValue NonConstantResult =
+      DAG->getNode(ISD::PARTIAL_REDUCE_SMLA, DL, MVT::v2i32,
+                   buildVector(MVT::v2i32, DL, {1, 2}),
+                   DAG->getBuildVector(MVT::v2i32, DL, MixedLHS),
+                   buildVector(MVT::v2i32, DL, {3, 4}));
+  EXPECT_EQ(NonConstantResult.getOpcode(), ISD::PARTIAL_REDUCE_SMLA);
+}

>From a2f857f340e0eb43c5986647540d33a3e1617508 Mon Sep 17 00:00:00 2001
From: Chennes Xu <xuchen359 at gmail.com>
Date: Sun, 19 Jul 2026 11:23:16 +0800
Subject: [PATCH 3/3] Move partial reduction constant folding test coverage to
 IR where possible

---
 .../RISCV/rvv/partial-reduction-add.ll        | 61 ++++++++++++++
 .../SelectionDAGNodeConstructionTest.cpp      | 80 ++++---------------
 2 files changed, 75 insertions(+), 66 deletions(-)

diff --git a/llvm/test/CodeGen/RISCV/rvv/partial-reduction-add.ll b/llvm/test/CodeGen/RISCV/rvv/partial-reduction-add.ll
index c98aad2b83ae8..4c2c2e81bac82 100644
--- a/llvm/test/CodeGen/RISCV/rvv/partial-reduction-add.ll
+++ b/llvm/test/CodeGen/RISCV/rvv/partial-reduction-add.ll
@@ -95,6 +95,67 @@ define <4 x i32> @partial_reduce_add_constants() {
   ret <4 x i32> %partial.reduce
 }
 
+; Accumulation wraps at the accumulator element width.
+define <2 x i32> @partial_reduce_add_constants_wraparound() {
+; CHECK-LABEL: partial_reduce_add_constants_wraparound:
+; CHECK:       # %bb.0:
+; CHECK-NEXT:    vsetivli zero, 2, e32, mf2, ta, ma
+; CHECK-NEXT:    vid.v v8
+; CHECK-NEXT:    lui a0, 524288
+; CHECK-NEXT:    addi a0, a0, 1
+; CHECK-NEXT:    vadd.vx v8, v8, a0
+; CHECK-NEXT:    ret
+  %partial.reduce = call <2 x i32> @llvm.vector.partial.reduce.add(
+      <2 x i32> <i32 2147483647, i32 -2147483648>,
+      <4 x i32> <i32 1, i32 1, i32 1, i32 1>)
+  ret <2 x i32> %partial.reduce
+}
+
+; A poison input only affects its result lane; other lanes still fold.
+define i32 @partial_reduce_add_constants_input_poison() {
+; CHECK-LABEL: partial_reduce_add_constants_input_poison:
+; CHECK:       # %bb.0:
+; CHECK-NEXT:    li a0, 104
+; CHECK-NEXT:    ret
+  %partial.reduce = call <2 x i32> @llvm.vector.partial.reduce.add(
+      <2 x i32> <i32 100, i32 200>,
+      <4 x i32> <i32 1, i32 poison, i32 3, i32 4>)
+  %unaffected = extractelement <2 x i32> %partial.reduce, i64 0
+  ret i32 %unaffected
+}
+
+; A poison accumulator lane does not prevent other lanes from folding.
+define i32 @partial_reduce_add_constants_accumulator_poison() {
+; CHECK-LABEL: partial_reduce_add_constants_accumulator_poison:
+; CHECK:       # %bb.0:
+; CHECK-NEXT:    li a0, 206
+; CHECK-NEXT:    ret
+  %partial.reduce = call <2 x i32> @llvm.vector.partial.reduce.add(
+      <2 x i32> <i32 poison, i32 200>,
+      <4 x i32> <i32 1, i32 2, i32 3, i32 4>)
+  %unaffected = extractelement <2 x i32> %partial.reduce, i64 1
+  ret i32 %unaffected
+}
+
+; Undef deliberately leaves the partial reduction unfolded.
+define <2 x i32> @partial_reduce_add_constants_undef() {
+; CHECK-LABEL: partial_reduce_add_constants_undef:
+; CHECK:       # %bb.0:
+; CHECK-NEXT:    li a0, 100
+; CHECK-NEXT:    vsetivli zero, 2, e32, mf2, ta, ma
+; CHECK-NEXT:    vid.v v8
+; CHECK-NEXT:    vmv.v.x v9, a0
+; CHECK-NEXT:    vmacc.vx v9, a0, v8
+; CHECK-NEXT:    vadd.vi v9, v9, 1
+; CHECK-NEXT:    vadd.vi v8, v8, 3
+; CHECK-NEXT:    vadd.vv v8, v8, v9
+; CHECK-NEXT:    ret
+  %partial.reduce = call <2 x i32> @llvm.vector.partial.reduce.add(
+      <2 x i32> <i32 100, i32 200>,
+      <4 x i32> <i32 1, i32 undef, i32 3, i32 4>)
+  ret <2 x i32> %partial.reduce
+}
+
 ; Ensure scalable splats do not fall through to generic scalar folding.
 define <vscale x 4 x i32> @partial_reduce_add_nxv4i32_splat_constants() {
 ; CHECK-LABEL: partial_reduce_add_nxv4i32_splat_constants:
diff --git a/llvm/unittests/CodeGen/SelectionDAGNodeConstructionTest.cpp b/llvm/unittests/CodeGen/SelectionDAGNodeConstructionTest.cpp
index 7ae37341cac81..a14c00188d831 100644
--- a/llvm/unittests/CodeGen/SelectionDAGNodeConstructionTest.cpp
+++ b/llvm/unittests/CodeGen/SelectionDAGNodeConstructionTest.cpp
@@ -389,25 +389,7 @@ TEST_F(SelectionDAGNodeConstructionTest, CTLS) {
 }
 
 TEST_F(SelectionDAGNodeConstructionTest,
-       FoldConstantPartialReduceMLASignedness) {
-  SDLoc DL;
-  SDValue Acc = buildVector(MVT::v2i32, DL, {100, 200});
-  SDValue LHS = buildVector(MVT::v8i8, DL, {-1, 2, -3, 4, -5, 6, -7, 8});
-  SDValue RHS = buildVector(MVT::v8i8, DL, {1, -2, 3, -4, 5, -6, 7, -8});
-
-  checkBuildVector(
-      DAG->getNode(ISD::PARTIAL_REDUCE_SMLA, DL, MVT::v2i32, Acc, LHS, RHS),
-      {16, 80});
-  checkBuildVector(
-      DAG->getNode(ISD::PARTIAL_REDUCE_UMLA, DL, MVT::v2i32, Acc, LHS, RHS),
-      {4112, 5200});
-  checkBuildVector(
-      DAG->getNode(ISD::PARTIAL_REDUCE_SUMLA, DL, MVT::v2i32, Acc, LHS, RHS),
-      {16, 5200});
-}
-
-TEST_F(SelectionDAGNodeConstructionTest,
-       FoldConstantPartialReduceMLAWidthAndOverflow) {
+       FoldConstantPartialReduceMLASignednessAndInputWidth) {
   SDLoc DL;
   SDValue PromotedLHS = buildVector(MVT::v4i8, MVT::i32, DL, {255, 128, 0, 0});
   SDValue PromotedRHS = buildVector(MVT::v4i8, MVT::i32, DL, {255, 255, 0, 0});
@@ -422,76 +404,42 @@ TEST_F(SelectionDAGNodeConstructionTest,
   checkBuildVector(DAG->getNode(ISD::PARTIAL_REDUCE_SUMLA, DL, MVT::v2i32,
                                 ZeroAcc, PromotedLHS, PromotedRHS),
                    {-255, -32640});
-
-  SDValue WrapAcc = buildVector(MVT::v2i32, DL, {INT32_MAX, INT32_MIN});
-  SDValue Ones = buildVector(MVT::v4i8, DL, {1, 1, 1, 1});
-  checkBuildVector(DAG->getNode(ISD::PARTIAL_REDUCE_SMLA, DL, MVT::v2i32,
-                                WrapAcc, Ones, Ones),
-                   {-2147483647, -2147483646});
 }
 
-TEST_F(SelectionDAGNodeConstructionTest, FoldConstantPartialReduceMLAPoison) {
+TEST_F(SelectionDAGNodeConstructionTest,
+       FoldConstantPartialReduceMLARHSPoison) {
   SDLoc DL;
   SDValue Acc = buildVector(MVT::v2i32, DL, {100, 200});
-  SDValue LHS = buildVector(MVT::v8i8, DL, {-1, 2, -3, 4, -5, 6, -7, 8});
-  SDValue RHS = buildVector(MVT::v8i8, DL, {1, -2, 3, -4, 5, -6, 7, -8});
-
-  SmallVector<SDValue, 8> PoisonLHS;
-  for (SDValue Elt : LHS->op_values())
-    PoisonLHS.push_back(Elt);
-  PoisonLHS[2] = DAG->getPOISON(MVT::i8);
-  SDValue PoisonResult =
-      DAG->getNode(ISD::PARTIAL_REDUCE_SMLA, DL, MVT::v2i32, Acc,
-                   DAG->getBuildVector(MVT::v8i8, DL, PoisonLHS), RHS);
-  ASSERT_EQ(PoisonResult.getOpcode(), ISD::BUILD_VECTOR);
-  EXPECT_EQ(PoisonResult.getOperand(0).getOpcode(), ISD::POISON);
-  checkConstant(PoisonResult.getOperand(1), 80);
-
-  SmallVector<SDValue, 8> PoisonRHS;
+  SDValue LHS = buildVector(MVT::v4i8, DL, {1, 2, 3, 4});
+  SDValue RHS = buildVector(MVT::v4i8, DL, {5, 6, 7, 8});
+  SmallVector<SDValue, 4> PoisonRHS;
   for (SDValue Elt : RHS->op_values())
     PoisonRHS.push_back(Elt);
-  PoisonRHS[3] = DAG->getPOISON(MVT::i8);
-  PoisonResult =
+  PoisonRHS[2] = DAG->getPOISON(MVT::i8);
+  SDValue PoisonResult =
       DAG->getNode(ISD::PARTIAL_REDUCE_SMLA, DL, MVT::v2i32, Acc, LHS,
-                   DAG->getBuildVector(MVT::v8i8, DL, PoisonRHS));
-  ASSERT_EQ(PoisonResult.getOpcode(), ISD::BUILD_VECTOR);
-  checkConstant(PoisonResult.getOperand(0), 16);
-  EXPECT_EQ(PoisonResult.getOperand(1).getOpcode(), ISD::POISON);
-
-  SmallVector<SDValue, 2> PoisonAcc;
-  for (SDValue Elt : Acc->op_values())
-    PoisonAcc.push_back(Elt);
-  PoisonAcc[0] = DAG->getPOISON(MVT::i32);
-  PoisonResult =
-      DAG->getNode(ISD::PARTIAL_REDUCE_SMLA, DL, MVT::v2i32,
-                   DAG->getBuildVector(MVT::v2i32, DL, PoisonAcc), LHS, RHS);
+                   DAG->getBuildVector(MVT::v4i8, DL, PoisonRHS));
   ASSERT_EQ(PoisonResult.getOpcode(), ISD::BUILD_VECTOR);
   EXPECT_EQ(PoisonResult.getOperand(0).getOpcode(), ISD::POISON);
-  checkConstant(PoisonResult.getOperand(1), 80);
+  checkConstant(PoisonResult.getOperand(1), 244);
 }
 
 TEST_F(SelectionDAGNodeConstructionTest, DontFoldPartialReduceMLA) {
   SDLoc DL;
   SDValue Acc = buildVector(MVT::v2i32, DL, {100, 200});
-  SDValue LHS = buildVector(MVT::v8i8, DL, {-1, 2, -3, 4, -5, 6, -7, 8});
-  SDValue RHS = buildVector(MVT::v8i8, DL, {1, -2, 3, -4, 5, -6, 7, -8});
+  SDValue LHS = buildVector(MVT::v4i8, DL, {1, 2, 3, 4});
+  SDValue RHS = buildVector(MVT::v4i8, DL, {5, 6, 7, 8});
 
-  SmallVector<SDValue, 8> SpecialLHS;
+  SmallVector<SDValue, 4> SpecialLHS;
   for (SDValue Elt : LHS->op_values())
     SpecialLHS.push_back(Elt);
   SpecialLHS[2] = DAG->getConstant(APInt(8, 1), DL, MVT::i8,
                                    /*isTarget=*/false, /*isOpaque=*/true);
   SDValue OpaqueResult =
       DAG->getNode(ISD::PARTIAL_REDUCE_SMLA, DL, MVT::v2i32, Acc,
-                   DAG->getBuildVector(MVT::v8i8, DL, SpecialLHS), RHS);
+                   DAG->getBuildVector(MVT::v4i8, DL, SpecialLHS), RHS);
   EXPECT_EQ(OpaqueResult.getOpcode(), ISD::PARTIAL_REDUCE_SMLA);
 
-  SpecialLHS[2] = DAG->getUNDEF(MVT::i8);
-  SDValue UndefResult =
-      DAG->getNode(ISD::PARTIAL_REDUCE_SMLA, DL, MVT::v2i32, Acc,
-                   DAG->getBuildVector(MVT::v8i8, DL, SpecialLHS), RHS);
-  EXPECT_EQ(UndefResult.getOpcode(), ISD::PARTIAL_REDUCE_SMLA);
-
   SDValue Variable = DAG->getCopyFromReg(DAG->getEntryNode(), DL,
                                          Register::index2VirtReg(2), MVT::i32);
   SmallVector<SDValue, 2> MixedLHS = {Variable,



More information about the llvm-commits mailing list