[llvm] [X86][LV] Add partial-reduction dot-product support (PR #205373)

via llvm-commits llvm-commits at lists.llvm.org
Tue Jun 23 09:25:19 PDT 2026


github-actions[bot] wrote:

<!--LLVM CODE FORMAT COMMENT: {clang-format}-->


:warning: C/C++ code formatter, clang-format found issues in your code. :warning:

<details>
<summary>
You can test this locally with the following command:
</summary>

``````````bash
git-clang-format --diff origin/main HEAD --extensions cpp,h -- llvm/lib/CodeGen/SelectionDAG/LegalizeTypes.h llvm/lib/CodeGen/SelectionDAG/LegalizeVectorTypes.cpp llvm/lib/CodeGen/SelectionDAG/TargetLowering.cpp llvm/lib/Target/X86/X86ISelLowering.cpp llvm/lib/Target/X86/X86ISelLowering.h llvm/lib/Target/X86/X86TargetTransformInfo.cpp llvm/lib/Target/X86/X86TargetTransformInfo.h llvm/lib/Transforms/Vectorize/LoopVectorizationPlanner.cpp llvm/lib/Transforms/Vectorize/LoopVectorizationPlanner.h llvm/lib/Transforms/Vectorize/VPlanTransforms.cpp llvm/lib/Transforms/Vectorize/VPlanUtils.cpp llvm/lib/Transforms/Vectorize/VPlanUtils.h --diff_from_common_commit
``````````

:warning:
The reproduction instructions above might return results for more than one PR
in a stack if you are using a stacked PR workflow. You can limit the results by
changing `origin/main` to the base branch/commit you want to compare against.
:warning:

</details>

<details>
<summary>
View the diff from clang-format here.
</summary>

``````````diff
diff --git a/llvm/lib/CodeGen/SelectionDAG/LegalizeVectorTypes.cpp b/llvm/lib/CodeGen/SelectionDAG/LegalizeVectorTypes.cpp
index a48328f6c..fc794062c 100644
--- a/llvm/lib/CodeGen/SelectionDAG/LegalizeVectorTypes.cpp
+++ b/llvm/lib/CodeGen/SelectionDAG/LegalizeVectorTypes.cpp
@@ -5626,10 +5626,9 @@ SDValue DAGTypeLegalizer::WidenVecRes_PARTIAL_REDUCE_MLA(SDNode *N) {
   assert(InputEC.hasKnownScalarFactor(AccEC) &&
          "partial-reduce input must be a multiple of the accumulator");
   unsigned ScaleFactor = InputEC.getKnownScalarFactor(AccEC);
-  EVT WidenInputVT =
-      EVT::getVectorVT(*DAG.getContext(), InputVT.getVectorElementType(),
-                       WidenAccVT.getVectorElementCount()
-                           .multiplyCoefficientBy(ScaleFactor));
+  EVT WidenInputVT = EVT::getVectorVT(
+      *DAG.getContext(), InputVT.getVectorElementType(),
+      WidenAccVT.getVectorElementCount().multiplyCoefficientBy(ScaleFactor));
 
   auto WidenInput = [&](SDValue V) {
     if (getTypeAction(V.getValueType()) == TargetLowering::TypeWidenVector)
diff --git a/llvm/lib/Target/X86/X86ISelLowering.cpp b/llvm/lib/Target/X86/X86ISelLowering.cpp
index fe75837ec..6fbb5782b 100644
--- a/llvm/lib/Target/X86/X86ISelLowering.cpp
+++ b/llvm/lib/Target/X86/X86ISelLowering.cpp
@@ -2859,10 +2859,10 @@ X86TargetLowering::X86TargetLowering(const X86TargetMachine &TM,
     if (Subtarget.useAVX512Regs())
       setPartialReduceMLAAction(ISD::PARTIAL_REDUCE_SUMLA, MVT::v16i32,
                                 MVT::v64i8, Custom);
-    setPartialReduceMLAAction(ISD::PARTIAL_REDUCE_SUMLA, MVT::v8i32,
-                              MVT::v32i8, Custom);
-    setPartialReduceMLAAction(ISD::PARTIAL_REDUCE_SUMLA, MVT::v4i32,
-                              MVT::v16i8, Custom);
+    setPartialReduceMLAAction(ISD::PARTIAL_REDUCE_SUMLA, MVT::v8i32, MVT::v32i8,
+                              Custom);
+    setPartialReduceMLAAction(ISD::PARTIAL_REDUCE_SUMLA, MVT::v4i32, MVT::v16i8,
+                              Custom);
   }
 
   // VNNI-INT8 i8: SMLA (VPDPBSSD), UMLA (VPDPBUUD).
@@ -2898,10 +2898,10 @@ X86TargetLowering::X86TargetLowering(const X86TargetMachine &TM,
     if (Subtarget.useAVX512Regs())
       setPartialReduceMLAAction(ISD::PARTIAL_REDUCE_SMLA, MVT::v16i32,
                                 MVT::v32i16, Custom);
-    setPartialReduceMLAAction(ISD::PARTIAL_REDUCE_SMLA, MVT::v8i32,
-                              MVT::v16i16, Custom);
-    setPartialReduceMLAAction(ISD::PARTIAL_REDUCE_SMLA, MVT::v4i32,
-                              MVT::v8i16, Custom);
+    setPartialReduceMLAAction(ISD::PARTIAL_REDUCE_SMLA, MVT::v8i32, MVT::v16i16,
+                              Custom);
+    setPartialReduceMLAAction(ISD::PARTIAL_REDUCE_SMLA, MVT::v4i32, MVT::v8i16,
+                              Custom);
   }
 
   // VNNI-INT16 i16: SUMLA (VPDPWSUD), UMLA (VPDPWUUD).
@@ -34672,9 +34672,8 @@ SDValue X86TargetLowering::visitMaskedStore(SelectionDAG &DAG, const SDLoc &DL,
   return DAG.getMemIntrinsicNode(X86ISD::CSTORE, DL, Tys, Ops, Ty, MMO);
 }
 
-SDValue
-X86TargetLowering::LowerPARTIAL_REDUCE_MLA(SDValue Op,
-                                            SelectionDAG &DAG) const {
+SDValue X86TargetLowering::LowerPARTIAL_REDUCE_MLA(SDValue Op,
+                                                   SelectionDAG &DAG) const {
   SDLoc DL(Op);
   SDValue Acc = Op.getOperand(0);
   SDValue LHS = Op.getOperand(1);
@@ -34686,13 +34685,11 @@ X86TargetLowering::LowerPARTIAL_REDUCE_MLA(SDValue Op,
   if (Opcode == ISD::PARTIAL_REDUCE_FMLA) {
     auto DpBuilder = [](SelectionDAG &DAG, const SDLoc &DL,
                         ArrayRef<SDValue> Ops) {
-      MVT AccVT =
-          MVT::getVectorVT(MVT::f32, Ops[0].getValueSizeInBits() / 32);
+      MVT AccVT = MVT::getVectorVT(MVT::f32, Ops[0].getValueSizeInBits() / 32);
       return DAG.getNode(X86ISD::DPBF16PS, DL, AccVT, Ops);
     };
     return SplitOpsAndApply(DAG, Subtarget, DL, AccVT, {Acc, LHS, RHS},
-                            DpBuilder, /*CheckBWI=*/false,
-                            Subtarget.hasBF16());
+                            DpBuilder, /*CheckBWI=*/false, Subtarget.hasBF16());
   }
 
   EVT InputVT = LHS.getValueType();
@@ -34702,7 +34699,7 @@ X86TargetLowering::LowerPARTIAL_REDUCE_MLA(SDValue Op,
   // operate on i32 register operands that contain packed i16/i8 elements).
   auto BitcastI16ToI32 = [&](SDValue V) {
     EVT VT = EVT::getVectorVT(*DAG.getContext(), MVT::i32,
-                               V.getValueType().getVectorNumElements() / 2);
+                              V.getValueType().getVectorNumElements() / 2);
     return DAG.getBitcast(VT, V);
   };
 
@@ -34750,8 +34747,8 @@ X86TargetLowering::LowerPARTIAL_REDUCE_MLA(SDValue Op,
       MVT VT = MVT::getVectorVT(MVT::i32, Ops[0].getValueSizeInBits() / 32);
       return DAG.getNode(ISDOpc, DL, VT, Ops);
     };
-    return SplitOpsAndApply(DAG, Subtarget, DL, AccVT,
-                            {Acc, CastLHS, CastRHS}, DpBuilder,
+    return SplitOpsAndApply(DAG, Subtarget, DL, AccVT, {Acc, CastLHS, CastRHS},
+                            DpBuilder,
                             /*CheckBWI=*/false, Allow512);
   }
 
diff --git a/llvm/lib/Target/X86/X86TargetTransformInfo.cpp b/llvm/lib/Target/X86/X86TargetTransformInfo.cpp
index 9622fd02d..3388f5979 100644
--- a/llvm/lib/Target/X86/X86TargetTransformInfo.cpp
+++ b/llvm/lib/Target/X86/X86TargetTransformInfo.cpp
@@ -5684,13 +5684,13 @@ InstructionCost X86TTIImpl::getPartialReductionCost(
     bool HasVNNI = ST->hasVNNI() || ST->hasAVXVNNI();
     bool HasVNNIINT8 = ST->hasAVXVNNIINT8() || ST->hasAVX10_2();
     bool IsSUMLA = OpAExtend != OpBExtend; // mixed sign -> VPDPBUSD
-    bool IsSMLA = OpAExtend == TTI::PR_SignExtend &&
-                  OpBExtend == TTI::PR_SignExtend;
-    bool IsUMLA = OpAExtend == TTI::PR_ZeroExtend &&
-                  OpBExtend == TTI::PR_ZeroExtend;
-    PartialReduceOpcode = IsSUMLA   ? ISD::PARTIAL_REDUCE_SUMLA
-                          : IsSMLA  ? ISD::PARTIAL_REDUCE_SMLA
-                                    : ISD::PARTIAL_REDUCE_UMLA;
+    bool IsSMLA =
+        OpAExtend == TTI::PR_SignExtend && OpBExtend == TTI::PR_SignExtend;
+    bool IsUMLA =
+        OpAExtend == TTI::PR_ZeroExtend && OpBExtend == TTI::PR_ZeroExtend;
+    PartialReduceOpcode = IsSUMLA  ? ISD::PARTIAL_REDUCE_SUMLA
+                          : IsSMLA ? ISD::PARTIAL_REDUCE_SMLA
+                                   : ISD::PARTIAL_REDUCE_UMLA;
     if (IsSUMLA && !HasVNNI)
       return Invalid;
     if ((IsSMLA || IsUMLA) && !HasVNNIINT8)
@@ -5707,14 +5707,14 @@ InstructionCost X86TTIImpl::getPartialReductionCost(
     // Check that the sign combination is supported by available instructions.
     bool HasVNNI = ST->hasVNNI() || ST->hasAVXVNNI();
     bool HasVNNIINT16 = ST->hasAVXVNNIINT16() || ST->hasAVX10_2();
-    bool IsSMLA = OpAExtend == TTI::PR_SignExtend &&
-                  OpBExtend == TTI::PR_SignExtend;
+    bool IsSMLA =
+        OpAExtend == TTI::PR_SignExtend && OpBExtend == TTI::PR_SignExtend;
     bool IsSUMLA = OpAExtend != OpBExtend; // mixed sign -> VPDPWSUD
-    bool IsUMLA = OpAExtend == TTI::PR_ZeroExtend &&
-                  OpBExtend == TTI::PR_ZeroExtend;
-    PartialReduceOpcode = IsSUMLA   ? ISD::PARTIAL_REDUCE_SUMLA
-                          : IsSMLA  ? ISD::PARTIAL_REDUCE_SMLA
-                                    : ISD::PARTIAL_REDUCE_UMLA;
+    bool IsUMLA =
+        OpAExtend == TTI::PR_ZeroExtend && OpBExtend == TTI::PR_ZeroExtend;
+    PartialReduceOpcode = IsSUMLA  ? ISD::PARTIAL_REDUCE_SUMLA
+                          : IsSMLA ? ISD::PARTIAL_REDUCE_SMLA
+                                   : ISD::PARTIAL_REDUCE_UMLA;
     if (IsSMLA && !HasVNNI)
       return Invalid;
     if ((IsSUMLA || IsUMLA) && !HasVNNIINT16)

``````````

</details>


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


More information about the llvm-commits mailing list