[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