[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