[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:19 PDT 2026


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

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

>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/2] [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/2] [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);
+}



More information about the llvm-commits mailing list