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

via llvm-commits llvm-commits at lists.llvm.org
Wed Aug 5 07:31:11 PDT 2026


================
@@ -357,3 +387,97 @@ TEST_F(SelectionDAGNodeConstructionTest, CTLS) {
   SDValue Ctlsi1 = DAG->getNode(ISD::CTLS, DL, MVT::i32, i1Op);
   EXPECT_TRUE(isNullConstant(Ctlsi1));
 }
+
+TEST_F(SelectionDAGNodeConstructionTest,
+       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});
+  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});
+}
+
+TEST_F(SelectionDAGNodeConstructionTest,
+       FoldConstantPartialReduceMLALegalizedType) {
+  SDLoc DL;
+  // v4i16 is a legal AArch64 vector type but i16 is not a legal scalar type, so
+  // after type legalization the folded constants must be created in the
+  // promoted (i32) type. Build the operands with their legalized scalar types
+  // before enabling the flag.
+  SDValue Acc = buildVector(MVT::v4i16, MVT::i32, DL, {100, -200, 300, -400});
+  SDValue LHS = buildVector(MVT::v8i8, MVT::i32, DL, {1, 2, 3, 4, 5, 6, 7, 8});
+  SDValue RHS =
+      buildVector(MVT::v8i8, MVT::i32, DL, {5, 6, 7, 8, 9, 10, 11, 12});
+
+  DAG->NewNodesMustHaveLegalTypes = true;
+  SDValue Result =
+      DAG->getNode(ISD::PARTIAL_REDUCE_SMLA, DL, MVT::v4i16, Acc, LHS, RHS);
+  ASSERT_EQ(Result.getOpcode(), ISD::BUILD_VECTOR);
+  // Signed lane values must survive promotion, so check the negative lanes too.
+  checkBuildVector(Result, {/*100 + 1*5 + 5*9=*/150, /*-200 + 2*6 + 6*10=*/-128,
+                            /*300 + 3*7 + 7*11=*/398,
+                            /*-400 + 4*8 + 8*12=*/-272});
+  for (SDValue Op : Result->op_values())
+    EXPECT_EQ(Op.getValueType(), MVT::i32);
+}
+
+TEST_F(SelectionDAGNodeConstructionTest,
+       FoldConstantPartialReduceMLARHSPoison) {
+  SDLoc DL;
+  SDValue Acc = buildVector(MVT::v2i32, DL, {100, 200});
+  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[2] = DAG->getPOISON(MVT::i8);
+  SDValue PoisonResult =
+      DAG->getNode(ISD::PARTIAL_REDUCE_SMLA, DL, MVT::v2i32, Acc, LHS,
+                   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), 244);
+}
+
+TEST_F(SelectionDAGNodeConstructionTest, DontFoldPartialReduceMLA) {
+  SDLoc DL;
+  SDValue Acc = buildVector(MVT::v2i32, DL, {100, 200});
+  SDValue LHS = buildVector(MVT::v4i8, DL, {1, 2, 3, 4});
+  SDValue RHS = buildVector(MVT::v4i8, DL, {5, 6, 7, 8});
+
+  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::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::v4i8, 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,
----------------
Chennesxu wrote:

Thanks for the review! All four addressed:

- Moved the fold below the `!VT.isVector()` early return.
- `AccEltBits >= InputEltBits` is guaranteed by the assert in `getNode`, so the extension can never truncate - switched to `sext`/`zext`. (The `.trunc(InputEltBits)` above is a separate thing: it handles narrow elements carried by wider promoted scalar operands.)
- Used the `SmallVector` range constructor - also cleaned up the same pattern in `DontFoldPartialReduceMLA`.
- `MixedLHS` is now a plain array.

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


More information about the llvm-commits mailing list