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

via llvm-commits llvm-commits at lists.llvm.org
Fri Jul 17 07:52:53 PDT 2026


llvmorg-github-actions[bot] wrote:


<!--LLVM PR SUMMARY COMMENT-->

@llvm/pr-subscribers-backend-risc-v

Author: Chennes (Chennesxu)

<details>
<summary>Changes</summary>

This PR implements constant folding for `ISD::PARTIAL_REDUCE_SMLA`, `ISD::PARTIAL_REDUCE_UMLA`, and `ISD::PARTIAL_REDUCE_SUMLA` when their operands are constant `BUILD_VECTOR`s. The fold computes the partial reduction using `APInt` arithmetic and returns a folded `BUILD_VECTOR`.

Input constants are truncated to their logical element width before being sign- or zero-extended as required by each opcode. Input lane `I` is accumulated into result lane `I % NumAccElts`, matching `TargetLowering::expandPartialReduceMLA`. The reduction order is deliberately unspecified, so this mapping is a valid refinement.

Poison is propagated only to the affected result lanes. Folding is skipped for undef and opaque constants. Unsupported operand forms return early because partial-reduce nodes have no scalar form and must not fall through to generic per-element folding.

The RISC-V test precommits the existing codegen. Direct SelectionDAG unit tests cover signedness, input-width, poison, and no-fold cases that cannot all be reliably represented through IR.

Fixes #<!-- -->209191

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


3 Files Affected:

- (modified) llvm/lib/CodeGen/SelectionDAG/SelectionDAG.cpp (+71) 
- (modified) llvm/test/CodeGen/RISCV/rvv/partial-reduction-add.ll (+31) 
- (modified) llvm/unittests/CodeGen/SelectionDAGNodeConstructionTest.cpp (+147-1) 


``````````diff
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 209d0b5149fd4..c98aad2b83ae8 100644
--- a/llvm/test/CodeGen/RISCV/rvv/partial-reduction-add.ll
+++ b/llvm/test/CodeGen/RISCV/rvv/partial-reduction-add.ll
@@ -76,3 +76,34 @@ entry:
   ret <vscale x 8 x i32> %partial.reduce
 }
 
+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:    vmv.v.x v9, a0
+; 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>,
+      <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
+}
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);
+}

``````````

</details>


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


More information about the llvm-commits mailing list