[llvm] [ISel][AArch64] Add CodeGen support for partial sub reductions. (PR #186809)
Sander de Smalen via llvm-commits
llvm-commits at lists.llvm.org
Mon Apr 13 03:57:13 PDT 2026
https://github.com/sdesmalen-arm updated https://github.com/llvm/llvm-project/pull/186809
>From c5de37fdfe7e9e7a61d8e7e8af42adb5115282ba Mon Sep 17 00:00:00 2001
From: Sander de Smalen <sander.desmalen at arm.com>
Date: Mon, 2 Mar 2026 16:43:20 +0000
Subject: [PATCH 1/3] [ISel] Add CodeGen support for partial sub reductions.
---
llvm/include/llvm/CodeGen/ISDOpcodes.h | 4 +
llvm/include/llvm/CodeGen/TargetLowering.h | 2 +
.../include/llvm/Target/TargetSelectionDAG.td | 4 +
llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp | 45 +-
.../SelectionDAG/LegalizeIntegerTypes.cpp | 6 +
.../SelectionDAG/LegalizeVectorOps.cpp | 2 +
.../SelectionDAG/LegalizeVectorTypes.cpp | 4 +
.../lib/CodeGen/SelectionDAG/SelectionDAG.cpp | 2 +
.../SelectionDAG/SelectionDAGDumper.cpp | 4 +
.../Target/AArch64/AArch64ISelLowering.cpp | 68 +--
.../lib/Target/AArch64/AArch64SVEInstrInfo.td | 13 +
.../CodeGen/AArch64/partial-reduction-sub.ll | 402 ++++++++++++++++++
12 files changed, 521 insertions(+), 35 deletions(-)
create mode 100644 llvm/test/CodeGen/AArch64/partial-reduction-sub.ll
diff --git a/llvm/include/llvm/CodeGen/ISDOpcodes.h b/llvm/include/llvm/CodeGen/ISDOpcodes.h
index fa578f733d4e8..7f4e128eed904 100644
--- a/llvm/include/llvm/CodeGen/ISDOpcodes.h
+++ b/llvm/include/llvm/CodeGen/ISDOpcodes.h
@@ -1540,6 +1540,10 @@ enum NodeType {
PARTIAL_REDUCE_SUMLA, // sext, zext
PARTIAL_REDUCE_FMLA, // fpext, fpext
+ /// Similar to PARTIAL_REDUCE_[US]MLA, using a subtract instead of add.
+ PARTIAL_REDUCE_SMLS, // sext, sext
+ PARTIAL_REDUCE_UMLS, // zext, zext
+
/// The `llvm.experimental.stackmap` intrinsic.
/// Operands: input chain, glue, <id>, <numShadowBytes>, [live0[, live1...]]
/// Outputs: output chain, glue
diff --git a/llvm/include/llvm/CodeGen/TargetLowering.h b/llvm/include/llvm/CodeGen/TargetLowering.h
index 4b60c3f905120..34363a4d40a98 100644
--- a/llvm/include/llvm/CodeGen/TargetLowering.h
+++ b/llvm/include/llvm/CodeGen/TargetLowering.h
@@ -1683,6 +1683,7 @@ class LLVM_ABI TargetLoweringBase {
LegalizeAction getPartialReduceMLAAction(unsigned Opc, EVT AccVT,
EVT InputVT) const {
assert(Opc == ISD::PARTIAL_REDUCE_SMLA || Opc == ISD::PARTIAL_REDUCE_UMLA ||
+ Opc == ISD::PARTIAL_REDUCE_SMLS || Opc == ISD::PARTIAL_REDUCE_UMLS ||
Opc == ISD::PARTIAL_REDUCE_SUMLA || Opc == ISD::PARTIAL_REDUCE_FMLA);
PartialReduceActionTypes Key = {Opc, AccVT.getSimpleVT().SimpleTy,
InputVT.getSimpleVT().SimpleTy};
@@ -2799,6 +2800,7 @@ class LLVM_ABI TargetLoweringBase {
void setPartialReduceMLAAction(unsigned Opc, MVT AccVT, MVT InputVT,
LegalizeAction Action) {
assert(Opc == ISD::PARTIAL_REDUCE_SMLA || Opc == ISD::PARTIAL_REDUCE_UMLA ||
+ Opc == ISD::PARTIAL_REDUCE_SMLS || Opc == ISD::PARTIAL_REDUCE_UMLS ||
Opc == ISD::PARTIAL_REDUCE_SUMLA || Opc == ISD::PARTIAL_REDUCE_FMLA);
assert(AccVT.isValid() && InputVT.isValid() &&
"setPartialReduceMLAAction types aren't valid");
diff --git a/llvm/include/llvm/Target/TargetSelectionDAG.td b/llvm/include/llvm/Target/TargetSelectionDAG.td
index d689b3c1beda9..6bebf856e9b2b 100644
--- a/llvm/include/llvm/Target/TargetSelectionDAG.td
+++ b/llvm/include/llvm/Target/TargetSelectionDAG.td
@@ -556,6 +556,10 @@ def partial_reduce_umla : SDNode<"ISD::PARTIAL_REDUCE_UMLA",
SDTPartialReduceMLA>;
def partial_reduce_smla : SDNode<"ISD::PARTIAL_REDUCE_SMLA",
SDTPartialReduceMLA>;
+def partial_reduce_umls : SDNode<"ISD::PARTIAL_REDUCE_UMLS",
+ SDTPartialReduceMLA>;
+def partial_reduce_smls : SDNode<"ISD::PARTIAL_REDUCE_SMLS",
+ SDTPartialReduceMLA>;
def partial_reduce_sumla : SDNode<"ISD::PARTIAL_REDUCE_SUMLA",
SDTPartialReduceMLA>;
def partial_reduce_fmla : SDNode<"ISD::PARTIAL_REDUCE_FMLA",
diff --git a/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp b/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp
index 0f4503ae27998..65d11d653703e 100644
--- a/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp
+++ b/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp
@@ -2067,6 +2067,8 @@ SDValue DAGCombiner::visit(SDNode *N) {
case ISD::EXPERIMENTAL_VECTOR_HISTOGRAM: return visitMHISTOGRAM(N);
case ISD::PARTIAL_REDUCE_SMLA:
case ISD::PARTIAL_REDUCE_UMLA:
+ case ISD::PARTIAL_REDUCE_SMLS:
+ case ISD::PARTIAL_REDUCE_UMLS:
case ISD::PARTIAL_REDUCE_SUMLA:
case ISD::PARTIAL_REDUCE_FMLA:
return visitPARTIAL_REDUCE_MLA(N);
@@ -13407,6 +13409,14 @@ SDValue DAGCombiner::foldPartialReduceMLAMulOp(SDNode *N) {
Opc = Op1->getOpcode();
}
+ bool IsMLS = false;
+ if (Opc == ISD::SUB &&
+ ISD::isConstantSplatVectorAllZeros(Op1->getOperand(0).getNode())) {
+ Op1 = Op1->getOperand(1);
+ Opc = Op1->getOpcode();
+ IsMLS = true;
+ }
+
if (Opc != ISD::MUL && Opc != ISD::FMUL && Opc != ISD::SHL)
return SDValue();
@@ -13468,6 +13478,10 @@ SDValue DAGCombiner::foldPartialReduceMLAMulOp(SDNode *N) {
unsigned NewOpcode = LHSOpcode == ISD::SIGN_EXTEND
? ISD::PARTIAL_REDUCE_SMLA
: ISD::PARTIAL_REDUCE_UMLA;
+ if (IsMLS)
+ NewOpcode = NewOpcode == ISD::PARTIAL_REDUCE_SMLA
+ ? ISD::PARTIAL_REDUCE_SMLS
+ : ISD::PARTIAL_REDUCE_UMLS;
// Only perform these combines if the target supports folding
// the extends into the operation.
@@ -13491,18 +13505,22 @@ SDValue DAGCombiner::foldPartialReduceMLAMulOp(SDNode *N) {
unsigned NewOpc;
if (LHSOpcode == ISD::SIGN_EXTEND && RHSOpcode == ISD::SIGN_EXTEND)
- NewOpc = ISD::PARTIAL_REDUCE_SMLA;
+ NewOpc = IsMLS ? ISD::PARTIAL_REDUCE_SMLS : ISD::PARTIAL_REDUCE_SMLA;
else if (LHSOpcode == ISD::ZERO_EXTEND && RHSOpcode == ISD::ZERO_EXTEND)
- NewOpc = ISD::PARTIAL_REDUCE_UMLA;
- else if (LHSOpcode == ISD::SIGN_EXTEND && RHSOpcode == ISD::ZERO_EXTEND)
+ NewOpc = IsMLS ? ISD::PARTIAL_REDUCE_UMLS : ISD::PARTIAL_REDUCE_UMLA;
+ else if (!IsMLS && LHSOpcode == ISD::SIGN_EXTEND &&
+ RHSOpcode == ISD::ZERO_EXTEND)
NewOpc = ISD::PARTIAL_REDUCE_SUMLA;
- else if (LHSOpcode == ISD::ZERO_EXTEND && RHSOpcode == ISD::SIGN_EXTEND) {
+ else if (!IsMLS && LHSOpcode == ISD::ZERO_EXTEND &&
+ RHSOpcode == ISD::SIGN_EXTEND) {
NewOpc = ISD::PARTIAL_REDUCE_SUMLA;
std::swap(LHSExtOp, RHSExtOp);
- } else if (LHSOpcode == ISD::FP_EXTEND && RHSOpcode == ISD::FP_EXTEND) {
+ } else if (!IsMLS && LHSOpcode == ISD::FP_EXTEND &&
+ RHSOpcode == ISD::FP_EXTEND) {
NewOpc = ISD::PARTIAL_REDUCE_FMLA;
} else
return SDValue();
+
// For a 2-stage extend the signedness of both of the extends must match
// If the mul has the same type, there is no outer extend, and thus we
// can simply use the inner extends to pick the result node.
@@ -13548,6 +13566,14 @@ SDValue DAGCombiner::foldPartialReduceAdd(SDNode *N) {
Op1Opcode = Op1->getOpcode();
}
+ bool IsMLS = false;
+ if (Op1Opcode == ISD::SUB &&
+ ISD::isConstantSplatVectorAllZeros(Op1->getOperand(0).getNode())) {
+ Op1 = Op1->getOperand(1);
+ Op1Opcode = Op1->getOpcode();
+ IsMLS = true;
+ }
+
if (!ISD::isExtOpcode(Op1Opcode) && Op1Opcode != ISD::FP_EXTEND)
return SDValue();
@@ -13559,10 +13585,11 @@ SDValue DAGCombiner::foldPartialReduceAdd(SDNode *N) {
Op1.getValueType().getVectorElementType() != AccElemVT)
return SDValue();
- unsigned NewOpcode = N->getOpcode() == ISD::PARTIAL_REDUCE_FMLA
- ? ISD::PARTIAL_REDUCE_FMLA
- : Op1IsSigned ? ISD::PARTIAL_REDUCE_SMLA
- : ISD::PARTIAL_REDUCE_UMLA;
+ unsigned NewOpcode =
+ N->getOpcode() == ISD::PARTIAL_REDUCE_FMLA ? ISD::PARTIAL_REDUCE_FMLA
+ : Op1IsSigned
+ ? (IsMLS ? ISD::PARTIAL_REDUCE_SMLS : ISD::PARTIAL_REDUCE_SMLA)
+ : (IsMLS ? ISD::PARTIAL_REDUCE_UMLS : ISD::PARTIAL_REDUCE_UMLA);
SDValue UnextOp1 = Op1.getOperand(0);
EVT UnextOp1VT = UnextOp1.getValueType();
diff --git a/llvm/lib/CodeGen/SelectionDAG/LegalizeIntegerTypes.cpp b/llvm/lib/CodeGen/SelectionDAG/LegalizeIntegerTypes.cpp
index 4a27f804d6720..99f9f9f15c777 100644
--- a/llvm/lib/CodeGen/SelectionDAG/LegalizeIntegerTypes.cpp
+++ b/llvm/lib/CodeGen/SelectionDAG/LegalizeIntegerTypes.cpp
@@ -171,6 +171,8 @@ void DAGTypeLegalizer::PromoteIntegerResult(SDNode *N, unsigned ResNo) {
case ISD::PARTIAL_REDUCE_UMLA:
case ISD::PARTIAL_REDUCE_SMLA:
+ case ISD::PARTIAL_REDUCE_UMLS:
+ case ISD::PARTIAL_REDUCE_SMLS:
case ISD::PARTIAL_REDUCE_SUMLA:
Res = PromoteIntRes_PARTIAL_REDUCE_MLA(N);
break;
@@ -2168,6 +2170,8 @@ bool DAGTypeLegalizer::PromoteIntegerOperand(SDNode *N, unsigned OpNo) {
break;
case ISD::PARTIAL_REDUCE_UMLA:
case ISD::PARTIAL_REDUCE_SMLA:
+ case ISD::PARTIAL_REDUCE_UMLS:
+ case ISD::PARTIAL_REDUCE_SMLS:
case ISD::PARTIAL_REDUCE_SUMLA:
Res = PromoteIntOp_PARTIAL_REDUCE_MLA(N);
break;
@@ -3007,10 +3011,12 @@ SDValue DAGTypeLegalizer::PromoteIntOp_PARTIAL_REDUCE_MLA(SDNode *N) {
SmallVector<SDValue, 1> NewOps(N->ops());
switch (N->getOpcode()) {
case ISD::PARTIAL_REDUCE_SMLA:
+ case ISD::PARTIAL_REDUCE_SMLS:
NewOps[1] = SExtPromotedInteger(N->getOperand(1));
NewOps[2] = SExtPromotedInteger(N->getOperand(2));
break;
case ISD::PARTIAL_REDUCE_UMLA:
+ case ISD::PARTIAL_REDUCE_UMLS:
NewOps[1] = ZExtPromotedInteger(N->getOperand(1));
NewOps[2] = ZExtPromotedInteger(N->getOperand(2));
break;
diff --git a/llvm/lib/CodeGen/SelectionDAG/LegalizeVectorOps.cpp b/llvm/lib/CodeGen/SelectionDAG/LegalizeVectorOps.cpp
index c00fbe79c6d64..3b40da3bdea43 100644
--- a/llvm/lib/CodeGen/SelectionDAG/LegalizeVectorOps.cpp
+++ b/llvm/lib/CodeGen/SelectionDAG/LegalizeVectorOps.cpp
@@ -535,6 +535,8 @@ SDValue VectorLegalizer::LegalizeOp(SDValue Op) {
}
case ISD::PARTIAL_REDUCE_UMLA:
case ISD::PARTIAL_REDUCE_SMLA:
+ case ISD::PARTIAL_REDUCE_UMLS:
+ case ISD::PARTIAL_REDUCE_SMLS:
case ISD::PARTIAL_REDUCE_SUMLA:
case ISD::PARTIAL_REDUCE_FMLA:
Action =
diff --git a/llvm/lib/CodeGen/SelectionDAG/LegalizeVectorTypes.cpp b/llvm/lib/CodeGen/SelectionDAG/LegalizeVectorTypes.cpp
index 564bf3b7f152e..2ed728b9a65de 100644
--- a/llvm/lib/CodeGen/SelectionDAG/LegalizeVectorTypes.cpp
+++ b/llvm/lib/CodeGen/SelectionDAG/LegalizeVectorTypes.cpp
@@ -1531,6 +1531,8 @@ void DAGTypeLegalizer::SplitVectorResult(SDNode *N, unsigned ResNo) {
break;
case ISD::PARTIAL_REDUCE_UMLA:
case ISD::PARTIAL_REDUCE_SMLA:
+ case ISD::PARTIAL_REDUCE_UMLS:
+ case ISD::PARTIAL_REDUCE_SMLS:
case ISD::PARTIAL_REDUCE_SUMLA:
case ISD::PARTIAL_REDUCE_FMLA:
SplitVecRes_PARTIAL_REDUCE_MLA(N, Lo, Hi);
@@ -3751,6 +3753,8 @@ bool DAGTypeLegalizer::SplitVectorOperand(SDNode *N, unsigned OpNo) {
break;
case ISD::PARTIAL_REDUCE_UMLA:
case ISD::PARTIAL_REDUCE_SMLA:
+ case ISD::PARTIAL_REDUCE_UMLS:
+ case ISD::PARTIAL_REDUCE_SMLS:
case ISD::PARTIAL_REDUCE_SUMLA:
case ISD::PARTIAL_REDUCE_FMLA:
Res = SplitVecOp_PARTIAL_REDUCE_MLA(N);
diff --git a/llvm/lib/CodeGen/SelectionDAG/SelectionDAG.cpp b/llvm/lib/CodeGen/SelectionDAG/SelectionDAG.cpp
index 8e06325c3a8d5..475970a4c2d22 100644
--- a/llvm/lib/CodeGen/SelectionDAG/SelectionDAG.cpp
+++ b/llvm/lib/CodeGen/SelectionDAG/SelectionDAG.cpp
@@ -8664,6 +8664,8 @@ SDValue SelectionDAG::getNode(unsigned Opcode, const SDLoc &DL, EVT VT,
}
case ISD::PARTIAL_REDUCE_UMLA:
case ISD::PARTIAL_REDUCE_SMLA:
+ case ISD::PARTIAL_REDUCE_UMLS:
+ case ISD::PARTIAL_REDUCE_SMLS:
case ISD::PARTIAL_REDUCE_SUMLA:
case ISD::PARTIAL_REDUCE_FMLA: {
[[maybe_unused]] EVT AccVT = N1.getValueType();
diff --git a/llvm/lib/CodeGen/SelectionDAG/SelectionDAGDumper.cpp b/llvm/lib/CodeGen/SelectionDAG/SelectionDAGDumper.cpp
index 7161dd299f830..77c4af81d8637 100644
--- a/llvm/lib/CodeGen/SelectionDAG/SelectionDAGDumper.cpp
+++ b/llvm/lib/CodeGen/SelectionDAG/SelectionDAGDumper.cpp
@@ -604,8 +604,12 @@ std::string SDNode::getOperationName(const SelectionDAG *G) const {
case ISD::PARTIAL_REDUCE_UMLA:
return "partial_reduce_umla";
+ case ISD::PARTIAL_REDUCE_UMLS:
+ return "partial_reduce_umls";
case ISD::PARTIAL_REDUCE_SMLA:
return "partial_reduce_smla";
+ case ISD::PARTIAL_REDUCE_SMLS:
+ return "partial_reduce_smls";
case ISD::PARTIAL_REDUCE_SUMLA:
return "partial_reduce_sumla";
case ISD::PARTIAL_REDUCE_FMLA:
diff --git a/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp b/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
index c1a6654f68b2b..754d7a02d57e9 100644
--- a/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
+++ b/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
@@ -2015,12 +2015,11 @@ AArch64TargetLowering::AArch64TargetLowering(const TargetMachine &TM,
if (Subtarget->isSVEorStreamingSVEAvailable()) {
// Mark known legal pairs as 'Legal' (these will expand to UDOT or SDOT).
// Other pairs will default to 'Expand'.
- static const unsigned MLAOps[] = {ISD::PARTIAL_REDUCE_SMLA,
- ISD::PARTIAL_REDUCE_UMLA};
- setPartialReduceMLAAction(MLAOps, MVT::nxv2i64, MVT::nxv8i16, Legal);
- setPartialReduceMLAAction(MLAOps, MVT::nxv4i32, MVT::nxv16i8, Legal);
-
- setPartialReduceMLAAction(MLAOps, MVT::nxv2i64, MVT::nxv16i8, Custom);
+ static const unsigned DotMLAOps[] = {ISD::PARTIAL_REDUCE_SMLA,
+ ISD::PARTIAL_REDUCE_UMLA};
+ setPartialReduceMLAAction(DotMLAOps, MVT::nxv2i64, MVT::nxv8i16, Legal);
+ setPartialReduceMLAAction(DotMLAOps, MVT::nxv4i32, MVT::nxv16i8, Legal);
+ setPartialReduceMLAAction(DotMLAOps, MVT::nxv2i64, MVT::nxv16i8, Custom);
if (Subtarget->hasMatMulInt8()) {
setPartialReduceMLAAction(ISD::PARTIAL_REDUCE_SUMLA, MVT::nxv4i32,
@@ -2030,10 +2029,13 @@ AArch64TargetLowering::AArch64TargetLowering(const TargetMachine &TM,
}
if (Subtarget->hasSVE2() || Subtarget->hasSME()) {
+ static const unsigned MLALBTOps[] = {
+ ISD::PARTIAL_REDUCE_SMLA, ISD::PARTIAL_REDUCE_UMLA,
+ ISD::PARTIAL_REDUCE_SMLS, ISD::PARTIAL_REDUCE_UMLS};
// Wide add types
- setPartialReduceMLAAction(MLAOps, MVT::nxv2i64, MVT::nxv4i32, Legal);
- setPartialReduceMLAAction(MLAOps, MVT::nxv4i32, MVT::nxv8i16, Legal);
- setPartialReduceMLAAction(MLAOps, MVT::nxv8i16, MVT::nxv16i8, Legal);
+ setPartialReduceMLAAction(MLALBTOps, MVT::nxv2i64, MVT::nxv4i32, Legal);
+ setPartialReduceMLAAction(MLALBTOps, MVT::nxv4i32, MVT::nxv8i16, Legal);
+ setPartialReduceMLAAction(MLALBTOps, MVT::nxv8i16, MVT::nxv16i8, Legal);
setOperationAction(ISD::CLMUL, {MVT::nxv16i8, MVT::nxv4i32}, Legal);
}
@@ -2123,15 +2125,18 @@ AArch64TargetLowering::AArch64TargetLowering(const TargetMachine &TM,
setOperationAction(ISD::EXPERIMENTAL_VECTOR_HISTOGRAM, MVT::nxv2i64,
Custom);
- static const unsigned MLAOps[] = {ISD::PARTIAL_REDUCE_SMLA,
- ISD::PARTIAL_REDUCE_UMLA};
+ static const unsigned DotMLAOps[] = {ISD::PARTIAL_REDUCE_SMLA,
+ ISD::PARTIAL_REDUCE_UMLA};
+ static const unsigned MLALBTOps[] = {
+ ISD::PARTIAL_REDUCE_SMLA, ISD::PARTIAL_REDUCE_UMLA,
+ ISD::PARTIAL_REDUCE_SMLS, ISD::PARTIAL_REDUCE_UMLS};
// Must be lowered to SVE instructions.
- setPartialReduceMLAAction(MLAOps, MVT::v2i64, MVT::v4i32, Custom);
- setPartialReduceMLAAction(MLAOps, MVT::v2i64, MVT::v8i16, Custom);
- setPartialReduceMLAAction(MLAOps, MVT::v2i64, MVT::v16i8, Custom);
- setPartialReduceMLAAction(MLAOps, MVT::v4i32, MVT::v8i16, Custom);
- setPartialReduceMLAAction(MLAOps, MVT::v4i32, MVT::v16i8, Custom);
- setPartialReduceMLAAction(MLAOps, MVT::v8i16, MVT::v16i8, Custom);
+ setPartialReduceMLAAction(DotMLAOps, MVT::v2i64, MVT::v8i16, Custom);
+ setPartialReduceMLAAction(DotMLAOps, MVT::v2i64, MVT::v16i8, Custom);
+ setPartialReduceMLAAction(DotMLAOps, MVT::v4i32, MVT::v16i8, Custom);
+ setPartialReduceMLAAction(MLALBTOps, MVT::v2i64, MVT::v4i32, Custom);
+ setPartialReduceMLAAction(MLALBTOps, MVT::v4i32, MVT::v8i16, Custom);
+ setPartialReduceMLAAction(MLALBTOps, MVT::v8i16, MVT::v16i8, Custom);
}
}
@@ -2412,23 +2417,26 @@ void AArch64TargetLowering::addTypeForFixedLengthSVE(MVT VT) {
bool PreferNEON = VT.is64BitVector() || VT.is128BitVector();
bool PreferSVE = !PreferNEON && Subtarget->isSVEAvailable();
- static const unsigned MLAOps[] = {ISD::PARTIAL_REDUCE_SMLA,
- ISD::PARTIAL_REDUCE_UMLA};
+ static const unsigned MLALBTOps[] = {
+ ISD::PARTIAL_REDUCE_SMLA, ISD::PARTIAL_REDUCE_UMLA,
+ ISD::PARTIAL_REDUCE_SMLS, ISD::PARTIAL_REDUCE_UMLS};
+ static const unsigned DotMLAOps[] = {ISD::PARTIAL_REDUCE_SMLA,
+ ISD::PARTIAL_REDUCE_UMLA};
unsigned NumElts = VT.getVectorNumElements();
if (VT.getVectorElementType() == MVT::i64) {
- setPartialReduceMLAAction(MLAOps, VT,
+ setPartialReduceMLAAction(DotMLAOps, VT,
MVT::getVectorVT(MVT::i8, NumElts * 8), Custom);
- setPartialReduceMLAAction(MLAOps, VT,
+ setPartialReduceMLAAction(DotMLAOps, VT,
MVT::getVectorVT(MVT::i16, NumElts * 4), Custom);
- setPartialReduceMLAAction(MLAOps, VT,
+ setPartialReduceMLAAction(MLALBTOps, VT,
MVT::getVectorVT(MVT::i32, NumElts * 2), Custom);
} else if (VT.getVectorElementType() == MVT::i32) {
- setPartialReduceMLAAction(MLAOps, VT,
+ setPartialReduceMLAAction(DotMLAOps, VT,
MVT::getVectorVT(MVT::i8, NumElts * 4), Custom);
- setPartialReduceMLAAction(MLAOps, VT,
+ setPartialReduceMLAAction(MLALBTOps, VT,
MVT::getVectorVT(MVT::i16, NumElts * 2), Custom);
} else if (VT.getVectorElementType() == MVT::i16) {
- setPartialReduceMLAAction(MLAOps, VT,
+ setPartialReduceMLAAction(MLALBTOps, VT,
MVT::getVectorVT(MVT::i8, NumElts * 2), Custom);
}
if (Subtarget->hasMatMulInt8()) {
@@ -8477,6 +8485,8 @@ SDValue AArch64TargetLowering::LowerOperation(SDValue Op,
return LowerVECTOR_HISTOGRAM(Op, DAG);
case ISD::PARTIAL_REDUCE_SMLA:
case ISD::PARTIAL_REDUCE_UMLA:
+ case ISD::PARTIAL_REDUCE_SMLS:
+ case ISD::PARTIAL_REDUCE_UMLS:
case ISD::PARTIAL_REDUCE_SUMLA:
case ISD::PARTIAL_REDUCE_FMLA:
return LowerPARTIAL_REDUCE_MLA(Op, DAG);
@@ -32643,10 +32653,13 @@ AArch64TargetLowering::LowerPARTIAL_REDUCE_MLA(SDValue Op,
EVT ResultVT = Op.getValueType();
EVT OrigResultVT = ResultVT;
EVT OpVT = LHS.getValueType();
+ bool IsMLS = Op.getOpcode() == ISD::PARTIAL_REDUCE_UMLS ||
+ Op.getOpcode() == ISD::PARTIAL_REDUCE_SMLS;
// We can handle this case natively by accumulating into a wider
// zero-padded vector.
if (ResultVT == MVT::v2i32 && OpVT == MVT::v16i8) {
+ assert(!IsMLS && "Cannot handle this case for sub-reductions");
SDValue ZeroVec = DAG.getConstant(0, DL, MVT::v4i32);
SDValue WideAcc = DAG.getInsertSubvector(DL, ZeroVec, Acc, 0);
SDValue Wide =
@@ -32674,12 +32687,15 @@ AArch64TargetLowering::LowerPARTIAL_REDUCE_MLA(SDValue Op,
return ConvertToScalable ? convertFromScalableVector(DAG, OrigResultVT, Op)
: Op;
+ assert(!IsMLS && "i8 -> i64 custom sub-reductions are not supported");
+
EVT DotVT = ResultVT.isScalableVector() ? MVT::nxv4i32 : MVT::v4i32;
SDValue DotNode = DAG.getNode(Op.getOpcode(), DL, DotVT,
DAG.getConstant(0, DL, DotVT), LHS, RHS);
SDValue Res;
- bool IsUnsigned = Op.getOpcode() == ISD::PARTIAL_REDUCE_UMLA;
+ bool IsUnsigned = Op.getOpcode() == ISD::PARTIAL_REDUCE_UMLA ||
+ Op.getOpcode() == ISD::PARTIAL_REDUCE_UMLS;
if (Subtarget->hasSVE2() || Subtarget->isStreamingSVEAvailable()) {
unsigned LoOpcode = IsUnsigned ? AArch64ISD::UADDWB : AArch64ISD::SADDWB;
unsigned HiOpcode = IsUnsigned ? AArch64ISD::UADDWT : AArch64ISD::SADDWT;
diff --git a/llvm/lib/Target/AArch64/AArch64SVEInstrInfo.td b/llvm/lib/Target/AArch64/AArch64SVEInstrInfo.td
index 926593022b537..0c8ccd2c5acf9 100644
--- a/llvm/lib/Target/AArch64/AArch64SVEInstrInfo.td
+++ b/llvm/lib/Target/AArch64/AArch64SVEInstrInfo.td
@@ -3855,6 +3855,19 @@ let Predicates = [HasSVE2_or_SME] in {
defm UMLSLB_ZZZ : sve2_int_mla_long<0b10110, "umlslb", int_aarch64_sve_umlslb>;
defm UMLSLT_ZZZ : sve2_int_mla_long<0b10111, "umlslt", int_aarch64_sve_umlslt>;
+ def : Pat<(nxv2i64 (partial_reduce_umls nxv2i64:$Acc, nxv4i32:$LHS, nxv4i32:$RHS)),
+ (UMLSLT_ZZZ_D (UMLSLB_ZZZ_D $Acc, $LHS, $RHS), $LHS, $RHS)>;
+ def : Pat<(nxv2i64 (partial_reduce_smls nxv2i64:$Acc, nxv4i32:$LHS, nxv4i32:$RHS)),
+ (SMLSLT_ZZZ_D (SMLSLB_ZZZ_D $Acc, $LHS, $RHS), $LHS, $RHS)>;
+ def : Pat<(nxv4i32 (partial_reduce_umls nxv4i32:$Acc, nxv8i16:$LHS, nxv8i16:$RHS)),
+ (UMLSLT_ZZZ_S (UMLSLB_ZZZ_S $Acc, $LHS, $RHS), $LHS, $RHS)>;
+ def : Pat<(nxv4i32 (partial_reduce_smls nxv4i32:$Acc, nxv8i16:$LHS, nxv8i16:$RHS)),
+ (SMLSLT_ZZZ_S (SMLSLB_ZZZ_S $Acc, $LHS, $RHS), $LHS, $RHS)>;
+ def : Pat<(nxv8i16 (partial_reduce_umls nxv8i16:$Acc, nxv16i8:$LHS, nxv16i8:$RHS)),
+ (UMLSLT_ZZZ_H (UMLSLB_ZZZ_H $Acc, $LHS, $RHS), $LHS, $RHS)>;
+ def : Pat<(nxv8i16 (partial_reduce_smls nxv8i16:$Acc, nxv16i8:$LHS, nxv16i8:$RHS)),
+ (SMLSLT_ZZZ_H (SMLSLB_ZZZ_H $Acc, $LHS, $RHS), $LHS, $RHS)>;
+
def : Pat<(nxv2i64 (partial_reduce_umla nxv2i64:$Acc, nxv4i32:$LHS, nxv4i32:$RHS)),
(UMLALT_ZZZ_D (UMLALB_ZZZ_D $Acc, $LHS, $RHS), $LHS, $RHS)>;
def : Pat<(nxv2i64 (partial_reduce_smla nxv2i64:$Acc, nxv4i32:$LHS, nxv4i32:$RHS)),
diff --git a/llvm/test/CodeGen/AArch64/partial-reduction-sub.ll b/llvm/test/CodeGen/AArch64/partial-reduction-sub.ll
new file mode 100644
index 0000000000000..fd71767bcea14
--- /dev/null
+++ b/llvm/test/CodeGen/AArch64/partial-reduction-sub.ll
@@ -0,0 +1,402 @@
+; NOTE: Assertions have been autogenerated by utils/update_llc_test_checks.py UTC_ARGS: --version 6
+; RUN: llc < %s | FileCheck %s
+
+target triple = "aarch64"
+
+;
+; Ensure base cases are supported.
+;
+
+define <vscale x 8 x i16> @umlslbt_i8_i16(<vscale x 8 x i16> %acc, <vscale x 16 x i8> %a, <vscale x 16 x i8> %b) #0 {
+; CHECK-LABEL: umlslbt_i8_i16:
+; CHECK: // %bb.0:
+; CHECK-NEXT: umlslb z0.h, z1.b, z2.b
+; CHECK-NEXT: umlslt z0.h, z1.b, z2.b
+; CHECK-NEXT: ret
+ %a.zext = zext <vscale x 16 x i8> %a to <vscale x 16 x i16>
+ %b.zext = zext <vscale x 16 x i8> %b to <vscale x 16 x i16>
+ %mul = mul <vscale x 16 x i16> %a.zext, %b.zext
+ %mul.neg = sub <vscale x 16 x i16> zeroinitializer, %mul
+ %res = call <vscale x 8 x i16> @llvm.vector.partial.reduce.add(<vscale x 8 x i16> %acc, <vscale x 16 x i16> %mul.neg)
+ ret <vscale x 8 x i16> %res
+}
+
+define <vscale x 8 x i16> @smlslbt_i8_i16(<vscale x 8 x i16> %acc, <vscale x 16 x i8> %a, <vscale x 16 x i8> %b) #0 {
+; CHECK-LABEL: smlslbt_i8_i16:
+; CHECK: // %bb.0:
+; CHECK-NEXT: smlslb z0.h, z1.b, z2.b
+; CHECK-NEXT: smlslt z0.h, z1.b, z2.b
+; CHECK-NEXT: ret
+ %a.sext = sext <vscale x 16 x i8> %a to <vscale x 16 x i16>
+ %b.sext = sext <vscale x 16 x i8> %b to <vscale x 16 x i16>
+ %mul = mul <vscale x 16 x i16> %a.sext, %b.sext
+ %mul.neg = sub <vscale x 16 x i16> zeroinitializer, %mul
+ %res = call <vscale x 8 x i16> @llvm.vector.partial.reduce.add(<vscale x 8 x i16> %acc, <vscale x 16 x i16> %mul.neg)
+ ret <vscale x 8 x i16> %res
+}
+
+define <vscale x 4 x i32> @umlslbt_i16_i32(<vscale x 4 x i32> %acc, <vscale x 8 x i16> %a, <vscale x 8 x i16> %b) #0 {
+; CHECK-LABEL: umlslbt_i16_i32:
+; CHECK: // %bb.0:
+; CHECK-NEXT: umlslb z0.s, z1.h, z2.h
+; CHECK-NEXT: umlslt z0.s, z1.h, z2.h
+; CHECK-NEXT: ret
+ %a.zext = zext <vscale x 8 x i16> %a to <vscale x 8 x i32>
+ %b.zext = zext <vscale x 8 x i16> %b to <vscale x 8 x i32>
+ %mul = mul <vscale x 8 x i32> %a.zext, %b.zext
+ %mul.neg = sub <vscale x 8 x i32> zeroinitializer, %mul
+ %res = call <vscale x 4 x i32> @llvm.vector.partial.reduce.add(<vscale x 4 x i32> %acc, <vscale x 8 x i32> %mul.neg)
+ ret <vscale x 4 x i32> %res
+}
+
+define <vscale x 4 x i32> @smlslbt_i16_i32(<vscale x 4 x i32> %acc, <vscale x 8 x i16> %a, <vscale x 8 x i16> %b) #0 {
+; CHECK-LABEL: smlslbt_i16_i32:
+; CHECK: // %bb.0:
+; CHECK-NEXT: smlslb z0.s, z1.h, z2.h
+; CHECK-NEXT: smlslt z0.s, z1.h, z2.h
+; CHECK-NEXT: ret
+ %a.sext = sext <vscale x 8 x i16> %a to <vscale x 8 x i32>
+ %b.sext = sext <vscale x 8 x i16> %b to <vscale x 8 x i32>
+ %mul = mul <vscale x 8 x i32> %a.sext, %b.sext
+ %mul.neg = sub <vscale x 8 x i32> zeroinitializer, %mul
+ %res = call <vscale x 4 x i32> @llvm.vector.partial.reduce.add(<vscale x 4 x i32> %acc, <vscale x 8 x i32> %mul.neg)
+ ret <vscale x 4 x i32> %res
+}
+
+define <vscale x 2 x i64> @umlslbt_i32_i64(<vscale x 2 x i64> %acc, <vscale x 4 x i32> %a, <vscale x 4 x i32> %b) #0 {
+; CHECK-LABEL: umlslbt_i32_i64:
+; CHECK: // %bb.0:
+; CHECK-NEXT: umlslb z0.d, z1.s, z2.s
+; CHECK-NEXT: umlslt z0.d, z1.s, z2.s
+; CHECK-NEXT: ret
+ %a.zext = zext <vscale x 4 x i32> %a to <vscale x 4 x i64>
+ %b.zext = zext <vscale x 4 x i32> %b to <vscale x 4 x i64>
+ %mul = mul <vscale x 4 x i64> %a.zext, %b.zext
+ %mul.neg = sub <vscale x 4 x i64> zeroinitializer, %mul
+ %res = call <vscale x 2 x i64> @llvm.vector.partial.reduce.add(<vscale x 2 x i64> %acc, <vscale x 4 x i64> %mul.neg)
+ ret <vscale x 2 x i64> %res
+}
+
+define <vscale x 2 x i64> @smlslbt_i32_i64(<vscale x 2 x i64> %acc, <vscale x 4 x i32> %a, <vscale x 4 x i32> %b) #0 {
+; CHECK-LABEL: smlslbt_i32_i64:
+; CHECK: // %bb.0:
+; CHECK-NEXT: smlslb z0.d, z1.s, z2.s
+; CHECK-NEXT: smlslt z0.d, z1.s, z2.s
+; CHECK-NEXT: ret
+ %a.sext = sext <vscale x 4 x i32> %a to <vscale x 4 x i64>
+ %b.sext = sext <vscale x 4 x i32> %b to <vscale x 4 x i64>
+ %mul = mul <vscale x 4 x i64> %a.sext, %b.sext
+ %mul.neg = sub <vscale x 4 x i64> zeroinitializer, %mul
+ %res = call <vscale x 2 x i64> @llvm.vector.partial.reduce.add(<vscale x 2 x i64> %acc, <vscale x 4 x i64> %mul.neg)
+ ret <vscale x 2 x i64> %res
+}
+
+;
+; Ensure fixed-length codegen for streaming-compatible functions.
+;
+
+define <8 x i16> @fixed_umlslbt_i8_i16(<8 x i16> %acc, <16 x i8> %a, <16 x i8> %b) #0 "aarch64_pstate_sm_compatible" {
+; CHECK-LABEL: fixed_umlslbt_i8_i16:
+; CHECK: // %bb.0:
+; CHECK-NEXT: // kill: def $q0 killed $q0 def $z0
+; CHECK-NEXT: // kill: def $q2 killed $q2 def $z2
+; CHECK-NEXT: // kill: def $q1 killed $q1 def $z1
+; CHECK-NEXT: umlslb z0.h, z1.b, z2.b
+; CHECK-NEXT: umlslt z0.h, z1.b, z2.b
+; CHECK-NEXT: // kill: def $q0 killed $q0 killed $z0
+; CHECK-NEXT: ret
+ %a.zext = zext <16 x i8> %a to <16 x i16>
+ %b.zext = zext <16 x i8> %b to <16 x i16>
+ %mul = mul <16 x i16> %a.zext, %b.zext
+ %mul.neg = sub <16 x i16> zeroinitializer, %mul
+ %res = call <8 x i16> @llvm.vector.partial.reduce.add(<8 x i16> %acc, <16 x i16> %mul.neg)
+ ret <8 x i16> %res
+}
+
+define <8 x i16> @fixed_smlslbt_i8_i16(<8 x i16> %acc, <16 x i8> %a, <16 x i8> %b) #0 "aarch64_pstate_sm_compatible" {
+; CHECK-LABEL: fixed_smlslbt_i8_i16:
+; CHECK: // %bb.0:
+; CHECK-NEXT: // kill: def $q0 killed $q0 def $z0
+; CHECK-NEXT: // kill: def $q2 killed $q2 def $z2
+; CHECK-NEXT: // kill: def $q1 killed $q1 def $z1
+; CHECK-NEXT: smlslb z0.h, z1.b, z2.b
+; CHECK-NEXT: smlslt z0.h, z1.b, z2.b
+; CHECK-NEXT: // kill: def $q0 killed $q0 killed $z0
+; CHECK-NEXT: ret
+ %a.sext = sext <16 x i8> %a to <16 x i16>
+ %b.sext = sext <16 x i8> %b to <16 x i16>
+ %mul = mul <16 x i16> %a.sext, %b.sext
+ %mul.neg = sub <16 x i16> zeroinitializer, %mul
+ %res = call <8 x i16> @llvm.vector.partial.reduce.add(<8 x i16> %acc, <16 x i16> %mul.neg)
+ ret <8 x i16> %res
+}
+
+define <4 x i32> @fixed_umlslbt_i16_i32(<4 x i32> %acc, <8 x i16> %a, <8 x i16> %b) #0 "aarch64_pstate_sm_compatible" {
+; CHECK-LABEL: fixed_umlslbt_i16_i32:
+; CHECK: // %bb.0:
+; CHECK-NEXT: // kill: def $q0 killed $q0 def $z0
+; CHECK-NEXT: // kill: def $q2 killed $q2 def $z2
+; CHECK-NEXT: // kill: def $q1 killed $q1 def $z1
+; CHECK-NEXT: umlslb z0.s, z1.h, z2.h
+; CHECK-NEXT: umlslt z0.s, z1.h, z2.h
+; CHECK-NEXT: // kill: def $q0 killed $q0 killed $z0
+; CHECK-NEXT: ret
+ %a.zext = zext <8 x i16> %a to <8 x i32>
+ %b.zext = zext <8 x i16> %b to <8 x i32>
+ %mul = mul <8 x i32> %a.zext, %b.zext
+ %mul.neg = sub <8 x i32> zeroinitializer, %mul
+ %res = call <4 x i32> @llvm.vector.partial.reduce.add(<4 x i32> %acc, <8 x i32> %mul.neg)
+ ret <4 x i32> %res
+}
+
+define <4 x i32> @fixed_smlslbt_i16_i32(<4 x i32> %acc, <8 x i16> %a, <8 x i16> %b) #0 "aarch64_pstate_sm_compatible" {
+; CHECK-LABEL: fixed_smlslbt_i16_i32:
+; CHECK: // %bb.0:
+; CHECK-NEXT: // kill: def $q0 killed $q0 def $z0
+; CHECK-NEXT: // kill: def $q2 killed $q2 def $z2
+; CHECK-NEXT: // kill: def $q1 killed $q1 def $z1
+; CHECK-NEXT: smlslb z0.s, z1.h, z2.h
+; CHECK-NEXT: smlslt z0.s, z1.h, z2.h
+; CHECK-NEXT: // kill: def $q0 killed $q0 killed $z0
+; CHECK-NEXT: ret
+ %a.sext = sext <8 x i16> %a to <8 x i32>
+ %b.sext = sext <8 x i16> %b to <8 x i32>
+ %mul = mul <8 x i32> %a.sext, %b.sext
+ %mul.neg = sub <8 x i32> zeroinitializer, %mul
+ %res = call <4 x i32> @llvm.vector.partial.reduce.add(<4 x i32> %acc, <8 x i32> %mul.neg)
+ ret <4 x i32> %res
+}
+
+define <2 x i64> @fixed_umlslbt_i32_i64(<2 x i64> %acc, <4 x i32> %a, <4 x i32> %b) #0 "aarch64_pstate_sm_compatible" {
+; CHECK-LABEL: fixed_umlslbt_i32_i64:
+; CHECK: // %bb.0:
+; CHECK-NEXT: // kill: def $q0 killed $q0 def $z0
+; CHECK-NEXT: // kill: def $q2 killed $q2 def $z2
+; CHECK-NEXT: // kill: def $q1 killed $q1 def $z1
+; CHECK-NEXT: umlslb z0.d, z1.s, z2.s
+; CHECK-NEXT: umlslt z0.d, z1.s, z2.s
+; CHECK-NEXT: // kill: def $q0 killed $q0 killed $z0
+; CHECK-NEXT: ret
+ %a.zext = zext <4 x i32> %a to <4 x i64>
+ %b.zext = zext <4 x i32> %b to <4 x i64>
+ %mul = mul <4 x i64> %a.zext, %b.zext
+ %mul.neg = sub <4 x i64> zeroinitializer, %mul
+ %res = call <2 x i64> @llvm.vector.partial.reduce.add(<2 x i64> %acc, <4 x i64> %mul.neg)
+ ret <2 x i64> %res
+}
+
+define <2 x i64> @fixed_smlslbt_i32_i64(<2 x i64> %acc, <4 x i32> %a, <4 x i32> %b) #0 "aarch64_pstate_sm_compatible" {
+; CHECK-LABEL: fixed_smlslbt_i32_i64:
+; CHECK: // %bb.0:
+; CHECK-NEXT: // kill: def $q0 killed $q0 def $z0
+; CHECK-NEXT: // kill: def $q2 killed $q2 def $z2
+; CHECK-NEXT: // kill: def $q1 killed $q1 def $z1
+; CHECK-NEXT: smlslb z0.d, z1.s, z2.s
+; CHECK-NEXT: smlslt z0.d, z1.s, z2.s
+; CHECK-NEXT: // kill: def $q0 killed $q0 killed $z0
+; CHECK-NEXT: ret
+ %a.sext = sext <4 x i32> %a to <4 x i64>
+ %b.sext = sext <4 x i32> %b to <4 x i64>
+ %mul = mul <4 x i64> %a.sext, %b.sext
+ %mul.neg = sub <4 x i64> zeroinitializer, %mul
+ %res = call <2 x i64> @llvm.vector.partial.reduce.add(<2 x i64> %acc, <4 x i64> %mul.neg)
+ ret <2 x i64> %res
+}
+
+;
+; Test type legalisation for sub-reductions.
+;
+
+define <vscale x 8 x i16> @legalization_split_i8_i16(<vscale x 8 x i16> %acc, <vscale x 32 x i8> %a, <vscale x 32 x i8> %b) #0 {
+; CHECK-LABEL: legalization_split_i8_i16:
+; CHECK: // %bb.0:
+; CHECK-NEXT: umlslb z0.h, z1.b, z3.b
+; CHECK-NEXT: umlslt z0.h, z1.b, z3.b
+; CHECK-NEXT: umlslb z0.h, z2.b, z4.b
+; CHECK-NEXT: umlslt z0.h, z2.b, z4.b
+; CHECK-NEXT: ret
+ %a.zext = zext <vscale x 32 x i8> %a to <vscale x 32 x i16>
+ %b.zext = zext <vscale x 32 x i8> %b to <vscale x 32 x i16>
+ %mul = mul <vscale x 32 x i16> %a.zext, %b.zext
+ %mul.neg = sub <vscale x 32 x i16> zeroinitializer, %mul
+ %res = call <vscale x 8 x i16> @llvm.vector.partial.reduce.add(<vscale x 8 x i16> %acc, <vscale x 32 x i16> %mul.neg)
+ ret <vscale x 8 x i16> %res
+}
+
+define <vscale x 4 x i16> @legalization_promote_acc_i8_i16(<vscale x 4 x i16> %acc, <vscale x 16 x i8> %a, <vscale x 16 x i8> %b) #0 {
+; CHECK-LABEL: legalization_promote_acc_i8_i16:
+; CHECK: // %bb.0:
+; CHECK-NEXT: uunpklo z3.h, z1.b
+; CHECK-NEXT: uunpklo z4.h, z2.b
+; CHECK-NEXT: movi v5.2d, #0000000000000000
+; CHECK-NEXT: ptrue p0.h
+; CHECK-NEXT: uunpkhi z1.h, z1.b
+; CHECK-NEXT: uunpkhi z2.h, z2.b
+; CHECK-NEXT: msb z3.h, p0/m, z4.h, z5.h
+; CHECK-NEXT: msb z1.h, p0/m, z2.h, z5.h
+; CHECK-NEXT: uaddwb z0.s, z0.s, z3.h
+; CHECK-NEXT: uaddwt z0.s, z0.s, z3.h
+; CHECK-NEXT: uaddwb z0.s, z0.s, z1.h
+; CHECK-NEXT: uaddwt z0.s, z0.s, z1.h
+; CHECK-NEXT: ret
+ %a.zext = zext <vscale x 16 x i8> %a to <vscale x 16 x i16>
+ %b.zext = zext <vscale x 16 x i8> %b to <vscale x 16 x i16>
+ %mul = mul <vscale x 16 x i16> %a.zext, %b.zext
+ %mul.neg = sub <vscale x 16 x i16> zeroinitializer, %mul
+ %res = call <vscale x 4 x i16> @llvm.vector.partial.reduce.add(<vscale x 4 x i16> %acc, <vscale x 16 x i16> %mul.neg)
+ ret <vscale x 4 x i16> %res
+}
+
+define <vscale x 8 x i16> @legalization_promote_mul_ops_i8_i16(<vscale x 8 x i16> %acc, <vscale x 8 x i8> %a, <vscale x 8 x i8> %b) #0 {
+; CHECK-LABEL: legalization_promote_mul_ops_i8_i16:
+; CHECK: // %bb.0:
+; CHECK-NEXT: and z1.h, z1.h, #0xff
+; CHECK-NEXT: and z2.h, z2.h, #0xff
+; CHECK-NEXT: ptrue p0.h
+; CHECK-NEXT: mls z0.h, p0/m, z1.h, z2.h
+; CHECK-NEXT: ret
+ %a.zext = zext <vscale x 8 x i8> %a to <vscale x 8 x i16>
+ %b.zext = zext <vscale x 8 x i8> %b to <vscale x 8 x i16>
+ %mul = mul <vscale x 8 x i16> %a.zext, %b.zext
+ %mul.neg = sub <vscale x 8 x i16> zeroinitializer, %mul
+ %res = call <vscale x 8 x i16> @llvm.vector.partial.reduce.add(<vscale x 8 x i16> %acc, <vscale x 8 x i16> %mul.neg)
+ ret <vscale x 8 x i16> %res
+}
+
+; Test that MLSB/T are still generated when there is no 'mul'.
+define <vscale x 2 x i64> @extended_sub_i32_i64(<vscale x 2 x i64> %acc, <vscale x 4 x i32> %a, <vscale x 4 x i32> %b) #0 {
+; CHECK-LABEL: extended_sub_i32_i64:
+; CHECK: // %bb.0:
+; CHECK-NEXT: mov z2.s, #1 // =0x1
+; CHECK-NEXT: smlslb z0.d, z1.s, z2.s
+; CHECK-NEXT: smlslt z0.d, z1.s, z2.s
+; CHECK-NEXT: ret
+ %a.sext = sext <vscale x 4 x i32> %a to <vscale x 4 x i64>
+ %a.sext.neg = sub <vscale x 4 x i64> zeroinitializer, %a.sext
+ %res = call <vscale x 2 x i64> @llvm.vector.partial.reduce.add(<vscale x 2 x i64> %acc, <vscale x 4 x i64> %a.sext.neg)
+ ret <vscale x 2 x i64> %res
+}
+
+; Test that 'predication' is still supported.
+define <vscale x 2 x i64> @predicated_smlslbt_i32_i64(<vscale x 4 x i1> %pred, <vscale x 2 x i64> %acc, <vscale x 4 x i32> %a, <vscale x 4 x i32> %b) #0 {
+; CHECK-LABEL: predicated_smlslbt_i32_i64:
+; CHECK: // %bb.0:
+; CHECK-NEXT: movi v3.2d, #0000000000000000
+; CHECK-NEXT: sel z2.s, p0, z2.s, z3.s
+; CHECK-NEXT: smlslb z0.d, z1.s, z2.s
+; CHECK-NEXT: smlslt z0.d, z1.s, z2.s
+; CHECK-NEXT: ret
+ %a.sext = sext <vscale x 4 x i32> %a to <vscale x 4 x i64>
+ %b.sext = sext <vscale x 4 x i32> %b to <vscale x 4 x i64>
+ %mul = mul <vscale x 4 x i64> %a.sext, %b.sext
+ %mul.neg = sub <vscale x 4 x i64> zeroinitializer, %mul
+ %mul.neg.sel = select <vscale x 4 x i1> %pred, <vscale x 4 x i64> %mul.neg, <vscale x 4 x i64> zeroinitializer
+ %res = call <vscale x 2 x i64> @llvm.vector.partial.reduce.add(<vscale x 2 x i64> %acc, <vscale x 4 x i64> %mul.neg.sel)
+ ret <vscale x 2 x i64> %res
+}
+
+; Test that MLSB/T is not generated when the extends are mixed.
+define <vscale x 2 x i64> @negative_test_mixed_extends(<vscale x 2 x i64> %acc, <vscale x 4 x i32> %a, <vscale x 4 x i32> %b) #0 {
+; CHECK-LABEL: negative_test_mixed_extends:
+; CHECK: // %bb.0:
+; CHECK-NEXT: sunpklo z3.d, z1.s
+; CHECK-NEXT: uunpklo z4.d, z2.s
+; CHECK-NEXT: ptrue p0.d
+; CHECK-NEXT: sunpkhi z1.d, z1.s
+; CHECK-NEXT: uunpkhi z2.d, z2.s
+; CHECK-NEXT: mls z0.d, p0/m, z3.d, z4.d
+; CHECK-NEXT: mls z0.d, p0/m, z1.d, z2.d
+; CHECK-NEXT: ret
+ %a.sext = sext <vscale x 4 x i32> %a to <vscale x 4 x i64>
+ %b.zext = zext <vscale x 4 x i32> %b to <vscale x 4 x i64>
+ %mul = mul <vscale x 4 x i64> %a.sext, %b.zext
+ %mul.neg = sub <vscale x 4 x i64> zeroinitializer, %mul
+ %res = call <vscale x 2 x i64> @llvm.vector.partial.reduce.add(<vscale x 2 x i64> %acc, <vscale x 4 x i64> %mul.neg)
+ ret <vscale x 2 x i64> %res
+}
+
+; There is no sub dot-reduction, so we can't handle natively.
+define <vscale x 4 x i32> @negative_test_no_sub_dot_inst(<vscale x 4 x i32> %acc, <vscale x 16 x i8> %a, <vscale x 16 x i8> %b) #0 {
+; CHECK-LABEL: negative_test_no_sub_dot_inst:
+; CHECK: // %bb.0:
+; CHECK-NEXT: uunpklo z3.h, z1.b
+; CHECK-NEXT: uunpklo z4.h, z2.b
+; CHECK-NEXT: ptrue p0.s
+; CHECK-NEXT: uunpkhi z1.h, z1.b
+; CHECK-NEXT: uunpkhi z2.h, z2.b
+; CHECK-NEXT: uunpklo z5.s, z3.h
+; CHECK-NEXT: uunpklo z6.s, z4.h
+; CHECK-NEXT: uunpkhi z3.s, z3.h
+; CHECK-NEXT: uunpkhi z4.s, z4.h
+; CHECK-NEXT: mls z0.s, p0/m, z5.s, z6.s
+; CHECK-NEXT: uunpklo z5.s, z1.h
+; CHECK-NEXT: uunpklo z6.s, z2.h
+; CHECK-NEXT: uunpkhi z1.s, z1.h
+; CHECK-NEXT: uunpkhi z2.s, z2.h
+; CHECK-NEXT: mls z0.s, p0/m, z3.s, z4.s
+; CHECK-NEXT: mls z0.s, p0/m, z5.s, z6.s
+; CHECK-NEXT: mls z0.s, p0/m, z1.s, z2.s
+; CHECK-NEXT: ret
+ %a.zext = zext <vscale x 16 x i8> %a to <vscale x 16 x i32>
+ %b.zext = zext <vscale x 16 x i8> %b to <vscale x 16 x i32>
+ %mul = mul <vscale x 16 x i32> %a.zext, %b.zext
+ %mul.neg = sub <vscale x 16 x i32> zeroinitializer, %mul
+ %res = call <vscale x 4 x i32> @llvm.vector.partial.reduce.add(<vscale x 4 x i32> %acc, <vscale x 16 x i32> %mul.neg)
+ ret <vscale x 4 x i32> %res
+}
+
+; There exists FMLSLB/T instructions, but those are not yet supported.
+define <vscale x 4 x float> @negative_test_unsupported_fmlslbt(<vscale x 4 x float> %acc, <vscale x 8 x half> %a, <vscale x 8 x half> %b) #0 {
+; CHECK-LABEL: negative_test_unsupported_fmlslbt:
+; CHECK: // %bb.0:
+; CHECK-NEXT: uunpklo z3.s, z1.h
+; CHECK-NEXT: uunpklo z4.s, z2.h
+; CHECK-NEXT: ptrue p0.s
+; CHECK-NEXT: uunpkhi z1.s, z1.h
+; CHECK-NEXT: uunpkhi z2.s, z2.h
+; CHECK-NEXT: fcvt z3.s, p0/m, z3.h
+; CHECK-NEXT: fcvt z4.s, p0/m, z4.h
+; CHECK-NEXT: fcvt z1.s, p0/m, z1.h
+; CHECK-NEXT: fcvt z2.s, p0/m, z2.h
+; CHECK-NEXT: fmul z3.s, z3.s, z4.s
+; CHECK-NEXT: fmul z1.s, z1.s, z2.s
+; CHECK-NEXT: fneg z3.s, p0/m, z3.s
+; CHECK-NEXT: fneg z1.s, p0/m, z1.s
+; CHECK-NEXT: fadd z0.s, z0.s, z3.s
+; CHECK-NEXT: fadd z0.s, z0.s, z1.s
+; CHECK-NEXT: ret
+ %a.sext = fpext <vscale x 8 x half> %a to <vscale x 8 x float>
+ %b.sext = fpext <vscale x 8 x half> %b to <vscale x 8 x float>
+ %mul = fmul fast <vscale x 8 x float> %a.sext, %b.sext
+ %mul.neg = fsub fast <vscale x 8 x float> zeroinitializer, %mul
+ %res = call fast <vscale x 4 x float> @llvm.vector.partial.reduce.fadd(<vscale x 4 x float> %acc, <vscale x 8 x float> %mul.neg)
+ ret <vscale x 4 x float> %res
+}
+
+
+; Make sure wider types are supported when vscale_range supports it
+define void @wide_fixed_umlslbt_i8_i16(ptr %acc.ptr, ptr %a.ptr, ptr %b.ptr, ptr %dest.ptr) #0 vscale_range(2,0) {
+; CHECK-LABEL: wide_fixed_umlslbt_i8_i16:
+; CHECK: // %bb.0:
+; CHECK-NEXT: ptrue p0.b, vl32
+; CHECK-NEXT: ptrue p1.h, vl16
+; CHECK-NEXT: ld1b { z0.b }, p0/z, [x1]
+; CHECK-NEXT: ld1b { z1.b }, p0/z, [x2]
+; CHECK-NEXT: ld1h { z2.h }, p1/z, [x0]
+; CHECK-NEXT: umlslb z2.h, z0.b, z1.b
+; CHECK-NEXT: umlslt z2.h, z0.b, z1.b
+; CHECK-NEXT: st1h { z2.h }, p1, [x3]
+; CHECK-NEXT: ret
+ %a = load <32 x i8>, ptr %a.ptr
+ %b = load <32 x i8>, ptr %b.ptr
+ %acc = load <16 x i16>, ptr %acc.ptr
+ %a.zext = zext <32 x i8> %a to <32 x i16>
+ %b.zext = zext <32 x i8> %b to <32 x i16>
+ %mul = mul <32 x i16> %a.zext, %b.zext
+ %mul.neg = sub <32 x i16> zeroinitializer, %mul
+ %res = call <16 x i16> @llvm.vector.partial.reduce.add(<16 x i16> %acc, <32 x i16> %mul.neg)
+ store <16 x i16> %res, ptr %dest.ptr
+ ret void
+}
+
+attributes #0 = { "target-features"="+sve2" }
>From e8016dabb3ca8905ca8aee6e085d139887d5f46c Mon Sep 17 00:00:00 2001
From: Sander de Smalen <sander.desmalen at arm.com>
Date: Fri, 10 Apr 2026 14:33:04 +0000
Subject: [PATCH 2/3] Implement sub-reductions by inversing acc/result
---
llvm/include/llvm/CodeGen/ISDOpcodes.h | 4 -
llvm/include/llvm/CodeGen/TargetLowering.h | 2 -
.../include/llvm/Target/TargetSelectionDAG.td | 4 -
llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp | 55 +++++++------
.../SelectionDAG/LegalizeIntegerTypes.cpp | 6 --
.../SelectionDAG/LegalizeVectorOps.cpp | 2 -
.../SelectionDAG/LegalizeVectorTypes.cpp | 4 -
.../lib/CodeGen/SelectionDAG/SelectionDAG.cpp | 2 -
.../SelectionDAG/SelectionDAGDumper.cpp | 4 -
.../Target/AArch64/AArch64ISelLowering.cpp | 68 ++++++---------
.../lib/Target/AArch64/AArch64SVEInstrInfo.td | 15 ++--
.../CodeGen/AArch64/partial-reduction-sub.ll | 82 +++----------------
12 files changed, 79 insertions(+), 169 deletions(-)
diff --git a/llvm/include/llvm/CodeGen/ISDOpcodes.h b/llvm/include/llvm/CodeGen/ISDOpcodes.h
index 7f4e128eed904..fa578f733d4e8 100644
--- a/llvm/include/llvm/CodeGen/ISDOpcodes.h
+++ b/llvm/include/llvm/CodeGen/ISDOpcodes.h
@@ -1540,10 +1540,6 @@ enum NodeType {
PARTIAL_REDUCE_SUMLA, // sext, zext
PARTIAL_REDUCE_FMLA, // fpext, fpext
- /// Similar to PARTIAL_REDUCE_[US]MLA, using a subtract instead of add.
- PARTIAL_REDUCE_SMLS, // sext, sext
- PARTIAL_REDUCE_UMLS, // zext, zext
-
/// The `llvm.experimental.stackmap` intrinsic.
/// Operands: input chain, glue, <id>, <numShadowBytes>, [live0[, live1...]]
/// Outputs: output chain, glue
diff --git a/llvm/include/llvm/CodeGen/TargetLowering.h b/llvm/include/llvm/CodeGen/TargetLowering.h
index 34363a4d40a98..4b60c3f905120 100644
--- a/llvm/include/llvm/CodeGen/TargetLowering.h
+++ b/llvm/include/llvm/CodeGen/TargetLowering.h
@@ -1683,7 +1683,6 @@ class LLVM_ABI TargetLoweringBase {
LegalizeAction getPartialReduceMLAAction(unsigned Opc, EVT AccVT,
EVT InputVT) const {
assert(Opc == ISD::PARTIAL_REDUCE_SMLA || Opc == ISD::PARTIAL_REDUCE_UMLA ||
- Opc == ISD::PARTIAL_REDUCE_SMLS || Opc == ISD::PARTIAL_REDUCE_UMLS ||
Opc == ISD::PARTIAL_REDUCE_SUMLA || Opc == ISD::PARTIAL_REDUCE_FMLA);
PartialReduceActionTypes Key = {Opc, AccVT.getSimpleVT().SimpleTy,
InputVT.getSimpleVT().SimpleTy};
@@ -2800,7 +2799,6 @@ class LLVM_ABI TargetLoweringBase {
void setPartialReduceMLAAction(unsigned Opc, MVT AccVT, MVT InputVT,
LegalizeAction Action) {
assert(Opc == ISD::PARTIAL_REDUCE_SMLA || Opc == ISD::PARTIAL_REDUCE_UMLA ||
- Opc == ISD::PARTIAL_REDUCE_SMLS || Opc == ISD::PARTIAL_REDUCE_UMLS ||
Opc == ISD::PARTIAL_REDUCE_SUMLA || Opc == ISD::PARTIAL_REDUCE_FMLA);
assert(AccVT.isValid() && InputVT.isValid() &&
"setPartialReduceMLAAction types aren't valid");
diff --git a/llvm/include/llvm/Target/TargetSelectionDAG.td b/llvm/include/llvm/Target/TargetSelectionDAG.td
index 6bebf856e9b2b..d689b3c1beda9 100644
--- a/llvm/include/llvm/Target/TargetSelectionDAG.td
+++ b/llvm/include/llvm/Target/TargetSelectionDAG.td
@@ -556,10 +556,6 @@ def partial_reduce_umla : SDNode<"ISD::PARTIAL_REDUCE_UMLA",
SDTPartialReduceMLA>;
def partial_reduce_smla : SDNode<"ISD::PARTIAL_REDUCE_SMLA",
SDTPartialReduceMLA>;
-def partial_reduce_umls : SDNode<"ISD::PARTIAL_REDUCE_UMLS",
- SDTPartialReduceMLA>;
-def partial_reduce_smls : SDNode<"ISD::PARTIAL_REDUCE_SMLS",
- SDTPartialReduceMLA>;
def partial_reduce_sumla : SDNode<"ISD::PARTIAL_REDUCE_SUMLA",
SDTPartialReduceMLA>;
def partial_reduce_fmla : SDNode<"ISD::PARTIAL_REDUCE_FMLA",
diff --git a/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp b/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp
index 65d11d653703e..270427f9cc03e 100644
--- a/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp
+++ b/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp
@@ -2067,8 +2067,6 @@ SDValue DAGCombiner::visit(SDNode *N) {
case ISD::EXPERIMENTAL_VECTOR_HISTOGRAM: return visitMHISTOGRAM(N);
case ISD::PARTIAL_REDUCE_SMLA:
case ISD::PARTIAL_REDUCE_UMLA:
- case ISD::PARTIAL_REDUCE_SMLS:
- case ISD::PARTIAL_REDUCE_UMLS:
case ISD::PARTIAL_REDUCE_SUMLA:
case ISD::PARTIAL_REDUCE_FMLA:
return visitPARTIAL_REDUCE_MLA(N);
@@ -13464,6 +13462,19 @@ SDValue DAGCombiner::foldPartialReduceMLAMulOp(SDNode *N) {
}
};
+ // Generate an MLA but invert the result/accumulator if this is a partial
+ // sub-reduction.
+ auto GetMLA = [&](unsigned Opc, SDValue Acc, SDValue LHS,
+ SDValue RHS) -> SDValue {
+ EVT AccVT = Acc.getValueType();
+ if (!IsMLS)
+ return DAG.getNode(Opc, DL, AccVT, Acc, LHS, RHS);
+ SDValue Zero = DAG.getConstant(0, DL, AccVT);
+ SDValue NewAcc = DAG.getNode(ISD::SUB, DL, AccVT, Zero, Acc);
+ SDValue MLA = DAG.getNode(Opc, DL, AccVT, NewAcc, LHS, RHS);
+ return DAG.getNode(ISD::SUB, DL, AccVT, Zero, MLA);
+ };
+
// partial_reduce_*mla(acc, mul(ext(x), splat(C)), splat(1))
// -> partial_reduce_*mla(acc, x, C)
APInt C;
@@ -13478,11 +13489,6 @@ SDValue DAGCombiner::foldPartialReduceMLAMulOp(SDNode *N) {
unsigned NewOpcode = LHSOpcode == ISD::SIGN_EXTEND
? ISD::PARTIAL_REDUCE_SMLA
: ISD::PARTIAL_REDUCE_UMLA;
- if (IsMLS)
- NewOpcode = NewOpcode == ISD::PARTIAL_REDUCE_SMLA
- ? ISD::PARTIAL_REDUCE_SMLS
- : ISD::PARTIAL_REDUCE_UMLS;
-
// Only perform these combines if the target supports folding
// the extends into the operation.
if (!TLI.isPartialReduceMLALegalOrCustom(
@@ -13492,7 +13498,7 @@ SDValue DAGCombiner::foldPartialReduceMLAMulOp(SDNode *N) {
SDValue C = DAG.getConstant(CTrunc, DL, LHSExtOpVT);
ApplyPredicate(C, LHSExtOp);
- return DAG.getNode(NewOpcode, DL, N->getValueType(0), Acc, LHSExtOp, C);
+ return GetMLA(NewOpcode, Acc, LHSExtOp, C);
}
unsigned RHSOpcode = RHS->getOpcode();
@@ -13505,18 +13511,15 @@ SDValue DAGCombiner::foldPartialReduceMLAMulOp(SDNode *N) {
unsigned NewOpc;
if (LHSOpcode == ISD::SIGN_EXTEND && RHSOpcode == ISD::SIGN_EXTEND)
- NewOpc = IsMLS ? ISD::PARTIAL_REDUCE_SMLS : ISD::PARTIAL_REDUCE_SMLA;
+ NewOpc = ISD::PARTIAL_REDUCE_SMLA;
else if (LHSOpcode == ISD::ZERO_EXTEND && RHSOpcode == ISD::ZERO_EXTEND)
- NewOpc = IsMLS ? ISD::PARTIAL_REDUCE_UMLS : ISD::PARTIAL_REDUCE_UMLA;
- else if (!IsMLS && LHSOpcode == ISD::SIGN_EXTEND &&
- RHSOpcode == ISD::ZERO_EXTEND)
+ NewOpc = ISD::PARTIAL_REDUCE_UMLA;
+ else if (LHSOpcode == ISD::SIGN_EXTEND && RHSOpcode == ISD::ZERO_EXTEND)
NewOpc = ISD::PARTIAL_REDUCE_SUMLA;
- else if (!IsMLS && LHSOpcode == ISD::ZERO_EXTEND &&
- RHSOpcode == ISD::SIGN_EXTEND) {
+ else if (LHSOpcode == ISD::ZERO_EXTEND && RHSOpcode == ISD::SIGN_EXTEND) {
NewOpc = ISD::PARTIAL_REDUCE_SUMLA;
std::swap(LHSExtOp, RHSExtOp);
- } else if (!IsMLS && LHSOpcode == ISD::FP_EXTEND &&
- RHSOpcode == ISD::FP_EXTEND) {
+ } else if (LHSOpcode == ISD::FP_EXTEND && RHSOpcode == ISD::FP_EXTEND) {
NewOpc = ISD::PARTIAL_REDUCE_FMLA;
} else
return SDValue();
@@ -13538,7 +13541,7 @@ SDValue DAGCombiner::foldPartialReduceMLAMulOp(SDNode *N) {
return SDValue();
ApplyPredicate(RHSExtOp, LHSExtOp);
- return DAG.getNode(NewOpc, DL, N->getValueType(0), Acc, LHSExtOp, RHSExtOp);
+ return GetMLA(NewOpc, Acc, LHSExtOp, RHSExtOp);
}
// partial.reduce.*mla(acc, *ext(op), splat(1))
@@ -13585,11 +13588,10 @@ SDValue DAGCombiner::foldPartialReduceAdd(SDNode *N) {
Op1.getValueType().getVectorElementType() != AccElemVT)
return SDValue();
- unsigned NewOpcode =
- N->getOpcode() == ISD::PARTIAL_REDUCE_FMLA ? ISD::PARTIAL_REDUCE_FMLA
- : Op1IsSigned
- ? (IsMLS ? ISD::PARTIAL_REDUCE_SMLS : ISD::PARTIAL_REDUCE_SMLA)
- : (IsMLS ? ISD::PARTIAL_REDUCE_UMLS : ISD::PARTIAL_REDUCE_UMLA);
+ unsigned NewOpcode = N->getOpcode() == ISD::PARTIAL_REDUCE_FMLA
+ ? ISD::PARTIAL_REDUCE_FMLA
+ : Op1IsSigned ? ISD::PARTIAL_REDUCE_SMLA
+ : ISD::PARTIAL_REDUCE_UMLA;
SDValue UnextOp1 = Op1.getOperand(0);
EVT UnextOp1VT = UnextOp1.getValueType();
@@ -13609,8 +13611,13 @@ SDValue DAGCombiner::foldPartialReduceAdd(SDNode *N) {
: DAG.getConstant(0, DL, UnextOp1VT);
Constant = DAG.getSelect(DL, UnextOp1VT, Pred, Constant, Zero);
}
- return DAG.getNode(NewOpcode, DL, N->getValueType(0), Acc, UnextOp1,
- Constant);
+ EVT AccVT = Acc.getValueType();
+ if (!IsMLS)
+ return DAG.getNode(NewOpcode, DL, AccVT, Acc, UnextOp1, Constant);
+ SDValue Zero = DAG.getConstant(0, DL, AccVT);
+ SDValue NegAcc = DAG.getNode(ISD::SUB, DL, AccVT, Zero, Acc);
+ SDValue MLA = DAG.getNode(NewOpcode, DL, AccVT, NegAcc, UnextOp1, Constant);
+ return DAG.getNode(ISD::SUB, DL, AccVT, Zero, MLA);
}
SDValue DAGCombiner::visitVP_STRIDED_LOAD(SDNode *N) {
diff --git a/llvm/lib/CodeGen/SelectionDAG/LegalizeIntegerTypes.cpp b/llvm/lib/CodeGen/SelectionDAG/LegalizeIntegerTypes.cpp
index 99f9f9f15c777..4a27f804d6720 100644
--- a/llvm/lib/CodeGen/SelectionDAG/LegalizeIntegerTypes.cpp
+++ b/llvm/lib/CodeGen/SelectionDAG/LegalizeIntegerTypes.cpp
@@ -171,8 +171,6 @@ void DAGTypeLegalizer::PromoteIntegerResult(SDNode *N, unsigned ResNo) {
case ISD::PARTIAL_REDUCE_UMLA:
case ISD::PARTIAL_REDUCE_SMLA:
- case ISD::PARTIAL_REDUCE_UMLS:
- case ISD::PARTIAL_REDUCE_SMLS:
case ISD::PARTIAL_REDUCE_SUMLA:
Res = PromoteIntRes_PARTIAL_REDUCE_MLA(N);
break;
@@ -2170,8 +2168,6 @@ bool DAGTypeLegalizer::PromoteIntegerOperand(SDNode *N, unsigned OpNo) {
break;
case ISD::PARTIAL_REDUCE_UMLA:
case ISD::PARTIAL_REDUCE_SMLA:
- case ISD::PARTIAL_REDUCE_UMLS:
- case ISD::PARTIAL_REDUCE_SMLS:
case ISD::PARTIAL_REDUCE_SUMLA:
Res = PromoteIntOp_PARTIAL_REDUCE_MLA(N);
break;
@@ -3011,12 +3007,10 @@ SDValue DAGTypeLegalizer::PromoteIntOp_PARTIAL_REDUCE_MLA(SDNode *N) {
SmallVector<SDValue, 1> NewOps(N->ops());
switch (N->getOpcode()) {
case ISD::PARTIAL_REDUCE_SMLA:
- case ISD::PARTIAL_REDUCE_SMLS:
NewOps[1] = SExtPromotedInteger(N->getOperand(1));
NewOps[2] = SExtPromotedInteger(N->getOperand(2));
break;
case ISD::PARTIAL_REDUCE_UMLA:
- case ISD::PARTIAL_REDUCE_UMLS:
NewOps[1] = ZExtPromotedInteger(N->getOperand(1));
NewOps[2] = ZExtPromotedInteger(N->getOperand(2));
break;
diff --git a/llvm/lib/CodeGen/SelectionDAG/LegalizeVectorOps.cpp b/llvm/lib/CodeGen/SelectionDAG/LegalizeVectorOps.cpp
index 3b40da3bdea43..c00fbe79c6d64 100644
--- a/llvm/lib/CodeGen/SelectionDAG/LegalizeVectorOps.cpp
+++ b/llvm/lib/CodeGen/SelectionDAG/LegalizeVectorOps.cpp
@@ -535,8 +535,6 @@ SDValue VectorLegalizer::LegalizeOp(SDValue Op) {
}
case ISD::PARTIAL_REDUCE_UMLA:
case ISD::PARTIAL_REDUCE_SMLA:
- case ISD::PARTIAL_REDUCE_UMLS:
- case ISD::PARTIAL_REDUCE_SMLS:
case ISD::PARTIAL_REDUCE_SUMLA:
case ISD::PARTIAL_REDUCE_FMLA:
Action =
diff --git a/llvm/lib/CodeGen/SelectionDAG/LegalizeVectorTypes.cpp b/llvm/lib/CodeGen/SelectionDAG/LegalizeVectorTypes.cpp
index 2ed728b9a65de..564bf3b7f152e 100644
--- a/llvm/lib/CodeGen/SelectionDAG/LegalizeVectorTypes.cpp
+++ b/llvm/lib/CodeGen/SelectionDAG/LegalizeVectorTypes.cpp
@@ -1531,8 +1531,6 @@ void DAGTypeLegalizer::SplitVectorResult(SDNode *N, unsigned ResNo) {
break;
case ISD::PARTIAL_REDUCE_UMLA:
case ISD::PARTIAL_REDUCE_SMLA:
- case ISD::PARTIAL_REDUCE_UMLS:
- case ISD::PARTIAL_REDUCE_SMLS:
case ISD::PARTIAL_REDUCE_SUMLA:
case ISD::PARTIAL_REDUCE_FMLA:
SplitVecRes_PARTIAL_REDUCE_MLA(N, Lo, Hi);
@@ -3753,8 +3751,6 @@ bool DAGTypeLegalizer::SplitVectorOperand(SDNode *N, unsigned OpNo) {
break;
case ISD::PARTIAL_REDUCE_UMLA:
case ISD::PARTIAL_REDUCE_SMLA:
- case ISD::PARTIAL_REDUCE_UMLS:
- case ISD::PARTIAL_REDUCE_SMLS:
case ISD::PARTIAL_REDUCE_SUMLA:
case ISD::PARTIAL_REDUCE_FMLA:
Res = SplitVecOp_PARTIAL_REDUCE_MLA(N);
diff --git a/llvm/lib/CodeGen/SelectionDAG/SelectionDAG.cpp b/llvm/lib/CodeGen/SelectionDAG/SelectionDAG.cpp
index 475970a4c2d22..8e06325c3a8d5 100644
--- a/llvm/lib/CodeGen/SelectionDAG/SelectionDAG.cpp
+++ b/llvm/lib/CodeGen/SelectionDAG/SelectionDAG.cpp
@@ -8664,8 +8664,6 @@ SDValue SelectionDAG::getNode(unsigned Opcode, const SDLoc &DL, EVT VT,
}
case ISD::PARTIAL_REDUCE_UMLA:
case ISD::PARTIAL_REDUCE_SMLA:
- case ISD::PARTIAL_REDUCE_UMLS:
- case ISD::PARTIAL_REDUCE_SMLS:
case ISD::PARTIAL_REDUCE_SUMLA:
case ISD::PARTIAL_REDUCE_FMLA: {
[[maybe_unused]] EVT AccVT = N1.getValueType();
diff --git a/llvm/lib/CodeGen/SelectionDAG/SelectionDAGDumper.cpp b/llvm/lib/CodeGen/SelectionDAG/SelectionDAGDumper.cpp
index 77c4af81d8637..7161dd299f830 100644
--- a/llvm/lib/CodeGen/SelectionDAG/SelectionDAGDumper.cpp
+++ b/llvm/lib/CodeGen/SelectionDAG/SelectionDAGDumper.cpp
@@ -604,12 +604,8 @@ std::string SDNode::getOperationName(const SelectionDAG *G) const {
case ISD::PARTIAL_REDUCE_UMLA:
return "partial_reduce_umla";
- case ISD::PARTIAL_REDUCE_UMLS:
- return "partial_reduce_umls";
case ISD::PARTIAL_REDUCE_SMLA:
return "partial_reduce_smla";
- case ISD::PARTIAL_REDUCE_SMLS:
- return "partial_reduce_smls";
case ISD::PARTIAL_REDUCE_SUMLA:
return "partial_reduce_sumla";
case ISD::PARTIAL_REDUCE_FMLA:
diff --git a/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp b/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
index 754d7a02d57e9..c1a6654f68b2b 100644
--- a/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
+++ b/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
@@ -2015,11 +2015,12 @@ AArch64TargetLowering::AArch64TargetLowering(const TargetMachine &TM,
if (Subtarget->isSVEorStreamingSVEAvailable()) {
// Mark known legal pairs as 'Legal' (these will expand to UDOT or SDOT).
// Other pairs will default to 'Expand'.
- static const unsigned DotMLAOps[] = {ISD::PARTIAL_REDUCE_SMLA,
- ISD::PARTIAL_REDUCE_UMLA};
- setPartialReduceMLAAction(DotMLAOps, MVT::nxv2i64, MVT::nxv8i16, Legal);
- setPartialReduceMLAAction(DotMLAOps, MVT::nxv4i32, MVT::nxv16i8, Legal);
- setPartialReduceMLAAction(DotMLAOps, MVT::nxv2i64, MVT::nxv16i8, Custom);
+ static const unsigned MLAOps[] = {ISD::PARTIAL_REDUCE_SMLA,
+ ISD::PARTIAL_REDUCE_UMLA};
+ setPartialReduceMLAAction(MLAOps, MVT::nxv2i64, MVT::nxv8i16, Legal);
+ setPartialReduceMLAAction(MLAOps, MVT::nxv4i32, MVT::nxv16i8, Legal);
+
+ setPartialReduceMLAAction(MLAOps, MVT::nxv2i64, MVT::nxv16i8, Custom);
if (Subtarget->hasMatMulInt8()) {
setPartialReduceMLAAction(ISD::PARTIAL_REDUCE_SUMLA, MVT::nxv4i32,
@@ -2029,13 +2030,10 @@ AArch64TargetLowering::AArch64TargetLowering(const TargetMachine &TM,
}
if (Subtarget->hasSVE2() || Subtarget->hasSME()) {
- static const unsigned MLALBTOps[] = {
- ISD::PARTIAL_REDUCE_SMLA, ISD::PARTIAL_REDUCE_UMLA,
- ISD::PARTIAL_REDUCE_SMLS, ISD::PARTIAL_REDUCE_UMLS};
// Wide add types
- setPartialReduceMLAAction(MLALBTOps, MVT::nxv2i64, MVT::nxv4i32, Legal);
- setPartialReduceMLAAction(MLALBTOps, MVT::nxv4i32, MVT::nxv8i16, Legal);
- setPartialReduceMLAAction(MLALBTOps, MVT::nxv8i16, MVT::nxv16i8, Legal);
+ setPartialReduceMLAAction(MLAOps, MVT::nxv2i64, MVT::nxv4i32, Legal);
+ setPartialReduceMLAAction(MLAOps, MVT::nxv4i32, MVT::nxv8i16, Legal);
+ setPartialReduceMLAAction(MLAOps, MVT::nxv8i16, MVT::nxv16i8, Legal);
setOperationAction(ISD::CLMUL, {MVT::nxv16i8, MVT::nxv4i32}, Legal);
}
@@ -2125,18 +2123,15 @@ AArch64TargetLowering::AArch64TargetLowering(const TargetMachine &TM,
setOperationAction(ISD::EXPERIMENTAL_VECTOR_HISTOGRAM, MVT::nxv2i64,
Custom);
- static const unsigned DotMLAOps[] = {ISD::PARTIAL_REDUCE_SMLA,
- ISD::PARTIAL_REDUCE_UMLA};
- static const unsigned MLALBTOps[] = {
- ISD::PARTIAL_REDUCE_SMLA, ISD::PARTIAL_REDUCE_UMLA,
- ISD::PARTIAL_REDUCE_SMLS, ISD::PARTIAL_REDUCE_UMLS};
+ static const unsigned MLAOps[] = {ISD::PARTIAL_REDUCE_SMLA,
+ ISD::PARTIAL_REDUCE_UMLA};
// Must be lowered to SVE instructions.
- setPartialReduceMLAAction(DotMLAOps, MVT::v2i64, MVT::v8i16, Custom);
- setPartialReduceMLAAction(DotMLAOps, MVT::v2i64, MVT::v16i8, Custom);
- setPartialReduceMLAAction(DotMLAOps, MVT::v4i32, MVT::v16i8, Custom);
- setPartialReduceMLAAction(MLALBTOps, MVT::v2i64, MVT::v4i32, Custom);
- setPartialReduceMLAAction(MLALBTOps, MVT::v4i32, MVT::v8i16, Custom);
- setPartialReduceMLAAction(MLALBTOps, MVT::v8i16, MVT::v16i8, Custom);
+ setPartialReduceMLAAction(MLAOps, MVT::v2i64, MVT::v4i32, Custom);
+ setPartialReduceMLAAction(MLAOps, MVT::v2i64, MVT::v8i16, Custom);
+ setPartialReduceMLAAction(MLAOps, MVT::v2i64, MVT::v16i8, Custom);
+ setPartialReduceMLAAction(MLAOps, MVT::v4i32, MVT::v8i16, Custom);
+ setPartialReduceMLAAction(MLAOps, MVT::v4i32, MVT::v16i8, Custom);
+ setPartialReduceMLAAction(MLAOps, MVT::v8i16, MVT::v16i8, Custom);
}
}
@@ -2417,26 +2412,23 @@ void AArch64TargetLowering::addTypeForFixedLengthSVE(MVT VT) {
bool PreferNEON = VT.is64BitVector() || VT.is128BitVector();
bool PreferSVE = !PreferNEON && Subtarget->isSVEAvailable();
- static const unsigned MLALBTOps[] = {
- ISD::PARTIAL_REDUCE_SMLA, ISD::PARTIAL_REDUCE_UMLA,
- ISD::PARTIAL_REDUCE_SMLS, ISD::PARTIAL_REDUCE_UMLS};
- static const unsigned DotMLAOps[] = {ISD::PARTIAL_REDUCE_SMLA,
- ISD::PARTIAL_REDUCE_UMLA};
+ static const unsigned MLAOps[] = {ISD::PARTIAL_REDUCE_SMLA,
+ ISD::PARTIAL_REDUCE_UMLA};
unsigned NumElts = VT.getVectorNumElements();
if (VT.getVectorElementType() == MVT::i64) {
- setPartialReduceMLAAction(DotMLAOps, VT,
+ setPartialReduceMLAAction(MLAOps, VT,
MVT::getVectorVT(MVT::i8, NumElts * 8), Custom);
- setPartialReduceMLAAction(DotMLAOps, VT,
+ setPartialReduceMLAAction(MLAOps, VT,
MVT::getVectorVT(MVT::i16, NumElts * 4), Custom);
- setPartialReduceMLAAction(MLALBTOps, VT,
+ setPartialReduceMLAAction(MLAOps, VT,
MVT::getVectorVT(MVT::i32, NumElts * 2), Custom);
} else if (VT.getVectorElementType() == MVT::i32) {
- setPartialReduceMLAAction(DotMLAOps, VT,
+ setPartialReduceMLAAction(MLAOps, VT,
MVT::getVectorVT(MVT::i8, NumElts * 4), Custom);
- setPartialReduceMLAAction(MLALBTOps, VT,
+ setPartialReduceMLAAction(MLAOps, VT,
MVT::getVectorVT(MVT::i16, NumElts * 2), Custom);
} else if (VT.getVectorElementType() == MVT::i16) {
- setPartialReduceMLAAction(MLALBTOps, VT,
+ setPartialReduceMLAAction(MLAOps, VT,
MVT::getVectorVT(MVT::i8, NumElts * 2), Custom);
}
if (Subtarget->hasMatMulInt8()) {
@@ -8485,8 +8477,6 @@ SDValue AArch64TargetLowering::LowerOperation(SDValue Op,
return LowerVECTOR_HISTOGRAM(Op, DAG);
case ISD::PARTIAL_REDUCE_SMLA:
case ISD::PARTIAL_REDUCE_UMLA:
- case ISD::PARTIAL_REDUCE_SMLS:
- case ISD::PARTIAL_REDUCE_UMLS:
case ISD::PARTIAL_REDUCE_SUMLA:
case ISD::PARTIAL_REDUCE_FMLA:
return LowerPARTIAL_REDUCE_MLA(Op, DAG);
@@ -32653,13 +32643,10 @@ AArch64TargetLowering::LowerPARTIAL_REDUCE_MLA(SDValue Op,
EVT ResultVT = Op.getValueType();
EVT OrigResultVT = ResultVT;
EVT OpVT = LHS.getValueType();
- bool IsMLS = Op.getOpcode() == ISD::PARTIAL_REDUCE_UMLS ||
- Op.getOpcode() == ISD::PARTIAL_REDUCE_SMLS;
// We can handle this case natively by accumulating into a wider
// zero-padded vector.
if (ResultVT == MVT::v2i32 && OpVT == MVT::v16i8) {
- assert(!IsMLS && "Cannot handle this case for sub-reductions");
SDValue ZeroVec = DAG.getConstant(0, DL, MVT::v4i32);
SDValue WideAcc = DAG.getInsertSubvector(DL, ZeroVec, Acc, 0);
SDValue Wide =
@@ -32687,15 +32674,12 @@ AArch64TargetLowering::LowerPARTIAL_REDUCE_MLA(SDValue Op,
return ConvertToScalable ? convertFromScalableVector(DAG, OrigResultVT, Op)
: Op;
- assert(!IsMLS && "i8 -> i64 custom sub-reductions are not supported");
-
EVT DotVT = ResultVT.isScalableVector() ? MVT::nxv4i32 : MVT::v4i32;
SDValue DotNode = DAG.getNode(Op.getOpcode(), DL, DotVT,
DAG.getConstant(0, DL, DotVT), LHS, RHS);
SDValue Res;
- bool IsUnsigned = Op.getOpcode() == ISD::PARTIAL_REDUCE_UMLA ||
- Op.getOpcode() == ISD::PARTIAL_REDUCE_UMLS;
+ bool IsUnsigned = Op.getOpcode() == ISD::PARTIAL_REDUCE_UMLA;
if (Subtarget->hasSVE2() || Subtarget->isStreamingSVEAvailable()) {
unsigned LoOpcode = IsUnsigned ? AArch64ISD::UADDWB : AArch64ISD::SADDWB;
unsigned HiOpcode = IsUnsigned ? AArch64ISD::UADDWT : AArch64ISD::SADDWT;
diff --git a/llvm/lib/Target/AArch64/AArch64SVEInstrInfo.td b/llvm/lib/Target/AArch64/AArch64SVEInstrInfo.td
index 0c8ccd2c5acf9..499940f227c3f 100644
--- a/llvm/lib/Target/AArch64/AArch64SVEInstrInfo.td
+++ b/llvm/lib/Target/AArch64/AArch64SVEInstrInfo.td
@@ -451,6 +451,9 @@ def AArch64fmul_p_oneuse : PatFrag<(ops node:$pred, node:$src1, node:$src2),
(AArch64fmul_p node:$pred, node:$src1, node:$src2)>;
+def AArch64ineg : PatFrags<(ops node:$op),
+ [(sub (SVEDup0), node:$op)]>;
+
def AArch64fabd_p : PatFrags<(ops node:$pg, node:$op1, node:$op2),
[(int_aarch64_sve_fabd_u node:$pg, node:$op1, node:$op2),
(AArch64fabs_mt node:$pg, (AArch64fsub_p node:$pg, node:$op1, node:$op2), undef)]>;
@@ -3855,17 +3858,17 @@ let Predicates = [HasSVE2_or_SME] in {
defm UMLSLB_ZZZ : sve2_int_mla_long<0b10110, "umlslb", int_aarch64_sve_umlslb>;
defm UMLSLT_ZZZ : sve2_int_mla_long<0b10111, "umlslt", int_aarch64_sve_umlslt>;
- def : Pat<(nxv2i64 (partial_reduce_umls nxv2i64:$Acc, nxv4i32:$LHS, nxv4i32:$RHS)),
+ def : Pat<(nxv2i64 (AArch64ineg (partial_reduce_umla (AArch64ineg nxv2i64:$Acc), nxv4i32:$LHS, nxv4i32:$RHS))),
(UMLSLT_ZZZ_D (UMLSLB_ZZZ_D $Acc, $LHS, $RHS), $LHS, $RHS)>;
- def : Pat<(nxv2i64 (partial_reduce_smls nxv2i64:$Acc, nxv4i32:$LHS, nxv4i32:$RHS)),
+ def : Pat<(nxv2i64 (AArch64ineg (partial_reduce_smla (AArch64ineg nxv2i64:$Acc), nxv4i32:$LHS, nxv4i32:$RHS))),
(SMLSLT_ZZZ_D (SMLSLB_ZZZ_D $Acc, $LHS, $RHS), $LHS, $RHS)>;
- def : Pat<(nxv4i32 (partial_reduce_umls nxv4i32:$Acc, nxv8i16:$LHS, nxv8i16:$RHS)),
+ def : Pat<(nxv4i32 (AArch64ineg (partial_reduce_umla (AArch64ineg nxv4i32:$Acc), nxv8i16:$LHS, nxv8i16:$RHS))),
(UMLSLT_ZZZ_S (UMLSLB_ZZZ_S $Acc, $LHS, $RHS), $LHS, $RHS)>;
- def : Pat<(nxv4i32 (partial_reduce_smls nxv4i32:$Acc, nxv8i16:$LHS, nxv8i16:$RHS)),
+ def : Pat<(nxv4i32 (AArch64ineg (partial_reduce_smla (AArch64ineg nxv4i32:$Acc), nxv8i16:$LHS, nxv8i16:$RHS))),
(SMLSLT_ZZZ_S (SMLSLB_ZZZ_S $Acc, $LHS, $RHS), $LHS, $RHS)>;
- def : Pat<(nxv8i16 (partial_reduce_umls nxv8i16:$Acc, nxv16i8:$LHS, nxv16i8:$RHS)),
+ def : Pat<(nxv8i16 (AArch64ineg (partial_reduce_umla (AArch64ineg nxv8i16:$Acc), nxv16i8:$LHS, nxv16i8:$RHS))),
(UMLSLT_ZZZ_H (UMLSLB_ZZZ_H $Acc, $LHS, $RHS), $LHS, $RHS)>;
- def : Pat<(nxv8i16 (partial_reduce_smls nxv8i16:$Acc, nxv16i8:$LHS, nxv16i8:$RHS)),
+ def : Pat<(nxv8i16 (AArch64ineg (partial_reduce_smla (AArch64ineg nxv8i16:$Acc), nxv16i8:$LHS, nxv16i8:$RHS))),
(SMLSLT_ZZZ_H (SMLSLB_ZZZ_H $Acc, $LHS, $RHS), $LHS, $RHS)>;
def : Pat<(nxv2i64 (partial_reduce_umla nxv2i64:$Acc, nxv4i32:$LHS, nxv4i32:$RHS)),
diff --git a/llvm/test/CodeGen/AArch64/partial-reduction-sub.ll b/llvm/test/CodeGen/AArch64/partial-reduction-sub.ll
index fd71767bcea14..a94665d35fdc1 100644
--- a/llvm/test/CodeGen/AArch64/partial-reduction-sub.ll
+++ b/llvm/test/CodeGen/AArch64/partial-reduction-sub.ll
@@ -207,13 +207,16 @@ define <2 x i64> @fixed_smlslbt_i32_i64(<2 x i64> %acc, <4 x i32> %a, <4 x i32>
; Test type legalisation for sub-reductions.
;
+; FIXME: The subr's could be removed with a DAG combine.
define <vscale x 8 x i16> @legalization_split_i8_i16(<vscale x 8 x i16> %acc, <vscale x 32 x i8> %a, <vscale x 32 x i8> %b) #0 {
; CHECK-LABEL: legalization_split_i8_i16:
; CHECK: // %bb.0:
-; CHECK-NEXT: umlslb z0.h, z1.b, z3.b
-; CHECK-NEXT: umlslt z0.h, z1.b, z3.b
-; CHECK-NEXT: umlslb z0.h, z2.b, z4.b
-; CHECK-NEXT: umlslt z0.h, z2.b, z4.b
+; CHECK-NEXT: subr z0.h, z0.h, #0 // =0x0
+; CHECK-NEXT: umlalb z0.h, z1.b, z3.b
+; CHECK-NEXT: umlalt z0.h, z1.b, z3.b
+; CHECK-NEXT: umlalb z0.h, z2.b, z4.b
+; CHECK-NEXT: umlalt z0.h, z2.b, z4.b
+; CHECK-NEXT: subr z0.h, z0.h, #0 // =0x0
; CHECK-NEXT: ret
%a.zext = zext <vscale x 32 x i8> %a to <vscale x 32 x i16>
%b.zext = zext <vscale x 32 x i8> %b to <vscale x 32 x i16>
@@ -226,18 +229,9 @@ define <vscale x 8 x i16> @legalization_split_i8_i16(<vscale x 8 x i16> %acc, <v
define <vscale x 4 x i16> @legalization_promote_acc_i8_i16(<vscale x 4 x i16> %acc, <vscale x 16 x i8> %a, <vscale x 16 x i8> %b) #0 {
; CHECK-LABEL: legalization_promote_acc_i8_i16:
; CHECK: // %bb.0:
-; CHECK-NEXT: uunpklo z3.h, z1.b
-; CHECK-NEXT: uunpklo z4.h, z2.b
-; CHECK-NEXT: movi v5.2d, #0000000000000000
-; CHECK-NEXT: ptrue p0.h
-; CHECK-NEXT: uunpkhi z1.h, z1.b
-; CHECK-NEXT: uunpkhi z2.h, z2.b
-; CHECK-NEXT: msb z3.h, p0/m, z4.h, z5.h
-; CHECK-NEXT: msb z1.h, p0/m, z2.h, z5.h
-; CHECK-NEXT: uaddwb z0.s, z0.s, z3.h
-; CHECK-NEXT: uaddwt z0.s, z0.s, z3.h
-; CHECK-NEXT: uaddwb z0.s, z0.s, z1.h
-; CHECK-NEXT: uaddwt z0.s, z0.s, z1.h
+; CHECK-NEXT: subr z0.s, z0.s, #0 // =0x0
+; CHECK-NEXT: udot z0.s, z1.b, z2.b
+; CHECK-NEXT: subr z0.s, z0.s, #0 // =0x0
; CHECK-NEXT: ret
%a.zext = zext <vscale x 16 x i8> %a to <vscale x 16 x i16>
%b.zext = zext <vscale x 16 x i8> %b to <vscale x 16 x i16>
@@ -247,22 +241,6 @@ define <vscale x 4 x i16> @legalization_promote_acc_i8_i16(<vscale x 4 x i16> %a
ret <vscale x 4 x i16> %res
}
-define <vscale x 8 x i16> @legalization_promote_mul_ops_i8_i16(<vscale x 8 x i16> %acc, <vscale x 8 x i8> %a, <vscale x 8 x i8> %b) #0 {
-; CHECK-LABEL: legalization_promote_mul_ops_i8_i16:
-; CHECK: // %bb.0:
-; CHECK-NEXT: and z1.h, z1.h, #0xff
-; CHECK-NEXT: and z2.h, z2.h, #0xff
-; CHECK-NEXT: ptrue p0.h
-; CHECK-NEXT: mls z0.h, p0/m, z1.h, z2.h
-; CHECK-NEXT: ret
- %a.zext = zext <vscale x 8 x i8> %a to <vscale x 8 x i16>
- %b.zext = zext <vscale x 8 x i8> %b to <vscale x 8 x i16>
- %mul = mul <vscale x 8 x i16> %a.zext, %b.zext
- %mul.neg = sub <vscale x 8 x i16> zeroinitializer, %mul
- %res = call <vscale x 8 x i16> @llvm.vector.partial.reduce.add(<vscale x 8 x i16> %acc, <vscale x 8 x i16> %mul.neg)
- ret <vscale x 8 x i16> %res
-}
-
; Test that MLSB/T are still generated when there is no 'mul'.
define <vscale x 2 x i64> @extended_sub_i32_i64(<vscale x 2 x i64> %acc, <vscale x 4 x i32> %a, <vscale x 4 x i32> %b) #0 {
; CHECK-LABEL: extended_sub_i32_i64:
@@ -295,47 +273,13 @@ define <vscale x 2 x i64> @predicated_smlslbt_i32_i64(<vscale x 4 x i1> %pred, <
ret <vscale x 2 x i64> %res
}
-; Test that MLSB/T is not generated when the extends are mixed.
-define <vscale x 2 x i64> @negative_test_mixed_extends(<vscale x 2 x i64> %acc, <vscale x 4 x i32> %a, <vscale x 4 x i32> %b) #0 {
-; CHECK-LABEL: negative_test_mixed_extends:
-; CHECK: // %bb.0:
-; CHECK-NEXT: sunpklo z3.d, z1.s
-; CHECK-NEXT: uunpklo z4.d, z2.s
-; CHECK-NEXT: ptrue p0.d
-; CHECK-NEXT: sunpkhi z1.d, z1.s
-; CHECK-NEXT: uunpkhi z2.d, z2.s
-; CHECK-NEXT: mls z0.d, p0/m, z3.d, z4.d
-; CHECK-NEXT: mls z0.d, p0/m, z1.d, z2.d
-; CHECK-NEXT: ret
- %a.sext = sext <vscale x 4 x i32> %a to <vscale x 4 x i64>
- %b.zext = zext <vscale x 4 x i32> %b to <vscale x 4 x i64>
- %mul = mul <vscale x 4 x i64> %a.sext, %b.zext
- %mul.neg = sub <vscale x 4 x i64> zeroinitializer, %mul
- %res = call <vscale x 2 x i64> @llvm.vector.partial.reduce.add(<vscale x 2 x i64> %acc, <vscale x 4 x i64> %mul.neg)
- ret <vscale x 2 x i64> %res
-}
-
; There is no sub dot-reduction, so we can't handle natively.
define <vscale x 4 x i32> @negative_test_no_sub_dot_inst(<vscale x 4 x i32> %acc, <vscale x 16 x i8> %a, <vscale x 16 x i8> %b) #0 {
; CHECK-LABEL: negative_test_no_sub_dot_inst:
; CHECK: // %bb.0:
-; CHECK-NEXT: uunpklo z3.h, z1.b
-; CHECK-NEXT: uunpklo z4.h, z2.b
-; CHECK-NEXT: ptrue p0.s
-; CHECK-NEXT: uunpkhi z1.h, z1.b
-; CHECK-NEXT: uunpkhi z2.h, z2.b
-; CHECK-NEXT: uunpklo z5.s, z3.h
-; CHECK-NEXT: uunpklo z6.s, z4.h
-; CHECK-NEXT: uunpkhi z3.s, z3.h
-; CHECK-NEXT: uunpkhi z4.s, z4.h
-; CHECK-NEXT: mls z0.s, p0/m, z5.s, z6.s
-; CHECK-NEXT: uunpklo z5.s, z1.h
-; CHECK-NEXT: uunpklo z6.s, z2.h
-; CHECK-NEXT: uunpkhi z1.s, z1.h
-; CHECK-NEXT: uunpkhi z2.s, z2.h
-; CHECK-NEXT: mls z0.s, p0/m, z3.s, z4.s
-; CHECK-NEXT: mls z0.s, p0/m, z5.s, z6.s
-; CHECK-NEXT: mls z0.s, p0/m, z1.s, z2.s
+; CHECK-NEXT: subr z0.s, z0.s, #0 // =0x0
+; CHECK-NEXT: udot z0.s, z1.b, z2.b
+; CHECK-NEXT: subr z0.s, z0.s, #0 // =0x0
; CHECK-NEXT: ret
%a.zext = zext <vscale x 16 x i8> %a to <vscale x 16 x i32>
%b.zext = zext <vscale x 16 x i8> %b to <vscale x 16 x i32>
>From 223fee8715b348b89883c27a3b84cdc38764cb6b Mon Sep 17 00:00:00 2001
From: Sander de Smalen <sander.desmalen at arm.com>
Date: Mon, 13 Apr 2026 10:37:54 +0000
Subject: [PATCH 3/3] Add support for [B]FMLSLB/T as well
---
llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp | 22 ++++-
.../Target/AArch64/AArch64ISelLowering.cpp | 3 -
.../lib/Target/AArch64/AArch64SVEInstrInfo.td | 14 ++-
.../CodeGen/AArch64/partial-reduction-sub.ll | 96 +++++++++++++------
llvm/test/CodeGen/AArch64/sve2p1-fdot.ll | 25 +----
.../AArch64/sve2p1-fixed-length-fdot.ll | 42 ++++----
6 files changed, 118 insertions(+), 84 deletions(-)
diff --git a/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp b/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp
index 270427f9cc03e..fad4ab29c3d54 100644
--- a/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp
+++ b/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp
@@ -13408,8 +13408,8 @@ SDValue DAGCombiner::foldPartialReduceMLAMulOp(SDNode *N) {
}
bool IsMLS = false;
- if (Opc == ISD::SUB &&
- ISD::isConstantSplatVectorAllZeros(Op1->getOperand(0).getNode())) {
+ if ((Opc == ISD::SUB && isZeroOrZeroSplat(Op1->getOperand(0))) ||
+ (Opc == ISD::FSUB && isZeroOrZeroSplatFP(Op1->getOperand(0)))) {
Op1 = Op1->getOperand(1);
Opc = Op1->getOpcode();
IsMLS = true;
@@ -13469,6 +13469,13 @@ SDValue DAGCombiner::foldPartialReduceMLAMulOp(SDNode *N) {
EVT AccVT = Acc.getValueType();
if (!IsMLS)
return DAG.getNode(Opc, DL, AccVT, Acc, LHS, RHS);
+
+ if (AccVT.isFloatingPoint()) {
+ SDValue NewAcc = DAG.getNode(ISD::FNEG, DL, AccVT, Acc);
+ SDValue MLA = DAG.getNode(Opc, DL, AccVT, NewAcc, LHS, RHS);
+ return DAG.getNode(ISD::FNEG, DL, AccVT, MLA);
+ }
+
SDValue Zero = DAG.getConstant(0, DL, AccVT);
SDValue NewAcc = DAG.getNode(ISD::SUB, DL, AccVT, Zero, Acc);
SDValue MLA = DAG.getNode(Opc, DL, AccVT, NewAcc, LHS, RHS);
@@ -13570,8 +13577,8 @@ SDValue DAGCombiner::foldPartialReduceAdd(SDNode *N) {
}
bool IsMLS = false;
- if (Op1Opcode == ISD::SUB &&
- ISD::isConstantSplatVectorAllZeros(Op1->getOperand(0).getNode())) {
+ if ((Op1Opcode == ISD::SUB && isZeroOrZeroSplat(Op1->getOperand(0))) ||
+ (Op1Opcode == ISD::FSUB && isZeroOrZeroSplatFP(Op1->getOperand(0)))) {
Op1 = Op1->getOperand(1);
Op1Opcode = Op1->getOpcode();
IsMLS = true;
@@ -13614,6 +13621,13 @@ SDValue DAGCombiner::foldPartialReduceAdd(SDNode *N) {
EVT AccVT = Acc.getValueType();
if (!IsMLS)
return DAG.getNode(NewOpcode, DL, AccVT, Acc, UnextOp1, Constant);
+
+ if (AccVT.isFloatingPoint()) {
+ SDValue NegAcc = DAG.getNode(ISD::FNEG, DL, AccVT, Acc);
+ SDValue MLA = DAG.getNode(NewOpcode, DL, AccVT, NegAcc, UnextOp1, Constant);
+ return DAG.getNode(ISD::FNEG, DL, AccVT, MLA);
+ }
+
SDValue Zero = DAG.getConstant(0, DL, AccVT);
SDValue NegAcc = DAG.getNode(ISD::SUB, DL, AccVT, Zero, Acc);
SDValue MLA = DAG.getNode(NewOpcode, DL, AccVT, NegAcc, UnextOp1, Constant);
diff --git a/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp b/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
index c1a6654f68b2b..924139cb842c7 100644
--- a/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
+++ b/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
@@ -2036,10 +2036,7 @@ AArch64TargetLowering::AArch64TargetLowering(const TargetMachine &TM,
setPartialReduceMLAAction(MLAOps, MVT::nxv8i16, MVT::nxv16i8, Legal);
setOperationAction(ISD::CLMUL, {MVT::nxv16i8, MVT::nxv4i32}, Legal);
- }
- // Handle floating-point partial reduction
- if (Subtarget->hasSVE2p1() || Subtarget->hasSME2()) {
setPartialReduceMLAAction(ISD::PARTIAL_REDUCE_FMLA, MVT::nxv4f32,
MVT::nxv8f16, Legal);
// We can use SVE2p1 fdot to emulate the fixed-length variant.
diff --git a/llvm/lib/Target/AArch64/AArch64SVEInstrInfo.td b/llvm/lib/Target/AArch64/AArch64SVEInstrInfo.td
index 499940f227c3f..c5ad2ab17224a 100644
--- a/llvm/lib/Target/AArch64/AArch64SVEInstrInfo.td
+++ b/llvm/lib/Target/AArch64/AArch64SVEInstrInfo.td
@@ -604,6 +604,9 @@ def AArch64fmulidx : PatFrags<(ops node:$op1, node:$op2, node:$idx),
def AArch64fsub : PatFrags<(ops node:$op1, node:$op2),
[(fsub node:$op1, node:$op2),
(AArch64fsub_p (SVEAllActive), node:$op1, node:$op2)]>;
+def AArch64fneg : PatFrags<(ops node:$op1),
+ [(fneg node:$op1),
+ (AArch64fneg_mt (SVEAllActive), node:$op1, (undef))]>;
def AArch64mul : PatFrag<(ops node:$op1, node:$op2),
(AArch64mul_p (SVEAnyPredicate), node:$op1, node:$op2)>;
@@ -2616,8 +2619,7 @@ let Predicates = [HasBF16, HasSVE_or_SME] in {
(BFMLALB_ZZZ nxv4f32:$acc, ZPR:$Zn, ZPR:$Zm)>;
def : Pat<(nxv4f32 (partial_reduce_fmla nxv4f32:$acc, nxv8bf16:$LHS, nxv8bf16:$RHS)),
- (BFMLALT_ZZZ (BFMLALB_ZZZ nxv4f32:$acc, ZPR:$LHS, ZPR:$RHS),
- ZPR:$LHS, ZPR:$RHS)>;
+ (BFMLALT_ZZZ (BFMLALB_ZZZ $acc, $LHS, $RHS), $LHS, $RHS)>;
defm BFCVT_ZPmZ : sve_bfloat_convert<"bfcvt", int_aarch64_sve_fcvt_bf16f32_v2, AArch64fcvtr_mt>;
defm BFCVTNT_ZPmZ : sve_bfloat_convert_top<"bfcvtnt", int_aarch64_sve_fcvtnt_bf16f32_v2>;
@@ -4170,6 +4172,11 @@ let Predicates = [HasSVE2_or_SME] in {
defm FMLSLB_ZZZ_SHH : sve2_fp_mla_long<0b010, "fmlslb", nxv4f32, nxv8f16, int_aarch64_sve_fmlslb>;
defm FMLSLT_ZZZ_SHH : sve2_fp_mla_long<0b011, "fmlslt", nxv4f32, nxv8f16, int_aarch64_sve_fmlslt>;
+ def : Pat<(nxv4f32 (partial_reduce_fmla nxv4f32:$acc, nxv8f16:$LHS, nxv8f16:$RHS)),
+ (FMLALT_ZZZ_SHH (FMLALB_ZZZ_SHH $acc, $LHS, $RHS), $LHS, $RHS)>;
+ def : Pat<(nxv4f32 (AArch64fneg (partial_reduce_fmla (AArch64fneg nxv4f32:$Acc), nxv8f16:$LHS, nxv8f16:$RHS))),
+ (FMLSLT_ZZZ_SHH (FMLSLB_ZZZ_SHH $Acc, $LHS, $RHS), $LHS, $RHS)>;
+
// SVE2 bitwise ternary operations
defm EOR3_ZZZZ : sve2_int_bitwise_ternary_op<0b000, "eor3", AArch64eor3>;
defm BCAX_ZZZZ : sve2_int_bitwise_ternary_op<0b010, "bcax", AArch64bcax>;
@@ -4371,6 +4378,9 @@ defm BFMLSLT_ZZZ_S : sve2_fp_mla_long<0b111, "bfmlslt", nxv4f32, nxv8bf16, int_a
defm BFMLSLB_ZZZI_S : sve2_fp_mla_long_by_indexed_elem<0b110, "bfmlslb", nxv4f32, nxv8bf16, int_aarch64_sve_bfmlslb_lane>;
defm BFMLSLT_ZZZI_S : sve2_fp_mla_long_by_indexed_elem<0b111, "bfmlslt", nxv4f32, nxv8bf16, int_aarch64_sve_bfmlslt_lane>;
+def : Pat<(nxv4f32 (AArch64fneg (partial_reduce_fmla (AArch64fneg nxv4f32:$Acc), nxv8bf16:$LHS, nxv8bf16:$RHS))),
+ (BFMLSLT_ZZZ_S (BFMLSLB_ZZZ_S $Acc, $LHS, $RHS), $LHS, $RHS)>;
+
defm SDOT_ZZZ_HtoS : sve2p1_two_way_dot_vv<"sdot", 0b0, int_aarch64_sve_sdot_x2>;
defm UDOT_ZZZ_HtoS : sve2p1_two_way_dot_vv<"udot", 0b1, int_aarch64_sve_udot_x2>;
defm SDOT_ZZZI_HtoS : sve2p1_two_way_dot_vvi<"sdot", 0b0, int_aarch64_sve_sdot_lane_x2>;
diff --git a/llvm/test/CodeGen/AArch64/partial-reduction-sub.ll b/llvm/test/CodeGen/AArch64/partial-reduction-sub.ll
index a94665d35fdc1..5286db229a4af 100644
--- a/llvm/test/CodeGen/AArch64/partial-reduction-sub.ll
+++ b/llvm/test/CodeGen/AArch64/partial-reduction-sub.ll
@@ -91,6 +91,35 @@ define <vscale x 2 x i64> @smlslbt_i32_i64(<vscale x 2 x i64> %acc, <vscale x 4
ret <vscale x 2 x i64> %res
}
+; requires +sve2p1
+define <vscale x 4 x float> @fmlslbt_bf16_f32(<vscale x 4 x float> %acc, <vscale x 8 x bfloat> %a, <vscale x 8 x bfloat> %b) "target-features"="+sve2p1,+bf16" {
+; CHECK-LABEL: fmlslbt_bf16_f32:
+; CHECK: // %bb.0:
+; CHECK-NEXT: bfmlslb z0.s, z1.h, z2.h
+; CHECK-NEXT: bfmlslt z0.s, z1.h, z2.h
+; CHECK-NEXT: ret
+ %a.sext = fpext <vscale x 8 x bfloat> %a to <vscale x 8 x float>
+ %b.sext = fpext <vscale x 8 x bfloat> %b to <vscale x 8 x float>
+ %mul = fmul fast <vscale x 8 x float> %a.sext, %b.sext
+ %mul.neg = fsub fast <vscale x 8 x float> zeroinitializer, %mul
+ %res = call fast <vscale x 4 x float> @llvm.vector.partial.reduce.fadd(<vscale x 4 x float> %acc, <vscale x 8 x float> %mul.neg)
+ ret <vscale x 4 x float> %res
+}
+
+define <vscale x 4 x float> @fmlslbt_f16_f32(<vscale x 4 x float> %acc, <vscale x 8 x half> %a, <vscale x 8 x half> %b) #0 {
+; CHECK-LABEL: fmlslbt_f16_f32:
+; CHECK: // %bb.0:
+; CHECK-NEXT: fmlslb z0.s, z1.h, z2.h
+; CHECK-NEXT: fmlslt z0.s, z1.h, z2.h
+; CHECK-NEXT: ret
+ %a.sext = fpext <vscale x 8 x half> %a to <vscale x 8 x float>
+ %b.sext = fpext <vscale x 8 x half> %b to <vscale x 8 x float>
+ %mul = fmul fast <vscale x 8 x float> %a.sext, %b.sext
+ %mul.neg = fsub fast <vscale x 8 x float> zeroinitializer, %mul
+ %res = call fast <vscale x 4 x float> @llvm.vector.partial.reduce.fadd(<vscale x 4 x float> %acc, <vscale x 8 x float> %mul.neg)
+ ret <vscale x 4 x float> %res
+}
+
;
; Ensure fixed-length codegen for streaming-compatible functions.
;
@@ -203,6 +232,42 @@ define <2 x i64> @fixed_smlslbt_i32_i64(<2 x i64> %acc, <4 x i32> %a, <4 x i32>
ret <2 x i64> %res
}
+; FIXME: This could use SVE2p1's bfmlslb/t
+define <4 x float> @fixed_fmlslbt_bf16_f32(<4 x float> %acc, <8 x bfloat> %a, <8 x bfloat> %b) "target-features"="+sve2p1,+bf16" {
+; CHECK-LABEL: fixed_fmlslbt_bf16_f32:
+; CHECK: // %bb.0:
+; CHECK-NEXT: fneg v0.4s, v0.4s
+; CHECK-NEXT: bfmlalb v0.4s, v1.8h, v2.8h
+; CHECK-NEXT: bfmlalt v0.4s, v1.8h, v2.8h
+; CHECK-NEXT: fneg v0.4s, v0.4s
+; CHECK-NEXT: ret
+ %a.sext = fpext <8 x bfloat> %a to <8 x float>
+ %b.sext = fpext <8 x bfloat> %b to <8 x float>
+ %mul = fmul fast <8 x float> %a.sext, %b.sext
+ %mul.neg = fsub fast <8 x float> zeroinitializer, %mul
+ %res = call fast <4 x float> @llvm.vector.partial.reduce.fadd(<4 x float> %acc, <8 x float> %mul.neg)
+ ret <4 x float> %res
+}
+
+; FIXME: This could use SVE2p1's fmlslb/t or NEON's fmlsl(2)
+define <4 x float> @fixed_fmlslbt_f16_f32(<4 x float> %acc, <8 x half> %a, <8 x half> %b) #0 {
+; CHECK-LABEL: fixed_fmlslbt_f16_f32:
+; CHECK: // %bb.0:
+; CHECK-NEXT: fneg v0.4s, v0.4s
+; CHECK-NEXT: // kill: def $q2 killed $q2 def $z2
+; CHECK-NEXT: // kill: def $q1 killed $q1 def $z1
+; CHECK-NEXT: fmlalb z0.s, z1.h, z2.h
+; CHECK-NEXT: fmlalt z0.s, z1.h, z2.h
+; CHECK-NEXT: fneg v0.4s, v0.4s
+; CHECK-NEXT: ret
+ %a.sext = fpext <8 x half> %a to <8 x float>
+ %b.sext = fpext <8 x half> %b to <8 x float>
+ %mul = fmul fast <8 x float> %a.sext, %b.sext
+ %mul.neg = fsub fast <8 x float> zeroinitializer, %mul
+ %res = call fast <4 x float> @llvm.vector.partial.reduce.fadd(<4 x float> %acc, <8 x float> %mul.neg)
+ ret <4 x float> %res
+}
+
;
; Test type legalisation for sub-reductions.
;
@@ -273,7 +338,7 @@ define <vscale x 2 x i64> @predicated_smlslbt_i32_i64(<vscale x 4 x i1> %pred, <
ret <vscale x 2 x i64> %res
}
-; There is no sub dot-reduction, so we can't handle natively.
+; There is no sub dot-reduction, so we need to introduce explicit instructions for negation.
define <vscale x 4 x i32> @negative_test_no_sub_dot_inst(<vscale x 4 x i32> %acc, <vscale x 16 x i8> %a, <vscale x 16 x i8> %b) #0 {
; CHECK-LABEL: negative_test_no_sub_dot_inst:
; CHECK: // %bb.0:
@@ -289,35 +354,6 @@ define <vscale x 4 x i32> @negative_test_no_sub_dot_inst(<vscale x 4 x i32> %acc
ret <vscale x 4 x i32> %res
}
-; There exists FMLSLB/T instructions, but those are not yet supported.
-define <vscale x 4 x float> @negative_test_unsupported_fmlslbt(<vscale x 4 x float> %acc, <vscale x 8 x half> %a, <vscale x 8 x half> %b) #0 {
-; CHECK-LABEL: negative_test_unsupported_fmlslbt:
-; CHECK: // %bb.0:
-; CHECK-NEXT: uunpklo z3.s, z1.h
-; CHECK-NEXT: uunpklo z4.s, z2.h
-; CHECK-NEXT: ptrue p0.s
-; CHECK-NEXT: uunpkhi z1.s, z1.h
-; CHECK-NEXT: uunpkhi z2.s, z2.h
-; CHECK-NEXT: fcvt z3.s, p0/m, z3.h
-; CHECK-NEXT: fcvt z4.s, p0/m, z4.h
-; CHECK-NEXT: fcvt z1.s, p0/m, z1.h
-; CHECK-NEXT: fcvt z2.s, p0/m, z2.h
-; CHECK-NEXT: fmul z3.s, z3.s, z4.s
-; CHECK-NEXT: fmul z1.s, z1.s, z2.s
-; CHECK-NEXT: fneg z3.s, p0/m, z3.s
-; CHECK-NEXT: fneg z1.s, p0/m, z1.s
-; CHECK-NEXT: fadd z0.s, z0.s, z3.s
-; CHECK-NEXT: fadd z0.s, z0.s, z1.s
-; CHECK-NEXT: ret
- %a.sext = fpext <vscale x 8 x half> %a to <vscale x 8 x float>
- %b.sext = fpext <vscale x 8 x half> %b to <vscale x 8 x float>
- %mul = fmul fast <vscale x 8 x float> %a.sext, %b.sext
- %mul.neg = fsub fast <vscale x 8 x float> zeroinitializer, %mul
- %res = call fast <vscale x 4 x float> @llvm.vector.partial.reduce.fadd(<vscale x 4 x float> %acc, <vscale x 8 x float> %mul.neg)
- ret <vscale x 4 x float> %res
-}
-
-
; Make sure wider types are supported when vscale_range supports it
define void @wide_fixed_umlslbt_i8_i16(ptr %acc.ptr, ptr %a.ptr, ptr %b.ptr, ptr %dest.ptr) #0 vscale_range(2,0) {
; CHECK-LABEL: wide_fixed_umlslbt_i8_i16:
diff --git a/llvm/test/CodeGen/AArch64/sve2p1-fdot.ll b/llvm/test/CodeGen/AArch64/sve2p1-fdot.ll
index 9dbe096ebdb57..d3954dc6bf8d6 100644
--- a/llvm/test/CodeGen/AArch64/sve2p1-fdot.ll
+++ b/llvm/test/CodeGen/AArch64/sve2p1-fdot.ll
@@ -9,19 +9,8 @@ target triple = "aarch64-linux-gnu"
define <vscale x 4 x float> @fdot_wide_nxv4f32(<vscale x 4 x float> %acc, <vscale x 8 x half> %a, <vscale x 8 x half> %b) {
; SVE2-LABEL: fdot_wide_nxv4f32:
; SVE2: // %bb.0: // %entry
-; SVE2-NEXT: uunpklo z3.s, z1.h
-; SVE2-NEXT: uunpklo z4.s, z2.h
-; SVE2-NEXT: ptrue p0.s
-; SVE2-NEXT: uunpkhi z1.s, z1.h
-; SVE2-NEXT: uunpkhi z2.s, z2.h
-; SVE2-NEXT: fcvt z3.s, p0/m, z3.h
-; SVE2-NEXT: fcvt z4.s, p0/m, z4.h
-; SVE2-NEXT: fcvt z1.s, p0/m, z1.h
-; SVE2-NEXT: fcvt z2.s, p0/m, z2.h
-; SVE2-NEXT: fmul z3.s, z3.s, z4.s
-; SVE2-NEXT: fmul z1.s, z1.s, z2.s
-; SVE2-NEXT: fadd z0.s, z0.s, z3.s
-; SVE2-NEXT: fadd z0.s, z0.s, z1.s
+; SVE2-NEXT: fmlalb z0.s, z1.h, z2.h
+; SVE2-NEXT: fmlalt z0.s, z1.h, z2.h
; SVE2-NEXT: ret
;
; SVE2P1-LABEL: fdot_wide_nxv4f32:
@@ -39,13 +28,9 @@ entry:
define <vscale x 4 x float> @fdot_splat_nxv4f32(<vscale x 4 x float> %acc, <vscale x 8 x half> %a) {
; SVE2-LABEL: fdot_splat_nxv4f32:
; SVE2: // %bb.0: // %entry
-; SVE2-NEXT: uunpklo z2.s, z1.h
-; SVE2-NEXT: ptrue p0.s
-; SVE2-NEXT: uunpkhi z1.s, z1.h
-; SVE2-NEXT: fcvt z2.s, p0/m, z2.h
-; SVE2-NEXT: fcvt z1.s, p0/m, z1.h
-; SVE2-NEXT: fadd z0.s, z0.s, z2.s
-; SVE2-NEXT: fadd z0.s, z0.s, z1.s
+; SVE2-NEXT: fmov z2.h, #1.00000000
+; SVE2-NEXT: fmlalb z0.s, z1.h, z2.h
+; SVE2-NEXT: fmlalt z0.s, z1.h, z2.h
; SVE2-NEXT: ret
;
; SVE2P1-LABEL: fdot_splat_nxv4f32:
diff --git a/llvm/test/CodeGen/AArch64/sve2p1-fixed-length-fdot.ll b/llvm/test/CodeGen/AArch64/sve2p1-fixed-length-fdot.ll
index 4463b072f69cf..76b39ba429a0d 100644
--- a/llvm/test/CodeGen/AArch64/sve2p1-fixed-length-fdot.ll
+++ b/llvm/test/CodeGen/AArch64/sve2p1-fixed-length-fdot.ll
@@ -7,17 +7,11 @@ target triple = "aarch64-linux-gnu"
define void @fdot_v4f32(ptr %accptr, ptr %aptr, ptr %bptr) {
; SVE2-LABEL: fdot_v4f32:
; SVE2: // %bb.0: // %entry
-; SVE2-NEXT: ldr q0, [x1]
-; SVE2-NEXT: ldr q1, [x2]
-; SVE2-NEXT: fcvtl v2.4s, v0.4h
-; SVE2-NEXT: fcvtl v3.4s, v1.4h
-; SVE2-NEXT: fcvtl2 v0.4s, v0.8h
-; SVE2-NEXT: fcvtl2 v1.4s, v1.8h
-; SVE2-NEXT: fmul v2.4s, v2.4s, v3.4s
-; SVE2-NEXT: ldr q3, [x0]
-; SVE2-NEXT: fmul v0.4s, v0.4s, v1.4s
-; SVE2-NEXT: fadd v1.4s, v3.4s, v2.4s
-; SVE2-NEXT: fadd v0.4s, v1.4s, v0.4s
+; SVE2-NEXT: ldr q0, [x0]
+; SVE2-NEXT: ldr q1, [x1]
+; SVE2-NEXT: ldr q2, [x2]
+; SVE2-NEXT: fmlalb z0.s, z1.h, z2.h
+; SVE2-NEXT: fmlalt z0.s, z1.h, z2.h
; SVE2-NEXT: str q0, [x0]
; SVE2-NEXT: ret
;
@@ -216,14 +210,12 @@ entry:
define <4 x float> @fixed_fdot_wide(<4 x float> %acc, <8 x half> %a, <8 x half> %b) {
; SVE2-LABEL: fixed_fdot_wide:
; SVE2: // %bb.0: // %entry
-; SVE2-NEXT: fcvtl v3.4s, v1.4h
-; SVE2-NEXT: fcvtl v4.4s, v2.4h
-; SVE2-NEXT: fcvtl2 v1.4s, v1.8h
-; SVE2-NEXT: fcvtl2 v2.4s, v2.8h
-; SVE2-NEXT: fmul v3.4s, v3.4s, v4.4s
-; SVE2-NEXT: fmul v1.4s, v1.4s, v2.4s
-; SVE2-NEXT: fadd v0.4s, v0.4s, v3.4s
-; SVE2-NEXT: fadd v0.4s, v0.4s, v1.4s
+; SVE2-NEXT: // kill: def $q0 killed $q0 def $z0
+; SVE2-NEXT: // kill: def $q2 killed $q2 def $z2
+; SVE2-NEXT: // kill: def $q1 killed $q1 def $z1
+; SVE2-NEXT: fmlalb z0.s, z1.h, z2.h
+; SVE2-NEXT: fmlalt z0.s, z1.h, z2.h
+; SVE2-NEXT: // kill: def $q0 killed $q0 killed $z0
; SVE2-NEXT: ret
;
; SVE2P1-LABEL: fixed_fdot_wide:
@@ -245,12 +237,12 @@ entry:
define <2 x float> @fixed_fdot(<2 x float> %acc, <4 x half> %a, <4 x half> %b) {
; SVE2-LABEL: fixed_fdot:
; SVE2: // %bb.0: // %entry
-; SVE2-NEXT: fcvtl v1.4s, v1.4h
-; SVE2-NEXT: fcvtl v2.4s, v2.4h
-; SVE2-NEXT: fmul v1.4s, v1.4s, v2.4s
-; SVE2-NEXT: fadd v0.2s, v0.2s, v1.2s
-; SVE2-NEXT: ext v1.16b, v1.16b, v1.16b, #8
-; SVE2-NEXT: fadd v0.2s, v1.2s, v0.2s
+; SVE2-NEXT: // kill: def $d0 killed $d0 def $z0
+; SVE2-NEXT: // kill: def $d2 killed $d2 def $z2
+; SVE2-NEXT: // kill: def $d1 killed $d1 def $z1
+; SVE2-NEXT: fmlalb z0.s, z1.h, z2.h
+; SVE2-NEXT: fmlalt z0.s, z1.h, z2.h
+; SVE2-NEXT: // kill: def $d0 killed $d0 killed $z0
; SVE2-NEXT: ret
;
; SVE2P1-LABEL: fixed_fdot:
More information about the llvm-commits
mailing list