[llvm] [SelectionDAG] Fold constant PARTIAL_REDUCE_SMLA/UMLA/SUMLA nodes (PR #210351)
Craig Topper via llvm-commits
llvm-commits at lists.llvm.org
Tue Aug 4 22:01:17 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,
----------------
topperc wrote:
Plain C array?
https://github.com/llvm/llvm-project/pull/210351
More information about the llvm-commits
mailing list