[llvm] [X86][CodeGen] Support partial-reduce dot products (PR #205373)

Zihao Wang via llvm-commits llvm-commits at lists.llvm.org
Tue Jul 14 06:38:48 PDT 2026


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

>From 8ebc1a169f53c5319842259ad90306b1e2db4560 Mon Sep 17 00:00:00 2001
From: zh Wang <rekind133 at outlook.com>
Date: Wed, 24 Jun 2026 12:13:59 +0800
Subject: [PATCH] [X86][CodeGen] Support partial-reduce dot products

Add SelectionDAG legalization support for partial-reduce MLA nodes and lower the X86 dot-product shapes to VNNI, AVX-VNNI, AVX-VNNI-INT8/16, AVX10.2, and BF16 instructions.

Also teach the X86 cost hook about the partial-reduction shapes so clients can query whether a given accumulator/input vector combination is supported before forming the operation.
---
 llvm/lib/CodeGen/SelectionDAG/LegalizeTypes.h |    1 +
 .../SelectionDAG/LegalizeVectorTypes.cpp      |   33 +
 .../CodeGen/SelectionDAG/TargetLowering.cpp   |    4 +
 llvm/lib/Target/X86/X86ISelLowering.cpp       |  244 +++-
 llvm/lib/Target/X86/X86ISelLowering.h         |    1 +
 .../lib/Target/X86/X86TargetTransformInfo.cpp |  118 ++
 llvm/lib/Target/X86/X86TargetTransformInfo.h  |    4 +-
 .../Analysis/CostModel/X86/partial-reduce.ll  |  211 ++++
 ...vector-partial-reduce-avx512-prefer-256.ll |   76 ++
 ...r-partial-reduce-avx512-vex-split-int16.ll |   42 +
 ...or-partial-reduce-avx512-vex-split-int8.ll |   42 +
 .../vector-partial-reduce-avx512-vex-split.ll |   43 +
 .../X86/vector-partial-reduce-dot-product.ll  | 1032 +++++++++++++++++
 .../X86/vector-partial-reduce-small-vf.ll     |   45 +
 14 files changed, 1891 insertions(+), 5 deletions(-)
 create mode 100644 llvm/test/Analysis/CostModel/X86/partial-reduce.ll
 create mode 100644 llvm/test/CodeGen/X86/vector-partial-reduce-avx512-prefer-256.ll
 create mode 100644 llvm/test/CodeGen/X86/vector-partial-reduce-avx512-vex-split-int16.ll
 create mode 100644 llvm/test/CodeGen/X86/vector-partial-reduce-avx512-vex-split-int8.ll
 create mode 100644 llvm/test/CodeGen/X86/vector-partial-reduce-avx512-vex-split.ll
 create mode 100644 llvm/test/CodeGen/X86/vector-partial-reduce-dot-product.ll
 create mode 100644 llvm/test/CodeGen/X86/vector-partial-reduce-small-vf.ll

diff --git a/llvm/lib/CodeGen/SelectionDAG/LegalizeTypes.h b/llvm/lib/CodeGen/SelectionDAG/LegalizeTypes.h
index 71d3e1c66be86..6b5f997155d91 100644
--- a/llvm/lib/CodeGen/SelectionDAG/LegalizeTypes.h
+++ b/llvm/lib/CodeGen/SelectionDAG/LegalizeTypes.h
@@ -1078,6 +1078,7 @@ class LLVM_LIBRARY_VISIBILITY DAGTypeLegalizer {
   SDValue WidenVecRes_VECTOR_SHUFFLE(ShuffleVectorSDNode *N);
   SDValue WidenVecRes_VECTOR_REVERSE(SDNode *N);
   SDValue WidenVecRes_GET_ACTIVE_LANE_MASK(SDNode *N);
+  SDValue WidenVecRes_PARTIAL_REDUCE_MLA(SDNode *N);
   void WidenVecRes_VECTOR_DEINTERLEAVE(SDNode *N);
 
   SDValue WidenVecRes_Ternary(SDNode *N);
diff --git a/llvm/lib/CodeGen/SelectionDAG/LegalizeVectorTypes.cpp b/llvm/lib/CodeGen/SelectionDAG/LegalizeVectorTypes.cpp
index 97aa765642ea7..fc794062c987f 100644
--- a/llvm/lib/CodeGen/SelectionDAG/LegalizeVectorTypes.cpp
+++ b/llvm/lib/CodeGen/SelectionDAG/LegalizeVectorTypes.cpp
@@ -5320,6 +5320,12 @@ void DAGTypeLegalizer::WidenVectorResult(SDNode *N, unsigned ResNo) {
   case ISD::GET_ACTIVE_LANE_MASK:
     Res = WidenVecRes_GET_ACTIVE_LANE_MASK(N);
     break;
+  case ISD::PARTIAL_REDUCE_UMLA:
+  case ISD::PARTIAL_REDUCE_SMLA:
+  case ISD::PARTIAL_REDUCE_SUMLA:
+  case ISD::PARTIAL_REDUCE_FMLA:
+    Res = WidenVecRes_PARTIAL_REDUCE_MLA(N);
+    break;
   case ISD::VECTOR_DEINTERLEAVE:
     WidenVecRes_VECTOR_DEINTERLEAVE(N);
     break;
@@ -5608,6 +5614,33 @@ SDValue DAGTypeLegalizer::WidenVecRes_Ternary(SDNode *N) {
                      {InOp1, InOp2, InOp3, Mask, N->getOperand(4)});
 }
 
+SDValue DAGTypeLegalizer::WidenVecRes_PARTIAL_REDUCE_MLA(SDNode *N) {
+  SDLoc DL(N);
+  EVT AccVT = N->getValueType(0);
+  EVT WidenAccVT = TLI.getTypeToTransformTo(*DAG.getContext(), AccVT);
+  SDValue Acc = GetWidenedVector(N->getOperand(0));
+
+  EVT InputVT = N->getOperand(1).getValueType();
+  ElementCount AccEC = AccVT.getVectorElementCount();
+  ElementCount InputEC = InputVT.getVectorElementCount();
+  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));
+
+  auto WidenInput = [&](SDValue V) {
+    if (getTypeAction(V.getValueType()) == TargetLowering::TypeWidenVector)
+      V = GetWidenedVector(V);
+    return ModifyToType(V, WidenInputVT);
+  };
+
+  return DAG.getNode(N->getOpcode(), DL, WidenAccVT, Acc,
+                     WidenInput(N->getOperand(1)),
+                     WidenInput(N->getOperand(2)));
+}
+
 SDValue DAGTypeLegalizer::WidenVecRes_Binary(SDNode *N) {
   // Binary op widening.
   SDLoc dl(N);
diff --git a/llvm/lib/CodeGen/SelectionDAG/TargetLowering.cpp b/llvm/lib/CodeGen/SelectionDAG/TargetLowering.cpp
index cc3deaa83f63b..e419aee0bf641 100644
--- a/llvm/lib/CodeGen/SelectionDAG/TargetLowering.cpp
+++ b/llvm/lib/CodeGen/SelectionDAG/TargetLowering.cpp
@@ -13770,6 +13770,10 @@ SDValue TargetLowering::expandPartialReduceMLA(SDNode *N,
   case ISD::PARTIAL_REDUCE_SMLA:
     ExtOpcLHS = ExtOpcRHS = ISD::SIGN_EXTEND;
     break;
+  case ISD::PARTIAL_REDUCE_SUMLA:
+    ExtOpcLHS = ISD::SIGN_EXTEND;
+    ExtOpcRHS = ISD::ZERO_EXTEND;
+    break;
   case ISD::PARTIAL_REDUCE_FMLA:
     ExtOpcLHS = ExtOpcRHS = ISD::FP_EXTEND;
     break;
diff --git a/llvm/lib/Target/X86/X86ISelLowering.cpp b/llvm/lib/Target/X86/X86ISelLowering.cpp
index b95ac78f50049..9b3f9b7db7e48 100644
--- a/llvm/lib/Target/X86/X86ISelLowering.cpp
+++ b/llvm/lib/Target/X86/X86ISelLowering.cpp
@@ -2833,6 +2833,113 @@ X86TargetLowering::X86TargetLowering(const X86TargetMachine &TM,
                        ISD::INTRINSIC_WO_CHAIN,
                        ISD::INTRINSIC_W_CHAIN});
 
+  // Set up partial reduction MLA actions for VNNI dot product support.
+  //
+  // i8 x i8 -> i32:
+  //   SUMLA (sext x zext) -> VPDPBUSD (VNNI / AVX-VNNI)
+  //   SMLA  (sext x sext) -> VPDPBSSD (AVX-VNNI-INT8 / AVX10.2)
+  //   UMLA  (zext x zext) -> VPDPBUUD (AVX-VNNI-INT8 / AVX10.2)
+  //
+  // i16 x i16 -> i32:
+  //   SMLA  (sext x sext) -> VPDPWSSD  (VNNI / AVX-VNNI)
+  //   SUMLA (sext x zext) -> VPDPWSUD  (AVX-VNNI-INT16 / AVX10.2)
+  //   UMLA  (zext x zext) -> VPDPWUUD  (AVX-VNNI-INT16 / AVX10.2)
+
+  // VNNI i8: only SUMLA (VPDPBUSD).
+  if (Subtarget.hasVNNI()) {
+    if (Subtarget.useAVX512Regs())
+      setPartialReduceMLAAction(ISD::PARTIAL_REDUCE_SUMLA, MVT::v16i32,
+                                MVT::v64i8, Custom);
+    if (Subtarget.hasVLX()) {
+      setPartialReduceMLAAction(ISD::PARTIAL_REDUCE_SUMLA, MVT::v8i32,
+                                MVT::v32i8, Custom);
+      setPartialReduceMLAAction(ISD::PARTIAL_REDUCE_SUMLA, MVT::v4i32,
+                                MVT::v16i8, Custom);
+    }
+  } else if (Subtarget.hasAVXVNNI()) {
+    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);
+  }
+
+  // VNNI-INT8 i8: SMLA (VPDPBSSD), UMLA (VPDPBUUD).
+  if (Subtarget.hasAVX10_2()) {
+    unsigned Int8MLAOps[] = {ISD::PARTIAL_REDUCE_SMLA,
+                             ISD::PARTIAL_REDUCE_UMLA};
+    if (Subtarget.useAVX512Regs())
+      setPartialReduceMLAAction(Int8MLAOps, MVT::v16i32, MVT::v64i8, Custom);
+    setPartialReduceMLAAction(Int8MLAOps, MVT::v8i32, MVT::v32i8, Custom);
+    setPartialReduceMLAAction(Int8MLAOps, MVT::v4i32, MVT::v16i8, Custom);
+  } else if (Subtarget.hasAVXVNNIINT8()) {
+    unsigned Int8MLAOps[] = {ISD::PARTIAL_REDUCE_SMLA,
+                             ISD::PARTIAL_REDUCE_UMLA};
+    // On AVX512-register targets, the vectorizer may produce 512-bit
+    // operations.
+    // Register 512-bit as Custom so SplitOpsAndApply can split to 256-bit.
+    if (Subtarget.useAVX512Regs())
+      setPartialReduceMLAAction(Int8MLAOps, MVT::v16i32, MVT::v64i8, Custom);
+    setPartialReduceMLAAction(Int8MLAOps, MVT::v8i32, MVT::v32i8, Custom);
+    setPartialReduceMLAAction(Int8MLAOps, MVT::v4i32, MVT::v16i8, Custom);
+  }
+
+  // VNNI i16: only SMLA (VPDPWSSD).
+  if (Subtarget.hasVNNI()) {
+    if (Subtarget.useAVX512Regs())
+      setPartialReduceMLAAction(ISD::PARTIAL_REDUCE_SMLA, MVT::v16i32,
+                                MVT::v32i16, Custom);
+    if (Subtarget.hasVLX()) {
+      setPartialReduceMLAAction(ISD::PARTIAL_REDUCE_SMLA, MVT::v8i32,
+                                MVT::v16i16, Custom);
+      setPartialReduceMLAAction(ISD::PARTIAL_REDUCE_SMLA, MVT::v4i32,
+                                MVT::v8i16, Custom);
+    }
+  } else if (Subtarget.hasAVXVNNI()) {
+    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);
+  }
+
+  // VNNI-INT16 i16: SUMLA (VPDPWSUD), UMLA (VPDPWUUD).
+  if (Subtarget.hasAVX10_2()) {
+    unsigned Int16MLAOps[] = {ISD::PARTIAL_REDUCE_SUMLA,
+                              ISD::PARTIAL_REDUCE_UMLA};
+    if (Subtarget.useAVX512Regs())
+      setPartialReduceMLAAction(Int16MLAOps, MVT::v16i32, MVT::v32i16, Custom);
+    setPartialReduceMLAAction(Int16MLAOps, MVT::v8i32, MVT::v16i16, Custom);
+    setPartialReduceMLAAction(Int16MLAOps, MVT::v4i32, MVT::v8i16, Custom);
+  } else if (Subtarget.hasAVXVNNIINT16()) {
+    unsigned Int16MLAOps[] = {ISD::PARTIAL_REDUCE_SUMLA,
+                              ISD::PARTIAL_REDUCE_UMLA};
+    // On AVX512-register targets, the vectorizer may produce 512-bit
+    // operations.
+    // Register 512-bit as Custom so SplitOpsAndApply can split to 256-bit.
+    if (Subtarget.useAVX512Regs())
+      setPartialReduceMLAAction(Int16MLAOps, MVT::v16i32, MVT::v32i16, Custom);
+    setPartialReduceMLAAction(Int16MLAOps, MVT::v8i32, MVT::v16i16, Custom);
+    setPartialReduceMLAAction(Int16MLAOps, MVT::v4i32, MVT::v8i16, Custom);
+  }
+
+  // BF16 dot product: bf16 x bf16 -> f32 (vdpbf16ps, scale=2).
+  if (Subtarget.hasBF16()) {
+    if (Subtarget.useAVX512Regs())
+      setPartialReduceMLAAction(ISD::PARTIAL_REDUCE_FMLA, MVT::v16f32,
+                                MVT::v32bf16, Custom);
+    if (Subtarget.hasVLX()) {
+      setPartialReduceMLAAction(ISD::PARTIAL_REDUCE_FMLA, MVT::v8f32,
+                                MVT::v16bf16, Custom);
+      setPartialReduceMLAAction(ISD::PARTIAL_REDUCE_FMLA, MVT::v4f32,
+                                MVT::v8bf16, Custom);
+    }
+  }
+
   computeRegisterProperties(Subtarget.getRegisterInfo());
 
   MaxStoresPerMemset = 16; // For @llvm.memset -> sequence of stores
@@ -3141,7 +3248,7 @@ static bool isX86CCSigned(X86::CondCode X86CC) {
 
 static X86::CondCode TranslateIntegerX86CC(ISD::CondCode SetCCOpcode) {
   switch (SetCCOpcode) {
-  // clang-format off
+    // clang-format off
   default: llvm_unreachable("Invalid integer condition!");
   case ISD::SETEQ:  return X86::COND_E;
   case ISD::SETGT:  return X86::COND_G;
@@ -34591,10 +34698,138 @@ 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 {
+  SDLoc DL(Op);
+  SDValue Acc = Op.getOperand(0);
+  SDValue LHS = Op.getOperand(1);
+  SDValue RHS = Op.getOperand(2);
+  EVT AccVT = Acc.getValueType();
+  unsigned Opcode = Op.getOpcode();
+
+  // BF16 dot product: bf16 x bf16 -> f32 (vdpbf16ps).
+  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);
+      return DAG.getNode(X86ISD::DPBF16PS, DL, AccVT, Ops);
+    };
+    return SplitOpsAndApply(DAG, Subtarget, DL, AccVT, {Acc, LHS, RHS},
+                            DpBuilder, /*CheckBWI=*/false, Subtarget.hasBF16());
+  }
+
+  EVT InputVT = LHS.getValueType();
+  EVT InputEltVT = InputVT.getVectorElementType();
+
+  // Helper to bitcast i16 inputs to i32 (VNNI dot product instructions
+  // 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);
+    return DAG.getBitcast(VT, V);
+  };
+
+  bool HasAVX512VNNI = Subtarget.hasVNNI();
+  bool HasVNNIINT8 = Subtarget.hasAVXVNNIINT8() || Subtarget.hasAVX10_2();
+  bool HasVNNIINT16 = Subtarget.hasAVXVNNIINT16() || Subtarget.hasAVX10_2();
+
+  // i16 x i16 -> i32.
+  if (InputEltVT == MVT::i16) {
+    unsigned ISDOpc;
+    SDValue Op1 = LHS, Op2 = RHS;
+    switch (Opcode) {
+    case ISD::PARTIAL_REDUCE_SMLA:
+      // VPDPWSSD: signed x signed (VNNI).
+      ISDOpc = X86ISD::VPDPWSSD;
+      break;
+    case ISD::PARTIAL_REDUCE_SUMLA:
+      // VPDPWSUD: signed x unsigned (VNNI-INT16).
+      // SUMLA: Op1=sext, Op2=zext -> matches VPDPWSUD(acc, signed, unsigned).
+      assert(HasVNNIINT16 && "SUMLA i16 should only be Custom with "
+                             "AVXVNNIINT16 or AVX10.2");
+      ISDOpc = X86ISD::VPDPWSUD;
+      break;
+    case ISD::PARTIAL_REDUCE_UMLA:
+      // VPDPWUUD: unsigned x unsigned (VNNI-INT16).
+      assert(HasVNNIINT16 && "UMLA i16 should only be Custom with "
+                             "AVXVNNIINT16 or AVX10.2");
+      ISDOpc = X86ISD::VPDPWUUD;
+      break;
+    default:
+      llvm_unreachable("unexpected opcode");
+    }
+
+    SDValue CastLHS = BitcastI16ToI32(Op1);
+    SDValue CastRHS = BitcastI16ToI32(Op2);
+
+    // Allow 512-bit only when the specific instruction has a 512-bit encoding:
+    // VPDPWSSD (SMLA): 512-bit available with AVX512-VNNI.
+    // VPDPWSUD/VPDPWUUD (SUMLA/UMLA): 512-bit available with AVX10.2 only.
+    bool Allow512 = (Opcode == ISD::PARTIAL_REDUCE_SMLA)
+                        ? HasAVX512VNNI
+                        : Subtarget.hasAVX10_2();
+    auto DpBuilder = [ISDOpc](SelectionDAG &DAG, const SDLoc &DL,
+                              ArrayRef<SDValue> Ops) {
+      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,
+                            /*CheckBWI=*/false, Allow512);
+  }
+
+  // i8 x i8 -> i32.
+  {
+    unsigned ISDOpc;
+    SDValue Op1, Op2;
+    switch (Opcode) {
+    case ISD::PARTIAL_REDUCE_SUMLA:
+      // VPDPBUSD: unsigned x signed (VNNI).
+      // SUMLA: Op1=sext, Op2=zext -> VPDPBUSD(acc, unsigned, signed).
+      ISDOpc = X86ISD::VPDPBUSD;
+      Op1 = RHS; // zext (unsigned)
+      Op2 = LHS; // sext (signed)
+      break;
+    case ISD::PARTIAL_REDUCE_SMLA:
+      // VPDPBSSD: signed x signed (VNNI-INT8).
+      assert(HasVNNIINT8 && "SMLA i8 should only be Custom with "
+                            "AVXVNNIINT8 or AVX10.2");
+      ISDOpc = X86ISD::VPDPBSSD;
+      Op1 = LHS;
+      Op2 = RHS;
+      break;
+    case ISD::PARTIAL_REDUCE_UMLA:
+      // VPDPBUUD: unsigned x unsigned (VNNI-INT8).
+      assert(HasVNNIINT8 && "UMLA i8 should only be Custom with "
+                            "AVXVNNIINT8 or AVX10.2");
+      ISDOpc = X86ISD::VPDPBUUD;
+      Op1 = LHS;
+      Op2 = RHS;
+      break;
+    default:
+      llvm_unreachable("unexpected opcode");
+    }
+
+    // Allow 512-bit only when the specific instruction has a 512-bit encoding:
+    // VPDPBUSD (SUMLA): 512-bit available with AVX512-VNNI.
+    // VPDPBSSD/VPDPBUUD (SMLA/UMLA): 512-bit available with AVX10.2 only.
+    bool Allow512 = (Opcode == ISD::PARTIAL_REDUCE_SUMLA)
+                        ? HasAVX512VNNI
+                        : Subtarget.hasAVX10_2();
+    auto DpBuilder = [ISDOpc](SelectionDAG &DAG, const SDLoc &DL,
+                              ArrayRef<SDValue> Ops) {
+      MVT VT = MVT::getVectorVT(MVT::i32, Ops[0].getValueSizeInBits() / 32);
+      return DAG.getNode(ISDOpc, DL, VT, Ops);
+    };
+    return SplitOpsAndApply(DAG, Subtarget, DL, AccVT, {Acc, Op1, Op2},
+                            DpBuilder, /*CheckBWI=*/false, Allow512);
+  }
+}
+
 /// Provide custom lowering hooks for some operations.
 SDValue X86TargetLowering::LowerOperation(SDValue Op, SelectionDAG &DAG) const {
   switch (Op.getOpcode()) {
-  // clang-format off
+    // clang-format off
   default: llvm_unreachable("Should not custom lower this!");
   case ISD::ATOMIC_FENCE:       return LowerATOMIC_FENCE(Op, Subtarget, DAG);
   case ISD::ATOMIC_CMP_SWAP_WITH_SUCCESS:
@@ -34759,6 +34994,11 @@ SDValue X86TargetLowering::LowerOperation(SDValue Op, SelectionDAG &DAG) const {
   case X86ISD::CVTPS2PH:        return LowerCVTPS2PH(Op, DAG);
   case ISD::PREFETCH:           return LowerPREFETCH(Op, Subtarget, DAG);
   case ISD::FLDEXP:             return LowerFLDEXP(Op, Subtarget, DAG);
+  case ISD::PARTIAL_REDUCE_SMLA:
+  case ISD::PARTIAL_REDUCE_UMLA:
+  case ISD::PARTIAL_REDUCE_SUMLA:
+  case ISD::PARTIAL_REDUCE_FMLA:
+                                return LowerPARTIAL_REDUCE_MLA(Op, DAG);
     // clang-format on
   }
 }
diff --git a/llvm/lib/Target/X86/X86ISelLowering.h b/llvm/lib/Target/X86/X86ISelLowering.h
index 0d05c5772a707..b1fdc5a42bf36 100644
--- a/llvm/lib/Target/X86/X86ISelLowering.h
+++ b/llvm/lib/Target/X86/X86ISelLowering.h
@@ -786,6 +786,7 @@ namespace llvm {
     SDValue LRINT_LLRINTHelper(SDNode *N, SelectionDAG &DAG) const;
 
     SDValue LowerBUILD_VECTOR(SDValue Op, SelectionDAG &DAG) const;
+    SDValue LowerPARTIAL_REDUCE_MLA(SDValue Op, SelectionDAG &DAG) const;
     SDValue LowerVSELECT(SDValue Op, SelectionDAG &DAG) const;
     SDValue LowerEXTRACT_VECTOR_ELT(SDValue Op, SelectionDAG &DAG) const;
     SDValue LowerINSERT_VECTOR_ELT(SDValue Op, SelectionDAG &DAG) const;
diff --git a/llvm/lib/Target/X86/X86TargetTransformInfo.cpp b/llvm/lib/Target/X86/X86TargetTransformInfo.cpp
index 8838fd7e71f02..c197e60179f02 100644
--- a/llvm/lib/Target/X86/X86TargetTransformInfo.cpp
+++ b/llvm/lib/Target/X86/X86TargetTransformInfo.cpp
@@ -5629,6 +5629,124 @@ X86TTIImpl::getAddressComputationCost(Type *PtrTy, ScalarEvolution *SE,
   return BaseT::getAddressComputationCost(PtrTy, SE, Ptr, CostKind);
 }
 
+InstructionCost X86TTIImpl::getPartialReductionCost(
+    unsigned Opcode, Type *InputTypeA, Type *InputTypeB, Type *AccumType,
+    ElementCount VF, TTI::PartialReductionExtendKind OpAExtend,
+    TTI::PartialReductionExtendKind OpBExtend, std::optional<unsigned> BinOp,
+    TTI::TargetCostKind CostKind, std::optional<FastMathFlags> FMF) const {
+  InstructionCost Invalid = InstructionCost::getInvalid();
+
+  if (!BinOp)
+    return Invalid;
+
+  // Both input types must match.
+  if (InputTypeB && InputTypeA != InputTypeB)
+    return Invalid;
+
+  // Floating-point partial reductions require reassoc and contract to allow
+  // the fusion into a single instruction (e.g. vdpbf16ps).
+  if (AccumType->isFloatingPointTy()) {
+    assert(FMF && "Missing FastMathFlags for floating-point partial reduction");
+    if (!FMF->allowReassoc() || !FMF->allowContract())
+      return Invalid;
+  } else {
+    assert(!FMF &&
+           "FastMathFlags only apply to floating-point partial reductions");
+  }
+
+  assert(OpBExtend != TTI::PR_None && InputTypeB &&
+         "Unexpected values for OpBExtend or InputTypeB");
+
+  unsigned ScaleFactor = 0;
+  unsigned PartialReduceOpcode = 0;
+
+  // BF16 dot product: bf16 x bf16 -> f32 (vdpbf16ps, scale=2).
+  if (Opcode == Instruction::FAdd && *BinOp == Instruction::FMul &&
+      InputTypeA->isBFloatTy() && AccumType->isFloatTy()) {
+    if (!ST->hasBF16())
+      return Invalid;
+    ScaleFactor = 2;
+    PartialReduceOpcode = ISD::PARTIAL_REDUCE_FMLA;
+    if (!VF.isKnownMultipleOf(2))
+      return Invalid;
+  }
+  // i8 x i8 -> i32 (vpdpbusd/vpdpbssd/vpdpbuud, scale=4).
+  else if (Opcode == Instruction::Add && *BinOp == Instruction::Mul &&
+           InputTypeA->isIntegerTy(8) && AccumType->isIntegerTy(32)) {
+    ScaleFactor = 4;
+    if (!VF.isKnownMultipleOf(4))
+      return Invalid;
+    if (OpAExtend == TTI::PR_None)
+      return Invalid;
+    // Check that the sign combination is supported by available instructions.
+    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;
+    if (IsSUMLA && !HasVNNI)
+      return Invalid;
+    if ((IsSMLA || IsUMLA) && !HasVNNIINT8)
+      return Invalid;
+  }
+  // i16 x i16 -> i32 (vpdpwssd/vpdpwsud/vpdpwuud, scale=2).
+  else if (Opcode == Instruction::Add && *BinOp == Instruction::Mul &&
+           InputTypeA->isIntegerTy(16) && AccumType->isIntegerTy(32)) {
+    ScaleFactor = 2;
+    if (!VF.isKnownMultipleOf(2))
+      return Invalid;
+    if (OpAExtend == TTI::PR_None)
+      return Invalid;
+    // 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 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;
+    if (IsSMLA && !HasVNNI)
+      return Invalid;
+    if ((IsSUMLA || IsUMLA) && !HasVNNIINT16)
+      return Invalid;
+  } else {
+    return Invalid;
+  }
+
+  // Check the exact partial-reduction shape has a lowering action. Type
+  // legalization alone may accept smaller accumulator vectors than the X86 dot
+  // product lowerings support.
+  Type *AccVecTy =
+      VectorType::get(AccumType, VF.divideCoefficientBy(ScaleFactor));
+  Type *InputVecTy = VectorType::get(InputTypeA, VF);
+  EVT AccVT = TLI->getValueType(DL, AccVecTy);
+  EVT InputVT = TLI->getValueType(DL, InputVecTy);
+  if (!AccVT.isSimple() || !InputVT.isSimple() ||
+      !TLI->isPartialReduceMLALegalOrCustom(PartialReduceOpcode, AccVT,
+                                            InputVT))
+    return Invalid;
+
+  auto LT = getTypeLegalizationCost(AccVecTy);
+  if (LT.second == MVT::INVALID_SIMPLE_VALUE_TYPE)
+    return Invalid;
+
+  unsigned CostOpcode =
+      AccumType->isFloatingPointTy() ? Instruction::FMul : Instruction::Mul;
+  InstructionCost DotProductCost =
+      getArithmeticInstrCost(CostOpcode, AccVecTy, CostKind);
+  if (!DotProductCost.isValid())
+    return Invalid;
+  return std::max(InstructionCost(LT.first), DotProductCost);
+}
+
 InstructionCost
 X86TTIImpl::getArithmeticReductionCost(unsigned Opcode, VectorType *ValTy,
                                        std::optional<FastMathFlags> FMF,
diff --git a/llvm/lib/Target/X86/X86TargetTransformInfo.h b/llvm/lib/Target/X86/X86TargetTransformInfo.h
index 22171f5469d98..c77522e1b4e14 100644
--- a/llvm/lib/Target/X86/X86TargetTransformInfo.h
+++ b/llvm/lib/Target/X86/X86TargetTransformInfo.h
@@ -157,9 +157,7 @@ class X86TTIImpl final : public BasicTTIImplBase<X86TTIImpl> {
       ElementCount VF, TTI::PartialReductionExtendKind OpAExtend,
       TTI::PartialReductionExtendKind OpBExtend, std::optional<unsigned> BinOp,
       TTI::TargetCostKind CostKind,
-      std::optional<FastMathFlags> FMF) const override {
-    return InstructionCost::getInvalid();
-  }
+      std::optional<FastMathFlags> FMF) const override;
 
   InstructionCost getMinMaxCost(Intrinsic::ID IID, Type *Ty,
                                 TTI::TargetCostKind CostKind,
diff --git a/llvm/test/Analysis/CostModel/X86/partial-reduce.ll b/llvm/test/Analysis/CostModel/X86/partial-reduce.ll
new file mode 100644
index 0000000000000..84f95109c3a6b
--- /dev/null
+++ b/llvm/test/Analysis/CostModel/X86/partial-reduce.ll
@@ -0,0 +1,211 @@
+; NOTE: Assertions have been autogenerated by utils/update_analyze_test_checks.py UTC_ARGS: --filter "Cost.of.*EXPRESSION" --version 6
+; RUN: opt -passes=loop-vectorize -enable-epilogue-vectorization=false \
+; RUN:     -debug-only=loop-vectorize -disable-output < %s 2>&1 | FileCheck %s
+
+; REQUIRES: asserts
+target triple = "x86_64-unknown-linux-gnu"
+
+define i32 @dot_s8u8_vf64(ptr readonly %a, ptr readonly %b) #0 {
+; CHECK-LABEL: 'dot_s8u8_vf64'
+; CHECK:  Cost of 1 for VF 64: EXPRESSION vp<[[VP8:%[0-9]+]]> = ir<%sum> + partial.reduce.add (mul nsw (ir<%load.a> sext to i32), (ir<%load.b> zext to i32))
+; CHECK:  Cost of 1 for VF 64: EXPRESSION vp<[[VP8]]> = ir<%sum> + partial.reduce.add (mul nsw (ir<%load.a> sext to i32), (ir<%load.b> zext to i32))
+;
+entry:
+  br label %loop
+
+loop:
+  %iv = phi i64 [ 0, %entry ], [ %iv.next, %loop ]
+  %sum = phi i32 [ 0, %entry ], [ %sum.next, %loop ]
+  %gep.a = getelementptr inbounds i8, ptr %a, i64 %iv
+  %load.a = load i8, ptr %gep.a, align 1
+  %ext.a = sext i8 %load.a to i32
+  %gep.b = getelementptr inbounds i8, ptr %b, i64 %iv
+  %load.b = load i8, ptr %gep.b, align 1
+  %ext.b = zext i8 %load.b to i32
+  %mul = mul nsw i32 %ext.a, %ext.b
+  %sum.next = add nsw i32 %sum, %mul
+  %iv.next = add nuw nsw i64 %iv, 1
+  %exitcond = icmp eq i64 %iv.next, 1024
+  br i1 %exitcond, label %exit, label %loop, !llvm.loop !0
+
+exit:
+  ret i32 %sum.next
+}
+
+define i32 @dot_s8s8_vf64(ptr readonly %a, ptr readonly %b) #1 {
+; CHECK-LABEL: 'dot_s8s8_vf64'
+; CHECK:  Cost of 1 for VF 64: EXPRESSION vp<[[VP8:%[0-9]+]]> = ir<%sum> + partial.reduce.add (mul nsw (ir<%load.a> sext to i32), (ir<%load.b> sext to i32))
+; CHECK:  Cost of 1 for VF 64: EXPRESSION vp<[[VP8]]> = ir<%sum> + partial.reduce.add (mul nsw (ir<%load.a> sext to i32), (ir<%load.b> sext to i32))
+;
+entry:
+  br label %loop
+
+loop:
+  %iv = phi i64 [ 0, %entry ], [ %iv.next, %loop ]
+  %sum = phi i32 [ 0, %entry ], [ %sum.next, %loop ]
+  %gep.a = getelementptr inbounds i8, ptr %a, i64 %iv
+  %load.a = load i8, ptr %gep.a, align 1
+  %ext.a = sext i8 %load.a to i32
+  %gep.b = getelementptr inbounds i8, ptr %b, i64 %iv
+  %load.b = load i8, ptr %gep.b, align 1
+  %ext.b = sext i8 %load.b to i32
+  %mul = mul nsw i32 %ext.a, %ext.b
+  %sum.next = add nsw i32 %sum, %mul
+  %iv.next = add nuw nsw i64 %iv, 1
+  %exitcond = icmp eq i64 %iv.next, 1024
+  br i1 %exitcond, label %exit, label %loop, !llvm.loop !1
+
+exit:
+  ret i32 %sum.next
+}
+
+define i32 @dot_u8u8_vf64(ptr readonly %a, ptr readonly %b) #1 {
+; CHECK-LABEL: 'dot_u8u8_vf64'
+; CHECK:  Cost of 1 for VF 64: EXPRESSION vp<[[VP8:%[0-9]+]]> = ir<%sum> + partial.reduce.add (mul (ir<%load.a> zext to i32), (ir<%load.b> zext to i32))
+; CHECK:  Cost of 1 for VF 64: EXPRESSION vp<[[VP8]]> = ir<%sum> + partial.reduce.add (mul (ir<%load.a> zext to i32), (ir<%load.b> zext to i32))
+;
+entry:
+  br label %loop
+
+loop:
+  %iv = phi i64 [ 0, %entry ], [ %iv.next, %loop ]
+  %sum = phi i32 [ 0, %entry ], [ %sum.next, %loop ]
+  %gep.a = getelementptr inbounds i8, ptr %a, i64 %iv
+  %load.a = load i8, ptr %gep.a, align 1
+  %ext.a = zext i8 %load.a to i32
+  %gep.b = getelementptr inbounds i8, ptr %b, i64 %iv
+  %load.b = load i8, ptr %gep.b, align 1
+  %ext.b = zext i8 %load.b to i32
+  %mul = mul i32 %ext.a, %ext.b
+  %sum.next = add i32 %sum, %mul
+  %iv.next = add nuw nsw i64 %iv, 1
+  %exitcond = icmp eq i64 %iv.next, 1024
+  br i1 %exitcond, label %exit, label %loop, !llvm.loop !2
+
+exit:
+  ret i32 %sum.next
+}
+
+define i32 @dot_s16s16_vf32(ptr readonly %a, ptr readonly %b) #0 {
+; CHECK-LABEL: 'dot_s16s16_vf32'
+; CHECK:  Cost of 1 for VF 32: EXPRESSION vp<[[VP8:%[0-9]+]]> = ir<%sum> + partial.reduce.add (mul nsw (ir<%load.a> sext to i32), (ir<%load.b> sext to i32))
+; CHECK:  Cost of 1 for VF 32: EXPRESSION vp<[[VP8]]> = ir<%sum> + partial.reduce.add (mul nsw (ir<%load.a> sext to i32), (ir<%load.b> sext to i32))
+;
+entry:
+  br label %loop
+
+loop:
+  %iv = phi i64 [ 0, %entry ], [ %iv.next, %loop ]
+  %sum = phi i32 [ 0, %entry ], [ %sum.next, %loop ]
+  %gep.a = getelementptr inbounds i16, ptr %a, i64 %iv
+  %load.a = load i16, ptr %gep.a, align 2
+  %ext.a = sext i16 %load.a to i32
+  %gep.b = getelementptr inbounds i16, ptr %b, i64 %iv
+  %load.b = load i16, ptr %gep.b, align 2
+  %ext.b = sext i16 %load.b to i32
+  %mul = mul nsw i32 %ext.a, %ext.b
+  %sum.next = add nsw i32 %sum, %mul
+  %iv.next = add nuw nsw i64 %iv, 1
+  %exitcond = icmp eq i64 %iv.next, 1024
+  br i1 %exitcond, label %exit, label %loop, !llvm.loop !3
+
+exit:
+  ret i32 %sum.next
+}
+
+define i32 @dot_s16u16_vf32(ptr readonly %a, ptr readonly %b) #2 {
+; CHECK-LABEL: 'dot_s16u16_vf32'
+; CHECK:  Cost of 1 for VF 32: EXPRESSION vp<[[VP8:%[0-9]+]]> = ir<%sum> + partial.reduce.add (mul nsw (ir<%load.a> sext to i32), (ir<%load.b> zext to i32))
+; CHECK:  Cost of 1 for VF 32: EXPRESSION vp<[[VP8]]> = ir<%sum> + partial.reduce.add (mul nsw (ir<%load.a> sext to i32), (ir<%load.b> zext to i32))
+;
+entry:
+  br label %loop
+
+loop:
+  %iv = phi i64 [ 0, %entry ], [ %iv.next, %loop ]
+  %sum = phi i32 [ 0, %entry ], [ %sum.next, %loop ]
+  %gep.a = getelementptr inbounds i16, ptr %a, i64 %iv
+  %load.a = load i16, ptr %gep.a, align 2
+  %ext.a = sext i16 %load.a to i32
+  %gep.b = getelementptr inbounds i16, ptr %b, i64 %iv
+  %load.b = load i16, ptr %gep.b, align 2
+  %ext.b = zext i16 %load.b to i32
+  %mul = mul nsw i32 %ext.a, %ext.b
+  %sum.next = add nsw i32 %sum, %mul
+  %iv.next = add nuw nsw i64 %iv, 1
+  %exitcond = icmp eq i64 %iv.next, 1024
+  br i1 %exitcond, label %exit, label %loop, !llvm.loop !4
+
+exit:
+  ret i32 %sum.next
+}
+
+define i32 @dot_u16u16_vf32(ptr readonly %a, ptr readonly %b) #2 {
+; CHECK-LABEL: 'dot_u16u16_vf32'
+; CHECK:  Cost of 1 for VF 32: EXPRESSION vp<[[VP8:%[0-9]+]]> = ir<%sum> + partial.reduce.add (mul (ir<%load.a> zext to i32), (ir<%load.b> zext to i32))
+; CHECK:  Cost of 1 for VF 32: EXPRESSION vp<[[VP8]]> = ir<%sum> + partial.reduce.add (mul (ir<%load.a> zext to i32), (ir<%load.b> zext to i32))
+;
+entry:
+  br label %loop
+
+loop:
+  %iv = phi i64 [ 0, %entry ], [ %iv.next, %loop ]
+  %sum = phi i32 [ 0, %entry ], [ %sum.next, %loop ]
+  %gep.a = getelementptr inbounds i16, ptr %a, i64 %iv
+  %load.a = load i16, ptr %gep.a, align 2
+  %ext.a = zext i16 %load.a to i32
+  %gep.b = getelementptr inbounds i16, ptr %b, i64 %iv
+  %load.b = load i16, ptr %gep.b, align 2
+  %ext.b = zext i16 %load.b to i32
+  %mul = mul i32 %ext.a, %ext.b
+  %sum.next = add i32 %sum, %mul
+  %iv.next = add nuw nsw i64 %iv, 1
+  %exitcond = icmp eq i64 %iv.next, 1024
+  br i1 %exitcond, label %exit, label %loop, !llvm.loop !5
+
+exit:
+  ret i32 %sum.next
+}
+
+define float @dot_bf16_vf32(ptr readonly %a, ptr readonly %b) #3 {
+; CHECK-LABEL: 'dot_bf16_vf32'
+; CHECK:  Cost of 1 for VF 32: EXPRESSION vp<[[VP8:%[0-9]+]]> = ir<%sum> + partial.reduce.fadd (mul contract (ir<%load.a> fpext to float), (ir<%load.b> fpext to float))
+; CHECK:  Cost of 1 for VF 32: EXPRESSION vp<[[VP8]]> = ir<%sum> + partial.reduce.fadd (mul contract (ir<%load.a> fpext to float), (ir<%load.b> fpext to float))
+;
+entry:
+  br label %loop
+
+loop:
+  %iv = phi i64 [ 0, %entry ], [ %iv.next, %loop ]
+  %sum = phi float [ 0.0, %entry ], [ %sum.next, %loop ]
+  %gep.a = getelementptr inbounds bfloat, ptr %a, i64 %iv
+  %load.a = load bfloat, ptr %gep.a, align 2
+  %ext.a = fpext bfloat %load.a to float
+  %gep.b = getelementptr inbounds bfloat, ptr %b, i64 %iv
+  %load.b = load bfloat, ptr %gep.b, align 2
+  %ext.b = fpext bfloat %load.b to float
+  %mul = fmul contract float %ext.a, %ext.b
+  %sum.next = fadd reassoc contract float %sum, %mul
+  %iv.next = add nuw nsw i64 %iv, 1
+  %exitcond = icmp eq i64 %iv.next, 1024
+  br i1 %exitcond, label %exit, label %loop, !llvm.loop !6
+
+exit:
+  ret float %sum.next
+}
+
+attributes #0 = { "target-features"="+avx2,+avx512bw,+avx512f,+avx512vl,+avxvnni,-avx512vnni" }
+attributes #1 = { "target-features"="+avx2,+avx512bw,+avx512f,+avx512vl,+avxvnniint8,-avx10.2" }
+attributes #2 = { "target-features"="+avx2,+avx512bw,+avx512f,+avx512vl,+avxvnniint16,-avx10.2" }
+attributes #3 = { "target-features"="+avx512bf16,+avx512bw,+avx512f,+avx512vl" }
+
+!0 = distinct !{!0, !7, !8}
+!1 = distinct !{!1, !7, !8}
+!2 = distinct !{!2, !7, !8}
+!3 = distinct !{!3, !7, !9}
+!4 = distinct !{!4, !7, !9}
+!5 = distinct !{!5, !7, !9}
+!6 = distinct !{!6, !7, !9}
+!7 = !{!"llvm.loop.interleave.count", i32 1}
+!8 = !{!"llvm.loop.vectorize.width", i32 64}
+!9 = !{!"llvm.loop.vectorize.width", i32 32}
diff --git a/llvm/test/CodeGen/X86/vector-partial-reduce-avx512-prefer-256.ll b/llvm/test/CodeGen/X86/vector-partial-reduce-avx512-prefer-256.ll
new file mode 100644
index 0000000000000..70475b40ee0ce
--- /dev/null
+++ b/llvm/test/CodeGen/X86/vector-partial-reduce-avx512-prefer-256.ll
@@ -0,0 +1,76 @@
+; NOTE: Assertions have been autogenerated by utils/update_llc_test_checks.py UTC_ARGS: --version 6
+; RUN: llc -mtriple=x86_64-unknown-linux-gnu -mattr=+avx10.2-512 < %s | FileCheck %s
+
+; When 512-bit vectors are disabled by the preferred vector width, generic
+; type legalization should split partial reductions to legal 256-bit halves.
+
+define <16 x i32> @partial_reduce_sumla_i8_v16i32_prefer256(<16 x i32> %acc, <64 x i8> %a, <64 x i8> %b) #0 {
+; CHECK-LABEL: partial_reduce_sumla_i8_v16i32_prefer256:
+; CHECK:       # %bb.0:
+; CHECK-NEXT:    vpdpbusd %ymm2, %ymm4, %ymm0
+; CHECK-NEXT:    vpdpbusd %ymm3, %ymm5, %ymm1
+; CHECK-NEXT:    retq
+  %a.sext = sext <64 x i8> %a to <64 x i32>
+  %b.zext = zext <64 x i8> %b to <64 x i32>
+  %mul = mul nsw <64 x i32> %a.sext, %b.zext
+  %res = call <16 x i32> @llvm.vector.partial.reduce.add.v16i32.v64i32(<16 x i32> %acc, <64 x i32> %mul)
+  ret <16 x i32> %res
+}
+
+define <16 x i32> @partial_reduce_smla_i16_v16i32_prefer256(<16 x i32> %acc, <32 x i16> %a, <32 x i16> %b) #0 {
+; CHECK-LABEL: partial_reduce_smla_i16_v16i32_prefer256:
+; CHECK:       # %bb.0:
+; CHECK-NEXT:    vpdpwssd %ymm4, %ymm2, %ymm0
+; CHECK-NEXT:    vpdpwssd %ymm5, %ymm3, %ymm1
+; CHECK-NEXT:    retq
+  %a.sext = sext <32 x i16> %a to <32 x i32>
+  %b.sext = sext <32 x i16> %b to <32 x i32>
+  %mul = mul nsw <32 x i32> %a.sext, %b.sext
+  %res = call <16 x i32> @llvm.vector.partial.reduce.add.v16i32.v32i32(<16 x i32> %acc, <32 x i32> %mul)
+  ret <16 x i32> %res
+}
+
+define <16 x i32> @partial_reduce_smla_i8_v16i32_prefer256(<16 x i32> %acc, <64 x i8> %a, <64 x i8> %b) #0 {
+; CHECK-LABEL: partial_reduce_smla_i8_v16i32_prefer256:
+; CHECK:       # %bb.0:
+; CHECK-NEXT:    vpdpbssd %ymm4, %ymm2, %ymm0
+; CHECK-NEXT:    vpdpbssd %ymm5, %ymm3, %ymm1
+; CHECK-NEXT:    retq
+  %a.sext = sext <64 x i8> %a to <64 x i32>
+  %b.sext = sext <64 x i8> %b to <64 x i32>
+  %mul = mul nsw <64 x i32> %a.sext, %b.sext
+  %res = call <16 x i32> @llvm.vector.partial.reduce.add.v16i32.v64i32(<16 x i32> %acc, <64 x i32> %mul)
+  ret <16 x i32> %res
+}
+
+define <16 x i32> @partial_reduce_sumla_i16_v16i32_prefer256(<16 x i32> %acc, <32 x i16> %a, <32 x i16> %b) #0 {
+; CHECK-LABEL: partial_reduce_sumla_i16_v16i32_prefer256:
+; CHECK:       # %bb.0:
+; CHECK-NEXT:    vpdpwsud %ymm4, %ymm2, %ymm0
+; CHECK-NEXT:    vpdpwsud %ymm5, %ymm3, %ymm1
+; CHECK-NEXT:    retq
+  %a.sext = sext <32 x i16> %a to <32 x i32>
+  %b.zext = zext <32 x i16> %b to <32 x i32>
+  %mul = mul nsw <32 x i32> %a.sext, %b.zext
+  %res = call <16 x i32> @llvm.vector.partial.reduce.add.v16i32.v32i32(<16 x i32> %acc, <32 x i32> %mul)
+  ret <16 x i32> %res
+}
+
+define <16 x float> @partial_reduce_fmla_bf16_v16f32_prefer256(<16 x float> %acc, <32 x bfloat> %a, <32 x bfloat> %b) #0 {
+; CHECK-LABEL: partial_reduce_fmla_bf16_v16f32_prefer256:
+; CHECK:       # %bb.0:
+; CHECK-NEXT:    vdpbf16ps %ymm4, %ymm2, %ymm0
+; CHECK-NEXT:    vdpbf16ps %ymm5, %ymm3, %ymm1
+; CHECK-NEXT:    retq
+  %a.ext = fpext <32 x bfloat> %a to <32 x float>
+  %b.ext = fpext <32 x bfloat> %b to <32 x float>
+  %mul = fmul <32 x float> %a.ext, %b.ext
+  %res = call <16 x float> @llvm.vector.partial.reduce.fadd.v16f32.v32f32(<16 x float> %acc, <32 x float> %mul)
+  ret <16 x float> %res
+}
+
+declare <16 x i32> @llvm.vector.partial.reduce.add.v16i32.v64i32(<16 x i32>, <64 x i32>)
+declare <16 x i32> @llvm.vector.partial.reduce.add.v16i32.v32i32(<16 x i32>, <32 x i32>)
+declare <16 x float> @llvm.vector.partial.reduce.fadd.v16f32.v32f32(<16 x float>, <32 x float>)
+
+attributes #0 = { "min-legal-vector-width"="256" "prefer-vector-width"="256" }
diff --git a/llvm/test/CodeGen/X86/vector-partial-reduce-avx512-vex-split-int16.ll b/llvm/test/CodeGen/X86/vector-partial-reduce-avx512-vex-split-int16.ll
new file mode 100644
index 0000000000000..443cd10ce8086
--- /dev/null
+++ b/llvm/test/CodeGen/X86/vector-partial-reduce-avx512-vex-split-int16.ll
@@ -0,0 +1,42 @@
+; NOTE: Assertions have been autogenerated by utils/update_llc_test_checks.py UTC_ARGS: --version 6
+; RUN: llc -mtriple=x86_64-unknown-linux-gnu -mattr=+avx512f,+avx512bw,+avx512vl,+avxvnniint16,-avx10.2 < %s | FileCheck %s
+
+; Targets with AVX512 registers but only 256-bit AVX-VNNI-INT16 instructions
+; should split 512-bit partial reductions to two 256-bit dot products instead
+; of expanding them as zmm arithmetic.
+
+define <16 x i32> @partial_reduce_sumla_i16_v16i32(<16 x i32> %acc, <32 x i16> %a, <32 x i16> %b) {
+; CHECK-LABEL: partial_reduce_sumla_i16_v16i32:
+; CHECK:       # %bb.0:
+; CHECK-NEXT:    vextractf64x4 $1, %zmm2, %ymm3
+; CHECK-NEXT:    vextractf64x4 $1, %zmm1, %ymm4
+; CHECK-NEXT:    vextractf64x4 $1, %zmm0, %ymm5
+; CHECK-NEXT:    vpdpwsud %ymm3, %ymm4, %ymm5
+; CHECK-NEXT:    vpdpwsud %ymm2, %ymm1, %ymm0
+; CHECK-NEXT:    vinsertf64x4 $1, %ymm5, %zmm0, %zmm0
+; CHECK-NEXT:    retq
+  %a.sext = sext <32 x i16> %a to <32 x i32>
+  %b.zext = zext <32 x i16> %b to <32 x i32>
+  %mul = mul nsw <32 x i32> %a.sext, %b.zext
+  %res = call <16 x i32> @llvm.vector.partial.reduce.add.v16i32.v32i32(<16 x i32> %acc, <32 x i32> %mul)
+  ret <16 x i32> %res
+}
+
+define <16 x i32> @partial_reduce_umla_i16_v16i32(<16 x i32> %acc, <32 x i16> %a, <32 x i16> %b) {
+; CHECK-LABEL: partial_reduce_umla_i16_v16i32:
+; CHECK:       # %bb.0:
+; CHECK-NEXT:    vextractf64x4 $1, %zmm2, %ymm3
+; CHECK-NEXT:    vextractf64x4 $1, %zmm1, %ymm4
+; CHECK-NEXT:    vextractf64x4 $1, %zmm0, %ymm5
+; CHECK-NEXT:    vpdpwuud %ymm3, %ymm4, %ymm5
+; CHECK-NEXT:    vpdpwuud %ymm2, %ymm1, %ymm0
+; CHECK-NEXT:    vinsertf64x4 $1, %ymm5, %zmm0, %zmm0
+; CHECK-NEXT:    retq
+  %a.zext = zext <32 x i16> %a to <32 x i32>
+  %b.zext = zext <32 x i16> %b to <32 x i32>
+  %mul = mul <32 x i32> %a.zext, %b.zext
+  %res = call <16 x i32> @llvm.vector.partial.reduce.add.v16i32.v32i32(<16 x i32> %acc, <32 x i32> %mul)
+  ret <16 x i32> %res
+}
+
+declare <16 x i32> @llvm.vector.partial.reduce.add.v16i32.v32i32(<16 x i32>, <32 x i32>)
diff --git a/llvm/test/CodeGen/X86/vector-partial-reduce-avx512-vex-split-int8.ll b/llvm/test/CodeGen/X86/vector-partial-reduce-avx512-vex-split-int8.ll
new file mode 100644
index 0000000000000..550e63957b14f
--- /dev/null
+++ b/llvm/test/CodeGen/X86/vector-partial-reduce-avx512-vex-split-int8.ll
@@ -0,0 +1,42 @@
+; NOTE: Assertions have been autogenerated by utils/update_llc_test_checks.py UTC_ARGS: --version 6
+; RUN: llc -mtriple=x86_64-unknown-linux-gnu -mattr=+avx512f,+avx512bw,+avx512vl,+avxvnniint8,-avx10.2 < %s | FileCheck %s
+
+; Targets with AVX512 registers but only 256-bit AVX-VNNI-INT8 instructions
+; should split 512-bit partial reductions to two 256-bit dot products instead
+; of expanding them as zmm arithmetic.
+
+define <16 x i32> @partial_reduce_smla_i8_v16i32(<16 x i32> %acc, <64 x i8> %a, <64 x i8> %b) {
+; CHECK-LABEL: partial_reduce_smla_i8_v16i32:
+; CHECK:       # %bb.0:
+; CHECK-NEXT:    vextractf64x4 $1, %zmm2, %ymm3
+; CHECK-NEXT:    vextractf64x4 $1, %zmm1, %ymm4
+; CHECK-NEXT:    vextractf64x4 $1, %zmm0, %ymm5
+; CHECK-NEXT:    vpdpbssd %ymm3, %ymm4, %ymm5
+; CHECK-NEXT:    vpdpbssd %ymm2, %ymm1, %ymm0
+; CHECK-NEXT:    vinsertf64x4 $1, %ymm5, %zmm0, %zmm0
+; CHECK-NEXT:    retq
+  %a.sext = sext <64 x i8> %a to <64 x i32>
+  %b.sext = sext <64 x i8> %b to <64 x i32>
+  %mul = mul nsw <64 x i32> %a.sext, %b.sext
+  %res = call <16 x i32> @llvm.vector.partial.reduce.add.v16i32.v64i32(<16 x i32> %acc, <64 x i32> %mul)
+  ret <16 x i32> %res
+}
+
+define <16 x i32> @partial_reduce_umla_i8_v16i32(<16 x i32> %acc, <64 x i8> %a, <64 x i8> %b) {
+; CHECK-LABEL: partial_reduce_umla_i8_v16i32:
+; CHECK:       # %bb.0:
+; CHECK-NEXT:    vextractf64x4 $1, %zmm2, %ymm3
+; CHECK-NEXT:    vextractf64x4 $1, %zmm1, %ymm4
+; CHECK-NEXT:    vextractf64x4 $1, %zmm0, %ymm5
+; CHECK-NEXT:    vpdpbuud %ymm3, %ymm4, %ymm5
+; CHECK-NEXT:    vpdpbuud %ymm2, %ymm1, %ymm0
+; CHECK-NEXT:    vinsertf64x4 $1, %ymm5, %zmm0, %zmm0
+; CHECK-NEXT:    retq
+  %a.zext = zext <64 x i8> %a to <64 x i32>
+  %b.zext = zext <64 x i8> %b to <64 x i32>
+  %mul = mul <64 x i32> %a.zext, %b.zext
+  %res = call <16 x i32> @llvm.vector.partial.reduce.add.v16i32.v64i32(<16 x i32> %acc, <64 x i32> %mul)
+  ret <16 x i32> %res
+}
+
+declare <16 x i32> @llvm.vector.partial.reduce.add.v16i32.v64i32(<16 x i32>, <64 x i32>)
diff --git a/llvm/test/CodeGen/X86/vector-partial-reduce-avx512-vex-split.ll b/llvm/test/CodeGen/X86/vector-partial-reduce-avx512-vex-split.ll
new file mode 100644
index 0000000000000..a0eb2249c0f2b
--- /dev/null
+++ b/llvm/test/CodeGen/X86/vector-partial-reduce-avx512-vex-split.ll
@@ -0,0 +1,43 @@
+; NOTE: Assertions have been autogenerated by utils/update_llc_test_checks.py UTC_ARGS: --version 6
+; RUN: llc -mtriple=x86_64-unknown-linux-gnu -mattr=+avx512f,+avx512bw,+avx512vl,+avxvnni,-avx512vnni,+fast-dpwssd < %s | FileCheck %s
+
+; Targets with AVX512 registers but only 256-bit VEX dot-product instructions
+; should split 512-bit partial reductions to two 256-bit dot products instead
+; of expanding them as zmm arithmetic.
+
+define <16 x i32> @partial_reduce_sumla_i8_v16i32(<16 x i32> %acc, <64 x i8> %a, <64 x i8> %b) {
+; CHECK-LABEL: partial_reduce_sumla_i8_v16i32:
+; CHECK:       # %bb.0:
+; CHECK-NEXT:    vextracti64x4 $1, %zmm2, %ymm3
+; CHECK-NEXT:    vextracti64x4 $1, %zmm1, %ymm4
+; CHECK-NEXT:    vextracti64x4 $1, %zmm0, %ymm5
+; CHECK-NEXT:    {vex} vpdpbusd %ymm3, %ymm4, %ymm5
+; CHECK-NEXT:    {vex} vpdpbusd %ymm2, %ymm1, %ymm0
+; CHECK-NEXT:    vinserti64x4 $1, %ymm5, %zmm0, %zmm0
+; CHECK-NEXT:    retq
+  %a.zext = zext <64 x i8> %a to <64 x i32>
+  %b.sext = sext <64 x i8> %b to <64 x i32>
+  %mul = mul nsw <64 x i32> %a.zext, %b.sext
+  %res = call <16 x i32> @llvm.vector.partial.reduce.add.v16i32.v64i32(<16 x i32> %acc, <64 x i32> %mul)
+  ret <16 x i32> %res
+}
+
+define <16 x i32> @partial_reduce_smla_i16_v16i32(<16 x i32> %acc, <32 x i16> %a, <32 x i16> %b) {
+; CHECK-LABEL: partial_reduce_smla_i16_v16i32:
+; CHECK:       # %bb.0:
+; CHECK-NEXT:    vextracti64x4 $1, %zmm2, %ymm3
+; CHECK-NEXT:    vextracti64x4 $1, %zmm1, %ymm4
+; CHECK-NEXT:    vextracti64x4 $1, %zmm0, %ymm5
+; CHECK-NEXT:    {vex} vpdpwssd %ymm3, %ymm4, %ymm5
+; CHECK-NEXT:    {vex} vpdpwssd %ymm2, %ymm1, %ymm0
+; CHECK-NEXT:    vinserti64x4 $1, %ymm5, %zmm0, %zmm0
+; CHECK-NEXT:    retq
+  %a.sext = sext <32 x i16> %a to <32 x i32>
+  %b.sext = sext <32 x i16> %b to <32 x i32>
+  %mul = mul nsw <32 x i32> %a.sext, %b.sext
+  %res = call <16 x i32> @llvm.vector.partial.reduce.add.v16i32.v32i32(<16 x i32> %acc, <32 x i32> %mul)
+  ret <16 x i32> %res
+}
+
+declare <16 x i32> @llvm.vector.partial.reduce.add.v16i32.v64i32(<16 x i32>, <64 x i32>)
+declare <16 x i32> @llvm.vector.partial.reduce.add.v16i32.v32i32(<16 x i32>, <32 x i32>)
diff --git a/llvm/test/CodeGen/X86/vector-partial-reduce-dot-product.ll b/llvm/test/CodeGen/X86/vector-partial-reduce-dot-product.ll
new file mode 100644
index 0000000000000..c06301ac60cc9
--- /dev/null
+++ b/llvm/test/CodeGen/X86/vector-partial-reduce-dot-product.ll
@@ -0,0 +1,1032 @@
+; NOTE: Assertions have been autogenerated by utils/update_llc_test_checks.py UTC_ARGS: --version 5
+; RUN: llc -mtriple=x86_64-unknown-linux-gnu -mattr=+avx512f,+avx512bw,+avx512vl,+avx512vnni,+avx512bf16 < %s | FileCheck %s --check-prefixes=AVX512VNNI
+; RUN: llc -mtriple=x86_64-unknown-linux-gnu -mattr=+avx,+avx2,+fma,+avxvnni < %s | FileCheck %s --check-prefixes=AVXVNNI
+; RUN: llc -mtriple=x86_64-unknown-linux-gnu -mattr=+avx512f,+avx512bw,+avx512vl,+avx512vnni,+avx512bf16,+avxvnniint8,+avxvnniint16 < %s | FileCheck %s --check-prefixes=AVXVNNIINT8INT16
+; RUN: llc -mtriple=x86_64-unknown-linux-gnu -mattr=+avx10.2-512 < %s | FileCheck %s --check-prefixes=AVX10
+
+; Test that PARTIAL_REDUCE_SUMLA (sext × zext) lowers to vpdpbusd.
+define <4 x i32> @partial_reduce_sumla_v4i32(<4 x i32> %acc, <16 x i8> %a, <16 x i8> %b) {
+; AVX512VNNI-LABEL: partial_reduce_sumla_v4i32:
+; AVX512VNNI:       # %bb.0:
+; AVX512VNNI-NEXT:    vpdpbusd %xmm2, %xmm1, %xmm0
+; AVX512VNNI-NEXT:    retq
+;
+; AVXVNNI-LABEL: partial_reduce_sumla_v4i32:
+; AVXVNNI:       # %bb.0:
+; AVXVNNI-NEXT:    {vex} vpdpbusd %xmm2, %xmm1, %xmm0
+; AVXVNNI-NEXT:    retq
+;
+; AVXVNNIINT8INT16-LABEL: partial_reduce_sumla_v4i32:
+; AVXVNNIINT8INT16:       # %bb.0:
+; AVXVNNIINT8INT16-NEXT:    vpdpbusd %xmm2, %xmm1, %xmm0
+; AVXVNNIINT8INT16-NEXT:    retq
+;
+; AVX10-LABEL: partial_reduce_sumla_v4i32:
+; AVX10:       # %bb.0:
+; AVX10-NEXT:    vpdpbusd %xmm2, %xmm1, %xmm0
+; AVX10-NEXT:    retq
+  %a.zext = zext <16 x i8> %a to <16 x i32>
+  %b.sext = sext <16 x i8> %b to <16 x i32>
+  %mul = mul nsw <16 x i32> %a.zext, %b.sext
+  %res = call <4 x i32> @llvm.vector.partial.reduce.add.v4i32.v16i32(<4 x i32> %acc, <16 x i32> %mul)
+  ret <4 x i32> %res
+}
+
+; Test 512-bit (zmm) partial reduction.
+define <16 x i32> @partial_reduce_sumla_v16i32(<16 x i32> %acc, <64 x i8> %a, <64 x i8> %b) {
+; AVX512VNNI-LABEL: partial_reduce_sumla_v16i32:
+; AVX512VNNI:       # %bb.0:
+; AVX512VNNI-NEXT:    vpdpbusd %zmm2, %zmm1, %zmm0
+; AVX512VNNI-NEXT:    retq
+;
+; AVXVNNI-LABEL: partial_reduce_sumla_v16i32:
+; AVXVNNI:       # %bb.0:
+; AVXVNNI-NEXT:    {vex} vpdpbusd %ymm4, %ymm2, %ymm0
+; AVXVNNI-NEXT:    {vex} vpdpbusd %ymm5, %ymm3, %ymm1
+; AVXVNNI-NEXT:    retq
+;
+; AVXVNNIINT8INT16-LABEL: partial_reduce_sumla_v16i32:
+; AVXVNNIINT8INT16:       # %bb.0:
+; AVXVNNIINT8INT16-NEXT:    vpdpbusd %zmm2, %zmm1, %zmm0
+; AVXVNNIINT8INT16-NEXT:    retq
+;
+; AVX10-LABEL: partial_reduce_sumla_v16i32:
+; AVX10:       # %bb.0:
+; AVX10-NEXT:    vpdpbusd %zmm2, %zmm1, %zmm0
+; AVX10-NEXT:    retq
+  %a.zext = zext <64 x i8> %a to <64 x i32>
+  %b.sext = sext <64 x i8> %b to <64 x i32>
+  %mul = mul nsw <64 x i32> %a.zext, %b.sext
+  %res = call <16 x i32> @llvm.vector.partial.reduce.add.v16i32.v64i32(<16 x i32> %acc, <64 x i32> %mul)
+  ret <16 x i32> %res
+}
+
+; Test 256-bit (ymm) partial reduction.
+define <8 x i32> @partial_reduce_sumla_v8i32(<8 x i32> %acc, <32 x i8> %a, <32 x i8> %b) {
+; AVX512VNNI-LABEL: partial_reduce_sumla_v8i32:
+; AVX512VNNI:       # %bb.0:
+; AVX512VNNI-NEXT:    vpdpbusd %ymm2, %ymm1, %ymm0
+; AVX512VNNI-NEXT:    retq
+;
+; AVXVNNI-LABEL: partial_reduce_sumla_v8i32:
+; AVXVNNI:       # %bb.0:
+; AVXVNNI-NEXT:    {vex} vpdpbusd %ymm2, %ymm1, %ymm0
+; AVXVNNI-NEXT:    retq
+;
+; AVXVNNIINT8INT16-LABEL: partial_reduce_sumla_v8i32:
+; AVXVNNIINT8INT16:       # %bb.0:
+; AVXVNNIINT8INT16-NEXT:    vpdpbusd %ymm2, %ymm1, %ymm0
+; AVXVNNIINT8INT16-NEXT:    retq
+;
+; AVX10-LABEL: partial_reduce_sumla_v8i32:
+; AVX10:       # %bb.0:
+; AVX10-NEXT:    vpdpbusd %ymm2, %ymm1, %ymm0
+; AVX10-NEXT:    retq
+  %a.zext = zext <32 x i8> %a to <32 x i32>
+  %b.sext = sext <32 x i8> %b to <32 x i32>
+  %mul = mul nsw <32 x i32> %a.zext, %b.sext
+  %res = call <8 x i32> @llvm.vector.partial.reduce.add.v8i32.v32i32(<8 x i32> %acc, <32 x i32> %mul)
+  ret <8 x i32> %res
+}
+
+declare <4 x i32> @llvm.vector.partial.reduce.add.v4i32.v16i32(<4 x i32>, <16 x i32>)
+declare <8 x i32> @llvm.vector.partial.reduce.add.v8i32.v32i32(<8 x i32>, <32 x i32>)
+declare <16 x i32> @llvm.vector.partial.reduce.add.v16i32.v64i32(<16 x i32>, <64 x i32>)
+
+; i16 x i16 -> i32 partial reduction tests (vpdpwssd)
+
+; Test 128-bit (xmm) i16 partial reduction: 8 x i16 -> 4 x i32.
+define <4 x i32> @partial_reduce_smla_i16_v4i32(<4 x i32> %acc, <8 x i16> %a, <8 x i16> %b) {
+; AVX512VNNI-LABEL: partial_reduce_smla_i16_v4i32:
+; AVX512VNNI:       # %bb.0:
+; AVX512VNNI-NEXT:    vpdpwssd %xmm2, %xmm1, %xmm0
+; AVX512VNNI-NEXT:    retq
+;
+; AVXVNNI-LABEL: partial_reduce_smla_i16_v4i32:
+; AVXVNNI:       # %bb.0:
+; AVXVNNI-NEXT:    {vex} vpdpwssd %xmm2, %xmm1, %xmm0
+; AVXVNNI-NEXT:    retq
+;
+; AVXVNNIINT8INT16-LABEL: partial_reduce_smla_i16_v4i32:
+; AVXVNNIINT8INT16:       # %bb.0:
+; AVXVNNIINT8INT16-NEXT:    vpdpwssd %xmm2, %xmm1, %xmm0
+; AVXVNNIINT8INT16-NEXT:    retq
+;
+; AVX10-LABEL: partial_reduce_smla_i16_v4i32:
+; AVX10:       # %bb.0:
+; AVX10-NEXT:    vpdpwssd %xmm2, %xmm1, %xmm0
+; AVX10-NEXT:    retq
+  %a.sext = sext <8 x i16> %a to <8 x i32>
+  %b.sext = sext <8 x i16> %b to <8 x i32>
+  %mul = mul nsw <8 x i32> %a.sext, %b.sext
+  %res = call <4 x i32> @llvm.vector.partial.reduce.add.v4i32.v8i32(<4 x i32> %acc, <8 x i32> %mul)
+  ret <4 x i32> %res
+}
+
+; Test 256-bit (ymm) i16 partial reduction: 16 x i16 -> 8 x i32.
+define <8 x i32> @partial_reduce_smla_i16_v8i32(<8 x i32> %acc, <16 x i16> %a, <16 x i16> %b) {
+; AVX512VNNI-LABEL: partial_reduce_smla_i16_v8i32:
+; AVX512VNNI:       # %bb.0:
+; AVX512VNNI-NEXT:    vpdpwssd %ymm2, %ymm1, %ymm0
+; AVX512VNNI-NEXT:    retq
+;
+; AVXVNNI-LABEL: partial_reduce_smla_i16_v8i32:
+; AVXVNNI:       # %bb.0:
+; AVXVNNI-NEXT:    {vex} vpdpwssd %ymm2, %ymm1, %ymm0
+; AVXVNNI-NEXT:    retq
+;
+; AVXVNNIINT8INT16-LABEL: partial_reduce_smla_i16_v8i32:
+; AVXVNNIINT8INT16:       # %bb.0:
+; AVXVNNIINT8INT16-NEXT:    vpdpwssd %ymm2, %ymm1, %ymm0
+; AVXVNNIINT8INT16-NEXT:    retq
+;
+; AVX10-LABEL: partial_reduce_smla_i16_v8i32:
+; AVX10:       # %bb.0:
+; AVX10-NEXT:    vpdpwssd %ymm2, %ymm1, %ymm0
+; AVX10-NEXT:    retq
+  %a.sext = sext <16 x i16> %a to <16 x i32>
+  %b.sext = sext <16 x i16> %b to <16 x i32>
+  %mul = mul nsw <16 x i32> %a.sext, %b.sext
+  %res = call <8 x i32> @llvm.vector.partial.reduce.add.v8i32.v16i32(<8 x i32> %acc, <16 x i32> %mul)
+  ret <8 x i32> %res
+}
+
+; Test 512-bit (zmm) i16 partial reduction: 32 x i16 -> 16 x i32.
+define <16 x i32> @partial_reduce_smla_i16_v16i32(<16 x i32> %acc, <32 x i16> %a, <32 x i16> %b) {
+; AVX512VNNI-LABEL: partial_reduce_smla_i16_v16i32:
+; AVX512VNNI:       # %bb.0:
+; AVX512VNNI-NEXT:    vpdpwssd %zmm2, %zmm1, %zmm0
+; AVX512VNNI-NEXT:    retq
+;
+; AVXVNNI-LABEL: partial_reduce_smla_i16_v16i32:
+; AVXVNNI:       # %bb.0:
+; AVXVNNI-NEXT:    {vex} vpdpwssd %ymm4, %ymm2, %ymm0
+; AVXVNNI-NEXT:    {vex} vpdpwssd %ymm5, %ymm3, %ymm1
+; AVXVNNI-NEXT:    retq
+;
+; AVXVNNIINT8INT16-LABEL: partial_reduce_smla_i16_v16i32:
+; AVXVNNIINT8INT16:       # %bb.0:
+; AVXVNNIINT8INT16-NEXT:    vpdpwssd %zmm2, %zmm1, %zmm0
+; AVXVNNIINT8INT16-NEXT:    retq
+;
+; AVX10-LABEL: partial_reduce_smla_i16_v16i32:
+; AVX10:       # %bb.0:
+; AVX10-NEXT:    vpdpwssd %zmm2, %zmm1, %zmm0
+; AVX10-NEXT:    retq
+  %a.sext = sext <32 x i16> %a to <32 x i32>
+  %b.sext = sext <32 x i16> %b to <32 x i32>
+  %mul = mul nsw <32 x i32> %a.sext, %b.sext
+  %res = call <16 x i32> @llvm.vector.partial.reduce.add.v16i32.v32i32(<16 x i32> %acc, <32 x i32> %mul)
+  ret <16 x i32> %res
+}
+
+; bf16 x bf16 -> f32 partial reduction tests (vdpbf16ps)
+
+; Test 128-bit (xmm) bf16 partial reduction: 8 x bf16 -> 4 x f32.
+define <4 x float> @partial_reduce_fmla_bf16_v4f32(<4 x float> %acc, <8 x bfloat> %a, <8 x bfloat> %b) {
+; AVX512VNNI-LABEL: partial_reduce_fmla_bf16_v4f32:
+; AVX512VNNI:       # %bb.0:
+; AVX512VNNI-NEXT:    vdpbf16ps %xmm2, %xmm1, %xmm0
+; AVX512VNNI-NEXT:    retq
+;
+; AVXVNNI-LABEL: partial_reduce_fmla_bf16_v4f32:
+; AVXVNNI:       # %bb.0:
+; AVXVNNI-NEXT:    vpmovzxwd {{.*#+}} ymm1 = xmm1[0],zero,xmm1[1],zero,xmm1[2],zero,xmm1[3],zero,xmm1[4],zero,xmm1[5],zero,xmm1[6],zero,xmm1[7],zero
+; AVXVNNI-NEXT:    vpslld $16, %ymm1, %ymm1
+; AVXVNNI-NEXT:    vpmovzxwd {{.*#+}} ymm2 = xmm2[0],zero,xmm2[1],zero,xmm2[2],zero,xmm2[3],zero,xmm2[4],zero,xmm2[5],zero,xmm2[6],zero,xmm2[7],zero
+; AVXVNNI-NEXT:    vpslld $16, %ymm2, %ymm2
+; AVXVNNI-NEXT:    vmulps %ymm2, %ymm1, %ymm1
+; AVXVNNI-NEXT:    vaddps %xmm1, %xmm0, %xmm0
+; AVXVNNI-NEXT:    vextractf128 $1, %ymm1, %xmm1
+; AVXVNNI-NEXT:    vaddps %xmm0, %xmm1, %xmm0
+; AVXVNNI-NEXT:    vzeroupper
+; AVXVNNI-NEXT:    retq
+;
+; AVXVNNIINT8INT16-LABEL: partial_reduce_fmla_bf16_v4f32:
+; AVXVNNIINT8INT16:       # %bb.0:
+; AVXVNNIINT8INT16-NEXT:    vdpbf16ps %xmm2, %xmm1, %xmm0
+; AVXVNNIINT8INT16-NEXT:    retq
+;
+; AVX10-LABEL: partial_reduce_fmla_bf16_v4f32:
+; AVX10:       # %bb.0:
+; AVX10-NEXT:    vdpbf16ps %xmm2, %xmm1, %xmm0
+; AVX10-NEXT:    retq
+  %a.ext = fpext <8 x bfloat> %a to <8 x float>
+  %b.ext = fpext <8 x bfloat> %b to <8 x float>
+  %mul = fmul <8 x float> %a.ext, %b.ext
+  %res = call <4 x float> @llvm.vector.partial.reduce.fadd.v4f32.v8f32(<4 x float> %acc, <8 x float> %mul)
+  ret <4 x float> %res
+}
+
+; Test 256-bit (ymm) bf16 partial reduction: 16 x bf16 -> 8 x f32.
+define <8 x float> @partial_reduce_fmla_bf16_v8f32(<8 x float> %acc, <16 x bfloat> %a, <16 x bfloat> %b) {
+; AVX512VNNI-LABEL: partial_reduce_fmla_bf16_v8f32:
+; AVX512VNNI:       # %bb.0:
+; AVX512VNNI-NEXT:    vdpbf16ps %ymm2, %ymm1, %ymm0
+; AVX512VNNI-NEXT:    retq
+;
+; AVXVNNI-LABEL: partial_reduce_fmla_bf16_v8f32:
+; AVXVNNI:       # %bb.0:
+; AVXVNNI-NEXT:    vpmovzxwd {{.*#+}} ymm3 = xmm1[0],zero,xmm1[1],zero,xmm1[2],zero,xmm1[3],zero,xmm1[4],zero,xmm1[5],zero,xmm1[6],zero,xmm1[7],zero
+; AVXVNNI-NEXT:    vpslld $16, %ymm3, %ymm3
+; AVXVNNI-NEXT:    vextracti128 $1, %ymm1, %xmm1
+; AVXVNNI-NEXT:    vpmovzxwd {{.*#+}} ymm1 = xmm1[0],zero,xmm1[1],zero,xmm1[2],zero,xmm1[3],zero,xmm1[4],zero,xmm1[5],zero,xmm1[6],zero,xmm1[7],zero
+; AVXVNNI-NEXT:    vpslld $16, %ymm1, %ymm1
+; AVXVNNI-NEXT:    vpmovzxwd {{.*#+}} ymm4 = xmm2[0],zero,xmm2[1],zero,xmm2[2],zero,xmm2[3],zero,xmm2[4],zero,xmm2[5],zero,xmm2[6],zero,xmm2[7],zero
+; AVXVNNI-NEXT:    vpslld $16, %ymm4, %ymm4
+; AVXVNNI-NEXT:    vmulps %ymm4, %ymm3, %ymm3
+; AVXVNNI-NEXT:    vextracti128 $1, %ymm2, %xmm2
+; AVXVNNI-NEXT:    vpmovzxwd {{.*#+}} ymm2 = xmm2[0],zero,xmm2[1],zero,xmm2[2],zero,xmm2[3],zero,xmm2[4],zero,xmm2[5],zero,xmm2[6],zero,xmm2[7],zero
+; AVXVNNI-NEXT:    vpslld $16, %ymm2, %ymm2
+; AVXVNNI-NEXT:    vmulps %ymm2, %ymm1, %ymm1
+; AVXVNNI-NEXT:    vaddps %ymm3, %ymm0, %ymm0
+; AVXVNNI-NEXT:    vaddps %ymm1, %ymm0, %ymm0
+; AVXVNNI-NEXT:    retq
+;
+; AVXVNNIINT8INT16-LABEL: partial_reduce_fmla_bf16_v8f32:
+; AVXVNNIINT8INT16:       # %bb.0:
+; AVXVNNIINT8INT16-NEXT:    vdpbf16ps %ymm2, %ymm1, %ymm0
+; AVXVNNIINT8INT16-NEXT:    retq
+;
+; AVX10-LABEL: partial_reduce_fmla_bf16_v8f32:
+; AVX10:       # %bb.0:
+; AVX10-NEXT:    vdpbf16ps %ymm2, %ymm1, %ymm0
+; AVX10-NEXT:    retq
+  %a.ext = fpext <16 x bfloat> %a to <16 x float>
+  %b.ext = fpext <16 x bfloat> %b to <16 x float>
+  %mul = fmul <16 x float> %a.ext, %b.ext
+  %res = call <8 x float> @llvm.vector.partial.reduce.fadd.v8f32.v16f32(<8 x float> %acc, <16 x float> %mul)
+  ret <8 x float> %res
+}
+
+; Test 512-bit (zmm) bf16 partial reduction: 32 x bf16 -> 16 x f32.
+define <16 x float> @partial_reduce_fmla_bf16_v16f32(<16 x float> %acc, <32 x bfloat> %a, <32 x bfloat> %b) {
+; AVX512VNNI-LABEL: partial_reduce_fmla_bf16_v16f32:
+; AVX512VNNI:       # %bb.0:
+; AVX512VNNI-NEXT:    vdpbf16ps %zmm2, %zmm1, %zmm0
+; AVX512VNNI-NEXT:    retq
+;
+; AVXVNNI-LABEL: partial_reduce_fmla_bf16_v16f32:
+; AVXVNNI:       # %bb.0:
+; AVXVNNI-NEXT:    vpmovzxwd {{.*#+}} ymm6 = xmm2[0],zero,xmm2[1],zero,xmm2[2],zero,xmm2[3],zero,xmm2[4],zero,xmm2[5],zero,xmm2[6],zero,xmm2[7],zero
+; AVXVNNI-NEXT:    vpslld $16, %ymm6, %ymm6
+; AVXVNNI-NEXT:    vextracti128 $1, %ymm2, %xmm2
+; AVXVNNI-NEXT:    vpmovzxwd {{.*#+}} ymm2 = xmm2[0],zero,xmm2[1],zero,xmm2[2],zero,xmm2[3],zero,xmm2[4],zero,xmm2[5],zero,xmm2[6],zero,xmm2[7],zero
+; AVXVNNI-NEXT:    vpslld $16, %ymm2, %ymm2
+; AVXVNNI-NEXT:    vpmovzxwd {{.*#+}} ymm7 = xmm3[0],zero,xmm3[1],zero,xmm3[2],zero,xmm3[3],zero,xmm3[4],zero,xmm3[5],zero,xmm3[6],zero,xmm3[7],zero
+; AVXVNNI-NEXT:    vpslld $16, %ymm7, %ymm7
+; AVXVNNI-NEXT:    vextracti128 $1, %ymm3, %xmm3
+; AVXVNNI-NEXT:    vpmovzxwd {{.*#+}} ymm3 = xmm3[0],zero,xmm3[1],zero,xmm3[2],zero,xmm3[3],zero,xmm3[4],zero,xmm3[5],zero,xmm3[6],zero,xmm3[7],zero
+; AVXVNNI-NEXT:    vpslld $16, %ymm3, %ymm3
+; AVXVNNI-NEXT:    vpmovzxwd {{.*#+}} ymm8 = xmm4[0],zero,xmm4[1],zero,xmm4[2],zero,xmm4[3],zero,xmm4[4],zero,xmm4[5],zero,xmm4[6],zero,xmm4[7],zero
+; AVXVNNI-NEXT:    vpslld $16, %ymm8, %ymm8
+; AVXVNNI-NEXT:    vmulps %ymm6, %ymm8, %ymm6
+; AVXVNNI-NEXT:    vextracti128 $1, %ymm4, %xmm4
+; AVXVNNI-NEXT:    vpmovzxwd {{.*#+}} ymm4 = xmm4[0],zero,xmm4[1],zero,xmm4[2],zero,xmm4[3],zero,xmm4[4],zero,xmm4[5],zero,xmm4[6],zero,xmm4[7],zero
+; AVXVNNI-NEXT:    vpslld $16, %ymm4, %ymm4
+; AVXVNNI-NEXT:    vmulps %ymm4, %ymm2, %ymm2
+; AVXVNNI-NEXT:    vpmovzxwd {{.*#+}} ymm4 = xmm5[0],zero,xmm5[1],zero,xmm5[2],zero,xmm5[3],zero,xmm5[4],zero,xmm5[5],zero,xmm5[6],zero,xmm5[7],zero
+; AVXVNNI-NEXT:    vpslld $16, %ymm4, %ymm4
+; AVXVNNI-NEXT:    vmulps %ymm4, %ymm7, %ymm4
+; AVXVNNI-NEXT:    vextracti128 $1, %ymm5, %xmm5
+; AVXVNNI-NEXT:    vpmovzxwd {{.*#+}} ymm5 = xmm5[0],zero,xmm5[1],zero,xmm5[2],zero,xmm5[3],zero,xmm5[4],zero,xmm5[5],zero,xmm5[6],zero,xmm5[7],zero
+; AVXVNNI-NEXT:    vpslld $16, %ymm5, %ymm5
+; AVXVNNI-NEXT:    vmulps %ymm5, %ymm3, %ymm3
+; AVXVNNI-NEXT:    vaddps %ymm6, %ymm0, %ymm0
+; AVXVNNI-NEXT:    vaddps %ymm2, %ymm0, %ymm0
+; AVXVNNI-NEXT:    vaddps %ymm4, %ymm1, %ymm1
+; AVXVNNI-NEXT:    vaddps %ymm3, %ymm1, %ymm1
+; AVXVNNI-NEXT:    retq
+;
+; AVXVNNIINT8INT16-LABEL: partial_reduce_fmla_bf16_v16f32:
+; AVXVNNIINT8INT16:       # %bb.0:
+; AVXVNNIINT8INT16-NEXT:    vdpbf16ps %zmm2, %zmm1, %zmm0
+; AVXVNNIINT8INT16-NEXT:    retq
+;
+; AVX10-LABEL: partial_reduce_fmla_bf16_v16f32:
+; AVX10:       # %bb.0:
+; AVX10-NEXT:    vdpbf16ps %zmm2, %zmm1, %zmm0
+; AVX10-NEXT:    retq
+  %a.ext = fpext <32 x bfloat> %a to <32 x float>
+  %b.ext = fpext <32 x bfloat> %b to <32 x float>
+  %mul = fmul <32 x float> %a.ext, %b.ext
+  %res = call <16 x float> @llvm.vector.partial.reduce.fadd.v16f32.v32f32(<16 x float> %acc, <32 x float> %mul)
+  ret <16 x float> %res
+}
+
+; VNNI-INT8 i8 x i8 -> i32 SMLA (sext x sext) tests: VPDPBSSD
+
+; Test 128-bit (xmm) SMLA i8.
+define <4 x i32> @partial_reduce_smla_i8_v4i32(<4 x i32> %acc, <16 x i8> %a, <16 x i8> %b) {
+; AVX512VNNI-LABEL: partial_reduce_smla_i8_v4i32:
+; AVX512VNNI:       # %bb.0:
+; AVX512VNNI-NEXT:    vpmovsxbd %xmm1, %zmm1
+; AVX512VNNI-NEXT:    vpmovsxbd %xmm2, %zmm2
+; AVX512VNNI-NEXT:    vpmulld %zmm2, %zmm1, %zmm1
+; AVX512VNNI-NEXT:    vpaddd %xmm1, %xmm0, %xmm0
+; AVX512VNNI-NEXT:    vextracti32x4 $3, %zmm1, %xmm2
+; AVX512VNNI-NEXT:    vpaddd %xmm0, %xmm2, %xmm0
+; AVX512VNNI-NEXT:    vextracti32x4 $2, %zmm1, %xmm2
+; AVX512VNNI-NEXT:    vextracti128 $1, %ymm1, %xmm1
+; AVX512VNNI-NEXT:    vpaddd %xmm2, %xmm1, %xmm1
+; AVX512VNNI-NEXT:    vpaddd %xmm0, %xmm1, %xmm0
+; AVX512VNNI-NEXT:    vzeroupper
+; AVX512VNNI-NEXT:    retq
+;
+; AVXVNNI-LABEL: partial_reduce_smla_i8_v4i32:
+; AVXVNNI:       # %bb.0:
+; AVXVNNI-NEXT:    vpmovsxbd %xmm1, %ymm3
+; AVXVNNI-NEXT:    vpshufd {{.*#+}} xmm1 = xmm1[2,3,2,3]
+; AVXVNNI-NEXT:    vpmovsxbd %xmm1, %ymm1
+; AVXVNNI-NEXT:    vpmovsxbd %xmm2, %ymm4
+; AVXVNNI-NEXT:    vpmulld %ymm4, %ymm3, %ymm3
+; AVXVNNI-NEXT:    vpshufd {{.*#+}} xmm2 = xmm2[2,3,2,3]
+; AVXVNNI-NEXT:    vpmovsxbd %xmm2, %ymm2
+; AVXVNNI-NEXT:    vpmulld %ymm2, %ymm1, %ymm1
+; AVXVNNI-NEXT:    vpaddd %xmm3, %xmm0, %xmm0
+; AVXVNNI-NEXT:    vextracti128 $1, %ymm3, %xmm2
+; AVXVNNI-NEXT:    vpaddd %xmm0, %xmm2, %xmm0
+; AVXVNNI-NEXT:    vpaddd %xmm1, %xmm0, %xmm0
+; AVXVNNI-NEXT:    vextracti128 $1, %ymm1, %xmm1
+; AVXVNNI-NEXT:    vpaddd %xmm0, %xmm1, %xmm0
+; AVXVNNI-NEXT:    vzeroupper
+; AVXVNNI-NEXT:    retq
+;
+; AVXVNNIINT8INT16-LABEL: partial_reduce_smla_i8_v4i32:
+; AVXVNNIINT8INT16:       # %bb.0:
+; AVXVNNIINT8INT16-NEXT:    vpdpbssd %xmm2, %xmm1, %xmm0
+; AVXVNNIINT8INT16-NEXT:    retq
+;
+; AVX10-LABEL: partial_reduce_smla_i8_v4i32:
+; AVX10:       # %bb.0:
+; AVX10-NEXT:    vpdpbssd %xmm2, %xmm1, %xmm0
+; AVX10-NEXT:    retq
+  %a.sext = sext <16 x i8> %a to <16 x i32>
+  %b.sext = sext <16 x i8> %b to <16 x i32>
+  %mul = mul nsw <16 x i32> %a.sext, %b.sext
+  %res = call <4 x i32> @llvm.vector.partial.reduce.add.v4i32.v16i32(<4 x i32> %acc, <16 x i32> %mul)
+  ret <4 x i32> %res
+}
+
+; Test 256-bit (ymm) SMLA i8.
+define <8 x i32> @partial_reduce_smla_i8_v8i32(<8 x i32> %acc, <32 x i8> %a, <32 x i8> %b) {
+; AVX512VNNI-LABEL: partial_reduce_smla_i8_v8i32:
+; AVX512VNNI:       # %bb.0:
+; AVX512VNNI-NEXT:    vpmovsxbd %xmm1, %zmm3
+; AVX512VNNI-NEXT:    vextracti128 $1, %ymm1, %xmm1
+; AVX512VNNI-NEXT:    vpmovsxbd %xmm1, %zmm1
+; AVX512VNNI-NEXT:    vpmovsxbd %xmm2, %zmm4
+; AVX512VNNI-NEXT:    vpmulld %zmm4, %zmm3, %zmm3
+; AVX512VNNI-NEXT:    vextracti128 $1, %ymm2, %xmm2
+; AVX512VNNI-NEXT:    vpmovsxbd %xmm2, %zmm2
+; AVX512VNNI-NEXT:    vpmulld %zmm2, %zmm1, %zmm1
+; AVX512VNNI-NEXT:    vpaddd %ymm3, %ymm0, %ymm0
+; AVX512VNNI-NEXT:    vextracti64x4 $1, %zmm3, %ymm2
+; AVX512VNNI-NEXT:    vpaddd %ymm0, %ymm2, %ymm0
+; AVX512VNNI-NEXT:    vpaddd %ymm1, %ymm0, %ymm0
+; AVX512VNNI-NEXT:    vextracti64x4 $1, %zmm1, %ymm1
+; AVX512VNNI-NEXT:    vpaddd %ymm0, %ymm1, %ymm0
+; AVX512VNNI-NEXT:    retq
+;
+; AVXVNNI-LABEL: partial_reduce_smla_i8_v8i32:
+; AVXVNNI:       # %bb.0:
+; AVXVNNI-NEXT:    vpmovsxbd %xmm1, %ymm3
+; AVXVNNI-NEXT:    vpshufd {{.*#+}} xmm4 = xmm1[2,3,2,3]
+; AVXVNNI-NEXT:    vpmovsxbd %xmm4, %ymm4
+; AVXVNNI-NEXT:    vextracti128 $1, %ymm1, %xmm1
+; AVXVNNI-NEXT:    vpmovsxbd %xmm1, %ymm5
+; AVXVNNI-NEXT:    vpshufd {{.*#+}} xmm1 = xmm1[2,3,2,3]
+; AVXVNNI-NEXT:    vpmovsxbd %xmm1, %ymm1
+; AVXVNNI-NEXT:    vpmovsxbd %xmm2, %ymm6
+; AVXVNNI-NEXT:    vpmulld %ymm6, %ymm3, %ymm3
+; AVXVNNI-NEXT:    vpshufd {{.*#+}} xmm6 = xmm2[2,3,2,3]
+; AVXVNNI-NEXT:    vpmovsxbd %xmm6, %ymm6
+; AVXVNNI-NEXT:    vpmulld %ymm6, %ymm4, %ymm4
+; AVXVNNI-NEXT:    vextracti128 $1, %ymm2, %xmm2
+; AVXVNNI-NEXT:    vpmovsxbd %xmm2, %ymm6
+; AVXVNNI-NEXT:    vpmulld %ymm6, %ymm5, %ymm5
+; AVXVNNI-NEXT:    vpshufd {{.*#+}} xmm2 = xmm2[2,3,2,3]
+; AVXVNNI-NEXT:    vpmovsxbd %xmm2, %ymm2
+; AVXVNNI-NEXT:    vpmulld %ymm2, %ymm1, %ymm1
+; AVXVNNI-NEXT:    vpaddd %ymm3, %ymm0, %ymm0
+; AVXVNNI-NEXT:    vpaddd %ymm4, %ymm0, %ymm0
+; AVXVNNI-NEXT:    vpaddd %ymm5, %ymm0, %ymm0
+; AVXVNNI-NEXT:    vpaddd %ymm1, %ymm0, %ymm0
+; AVXVNNI-NEXT:    retq
+;
+; AVXVNNIINT8INT16-LABEL: partial_reduce_smla_i8_v8i32:
+; AVXVNNIINT8INT16:       # %bb.0:
+; AVXVNNIINT8INT16-NEXT:    vpdpbssd %ymm2, %ymm1, %ymm0
+; AVXVNNIINT8INT16-NEXT:    retq
+;
+; AVX10-LABEL: partial_reduce_smla_i8_v8i32:
+; AVX10:       # %bb.0:
+; AVX10-NEXT:    vpdpbssd %ymm2, %ymm1, %ymm0
+; AVX10-NEXT:    retq
+  %a.sext = sext <32 x i8> %a to <32 x i32>
+  %b.sext = sext <32 x i8> %b to <32 x i32>
+  %mul = mul nsw <32 x i32> %a.sext, %b.sext
+  %res = call <8 x i32> @llvm.vector.partial.reduce.add.v8i32.v32i32(<8 x i32> %acc, <32 x i32> %mul)
+  ret <8 x i32> %res
+}
+
+; Test 512-bit (zmm) SMLA i8.
+; VEX (AVXVNNIINT8): splits to 2x256-bit ymm.
+; EVEX (AVX10.2): native 512-bit zmm.
+define <16 x i32> @partial_reduce_smla_i8_v16i32(<16 x i32> %acc, <64 x i8> %a, <64 x i8> %b) {
+; AVX512VNNI-LABEL: partial_reduce_smla_i8_v16i32:
+; AVX512VNNI:       # %bb.0:
+; AVX512VNNI-NEXT:    vpmovsxbd %xmm1, %zmm3
+; AVX512VNNI-NEXT:    vextracti128 $1, %ymm1, %xmm4
+; AVX512VNNI-NEXT:    vpmovsxbd %xmm4, %zmm4
+; AVX512VNNI-NEXT:    vextracti64x4 $1, %zmm1, %ymm1
+; AVX512VNNI-NEXT:    vpmovsxbd %xmm1, %zmm5
+; AVX512VNNI-NEXT:    vextracti128 $1, %ymm1, %xmm1
+; AVX512VNNI-NEXT:    vpmovsxbd %xmm1, %zmm1
+; AVX512VNNI-NEXT:    vpmovsxbd %xmm2, %zmm6
+; AVX512VNNI-NEXT:    vpmulld %zmm6, %zmm3, %zmm3
+; AVX512VNNI-NEXT:    vextracti128 $1, %ymm2, %xmm6
+; AVX512VNNI-NEXT:    vpmovsxbd %xmm6, %zmm6
+; AVX512VNNI-NEXT:    vpmulld %zmm6, %zmm4, %zmm4
+; AVX512VNNI-NEXT:    vextracti64x4 $1, %zmm2, %ymm2
+; AVX512VNNI-NEXT:    vpmovsxbd %xmm2, %zmm6
+; AVX512VNNI-NEXT:    vpmulld %zmm6, %zmm5, %zmm5
+; AVX512VNNI-NEXT:    vextracti128 $1, %ymm2, %xmm2
+; AVX512VNNI-NEXT:    vpmovsxbd %xmm2, %zmm2
+; AVX512VNNI-NEXT:    vpmulld %zmm2, %zmm1, %zmm1
+; AVX512VNNI-NEXT:    vpaddd %zmm3, %zmm0, %zmm0
+; AVX512VNNI-NEXT:    vpaddd %zmm4, %zmm0, %zmm0
+; AVX512VNNI-NEXT:    vpaddd %zmm5, %zmm0, %zmm0
+; AVX512VNNI-NEXT:    vpaddd %zmm1, %zmm0, %zmm0
+; AVX512VNNI-NEXT:    retq
+;
+; AVXVNNI-LABEL: partial_reduce_smla_i8_v16i32:
+; AVXVNNI:       # %bb.0:
+; AVXVNNI-NEXT:    vpmovsxbd %xmm2, %ymm6
+; AVXVNNI-NEXT:    vpshufd {{.*#+}} xmm7 = xmm2[2,3,2,3]
+; AVXVNNI-NEXT:    vpmovsxbd %xmm7, %ymm7
+; AVXVNNI-NEXT:    vextracti128 $1, %ymm2, %xmm2
+; AVXVNNI-NEXT:    vpmovsxbd %xmm2, %ymm8
+; AVXVNNI-NEXT:    vpshufd {{.*#+}} xmm2 = xmm2[2,3,2,3]
+; AVXVNNI-NEXT:    vpmovsxbd %xmm2, %ymm2
+; AVXVNNI-NEXT:    vpmovsxbd %xmm3, %ymm9
+; AVXVNNI-NEXT:    vpshufd {{.*#+}} xmm10 = xmm3[2,3,2,3]
+; AVXVNNI-NEXT:    vpmovsxbd %xmm10, %ymm10
+; AVXVNNI-NEXT:    vextracti128 $1, %ymm3, %xmm3
+; AVXVNNI-NEXT:    vpmovsxbd %xmm3, %ymm11
+; AVXVNNI-NEXT:    vpshufd {{.*#+}} xmm3 = xmm3[2,3,2,3]
+; AVXVNNI-NEXT:    vpmovsxbd %xmm3, %ymm3
+; AVXVNNI-NEXT:    vpmovsxbd %xmm4, %ymm12
+; AVXVNNI-NEXT:    vpmulld %ymm12, %ymm6, %ymm6
+; AVXVNNI-NEXT:    vpshufd {{.*#+}} xmm12 = xmm4[2,3,2,3]
+; AVXVNNI-NEXT:    vpmovsxbd %xmm12, %ymm12
+; AVXVNNI-NEXT:    vpmulld %ymm12, %ymm7, %ymm7
+; AVXVNNI-NEXT:    vextracti128 $1, %ymm4, %xmm4
+; AVXVNNI-NEXT:    vpmovsxbd %xmm4, %ymm12
+; AVXVNNI-NEXT:    vpmulld %ymm12, %ymm8, %ymm8
+; AVXVNNI-NEXT:    vpshufd {{.*#+}} xmm4 = xmm4[2,3,2,3]
+; AVXVNNI-NEXT:    vpmovsxbd %xmm4, %ymm4
+; AVXVNNI-NEXT:    vpmulld %ymm4, %ymm2, %ymm2
+; AVXVNNI-NEXT:    vpmovsxbd %xmm5, %ymm4
+; AVXVNNI-NEXT:    vpmulld %ymm4, %ymm9, %ymm4
+; AVXVNNI-NEXT:    vpshufd {{.*#+}} xmm9 = xmm5[2,3,2,3]
+; AVXVNNI-NEXT:    vpmovsxbd %xmm9, %ymm9
+; AVXVNNI-NEXT:    vpmulld %ymm9, %ymm10, %ymm9
+; AVXVNNI-NEXT:    vextracti128 $1, %ymm5, %xmm5
+; AVXVNNI-NEXT:    vpmovsxbd %xmm5, %ymm10
+; AVXVNNI-NEXT:    vpmulld %ymm10, %ymm11, %ymm10
+; AVXVNNI-NEXT:    vpshufd {{.*#+}} xmm5 = xmm5[2,3,2,3]
+; AVXVNNI-NEXT:    vpmovsxbd %xmm5, %ymm5
+; AVXVNNI-NEXT:    vpmulld %ymm5, %ymm3, %ymm3
+; AVXVNNI-NEXT:    vpaddd %ymm6, %ymm0, %ymm0
+; AVXVNNI-NEXT:    vpaddd %ymm7, %ymm0, %ymm0
+; AVXVNNI-NEXT:    vpaddd %ymm0, %ymm8, %ymm0
+; AVXVNNI-NEXT:    vpaddd %ymm2, %ymm0, %ymm0
+; AVXVNNI-NEXT:    vpaddd %ymm4, %ymm1, %ymm1
+; AVXVNNI-NEXT:    vpaddd %ymm1, %ymm9, %ymm1
+; AVXVNNI-NEXT:    vpaddd %ymm1, %ymm10, %ymm1
+; AVXVNNI-NEXT:    vpaddd %ymm3, %ymm1, %ymm1
+; AVXVNNI-NEXT:    retq
+;
+; AVXVNNIINT8INT16-LABEL: partial_reduce_smla_i8_v16i32:
+; AVXVNNIINT8INT16:       # %bb.0:
+; AVXVNNIINT8INT16-NEXT:    vextractf64x4 $1, %zmm2, %ymm3
+; AVXVNNIINT8INT16-NEXT:    vextractf64x4 $1, %zmm1, %ymm4
+; AVXVNNIINT8INT16-NEXT:    vextractf64x4 $1, %zmm0, %ymm5
+; AVXVNNIINT8INT16-NEXT:    vpdpbssd %ymm3, %ymm4, %ymm5
+; AVXVNNIINT8INT16-NEXT:    vpdpbssd %ymm2, %ymm1, %ymm0
+; AVXVNNIINT8INT16-NEXT:    vinsertf64x4 $1, %ymm5, %zmm0, %zmm0
+; AVXVNNIINT8INT16-NEXT:    retq
+;
+; AVX10-LABEL: partial_reduce_smla_i8_v16i32:
+; AVX10:       # %bb.0:
+; AVX10-NEXT:    vpdpbssd %zmm2, %zmm1, %zmm0
+; AVX10-NEXT:    retq
+  %a.sext = sext <64 x i8> %a to <64 x i32>
+  %b.sext = sext <64 x i8> %b to <64 x i32>
+  %mul = mul nsw <64 x i32> %a.sext, %b.sext
+  %res = call <16 x i32> @llvm.vector.partial.reduce.add.v16i32.v64i32(<16 x i32> %acc, <64 x i32> %mul)
+  ret <16 x i32> %res
+}
+
+; VNNI-INT8 i8 x i8 -> i32 UMLA (zext x zext) tests: VPDPBUUD
+
+; Test 128-bit (xmm) UMLA i8.
+define <4 x i32> @partial_reduce_umla_i8_v4i32(<4 x i32> %acc, <16 x i8> %a, <16 x i8> %b) {
+; AVX512VNNI-LABEL: partial_reduce_umla_i8_v4i32:
+; AVX512VNNI:       # %bb.0:
+; AVX512VNNI-NEXT:    vpmovzxbd {{.*#+}} zmm1 = xmm1[0],zero,zero,zero,xmm1[1],zero,zero,zero,xmm1[2],zero,zero,zero,xmm1[3],zero,zero,zero,xmm1[4],zero,zero,zero,xmm1[5],zero,zero,zero,xmm1[6],zero,zero,zero,xmm1[7],zero,zero,zero,xmm1[8],zero,zero,zero,xmm1[9],zero,zero,zero,xmm1[10],zero,zero,zero,xmm1[11],zero,zero,zero,xmm1[12],zero,zero,zero,xmm1[13],zero,zero,zero,xmm1[14],zero,zero,zero,xmm1[15],zero,zero,zero
+; AVX512VNNI-NEXT:    vpmovzxbd {{.*#+}} zmm2 = xmm2[0],zero,zero,zero,xmm2[1],zero,zero,zero,xmm2[2],zero,zero,zero,xmm2[3],zero,zero,zero,xmm2[4],zero,zero,zero,xmm2[5],zero,zero,zero,xmm2[6],zero,zero,zero,xmm2[7],zero,zero,zero,xmm2[8],zero,zero,zero,xmm2[9],zero,zero,zero,xmm2[10],zero,zero,zero,xmm2[11],zero,zero,zero,xmm2[12],zero,zero,zero,xmm2[13],zero,zero,zero,xmm2[14],zero,zero,zero,xmm2[15],zero,zero,zero
+; AVX512VNNI-NEXT:    vpmaddwd %zmm2, %zmm1, %zmm1
+; AVX512VNNI-NEXT:    vpaddd %xmm1, %xmm0, %xmm0
+; AVX512VNNI-NEXT:    vextracti32x4 $3, %zmm1, %xmm2
+; AVX512VNNI-NEXT:    vpaddd %xmm0, %xmm2, %xmm0
+; AVX512VNNI-NEXT:    vextracti32x4 $2, %zmm1, %xmm2
+; AVX512VNNI-NEXT:    vextracti128 $1, %ymm1, %xmm1
+; AVX512VNNI-NEXT:    vpaddd %xmm2, %xmm1, %xmm1
+; AVX512VNNI-NEXT:    vpaddd %xmm0, %xmm1, %xmm0
+; AVX512VNNI-NEXT:    vzeroupper
+; AVX512VNNI-NEXT:    retq
+;
+; AVXVNNI-LABEL: partial_reduce_umla_i8_v4i32:
+; AVXVNNI:       # %bb.0:
+; AVXVNNI-NEXT:    vpmovzxbd {{.*#+}} ymm3 = xmm1[0],zero,zero,zero,xmm1[1],zero,zero,zero,xmm1[2],zero,zero,zero,xmm1[3],zero,zero,zero,xmm1[4],zero,zero,zero,xmm1[5],zero,zero,zero,xmm1[6],zero,zero,zero,xmm1[7],zero,zero,zero
+; AVXVNNI-NEXT:    vpshufd {{.*#+}} xmm1 = xmm1[2,3,2,3]
+; AVXVNNI-NEXT:    vpmovzxbd {{.*#+}} ymm1 = xmm1[0],zero,zero,zero,xmm1[1],zero,zero,zero,xmm1[2],zero,zero,zero,xmm1[3],zero,zero,zero,xmm1[4],zero,zero,zero,xmm1[5],zero,zero,zero,xmm1[6],zero,zero,zero,xmm1[7],zero,zero,zero
+; AVXVNNI-NEXT:    vpmovzxbd {{.*#+}} ymm4 = xmm2[0],zero,zero,zero,xmm2[1],zero,zero,zero,xmm2[2],zero,zero,zero,xmm2[3],zero,zero,zero,xmm2[4],zero,zero,zero,xmm2[5],zero,zero,zero,xmm2[6],zero,zero,zero,xmm2[7],zero,zero,zero
+; AVXVNNI-NEXT:    vpmaddwd %ymm4, %ymm3, %ymm3
+; AVXVNNI-NEXT:    vpshufd {{.*#+}} xmm2 = xmm2[2,3,2,3]
+; AVXVNNI-NEXT:    vpmovzxbd {{.*#+}} ymm2 = xmm2[0],zero,zero,zero,xmm2[1],zero,zero,zero,xmm2[2],zero,zero,zero,xmm2[3],zero,zero,zero,xmm2[4],zero,zero,zero,xmm2[5],zero,zero,zero,xmm2[6],zero,zero,zero,xmm2[7],zero,zero,zero
+; AVXVNNI-NEXT:    vpmaddwd %ymm2, %ymm1, %ymm1
+; AVXVNNI-NEXT:    vpaddd %xmm3, %xmm0, %xmm0
+; AVXVNNI-NEXT:    vextracti128 $1, %ymm3, %xmm2
+; AVXVNNI-NEXT:    vpaddd %xmm0, %xmm2, %xmm0
+; AVXVNNI-NEXT:    vpaddd %xmm1, %xmm0, %xmm0
+; AVXVNNI-NEXT:    vextracti128 $1, %ymm1, %xmm1
+; AVXVNNI-NEXT:    vpaddd %xmm0, %xmm1, %xmm0
+; AVXVNNI-NEXT:    vzeroupper
+; AVXVNNI-NEXT:    retq
+;
+; AVXVNNIINT8INT16-LABEL: partial_reduce_umla_i8_v4i32:
+; AVXVNNIINT8INT16:       # %bb.0:
+; AVXVNNIINT8INT16-NEXT:    vpdpbuud %xmm2, %xmm1, %xmm0
+; AVXVNNIINT8INT16-NEXT:    retq
+;
+; AVX10-LABEL: partial_reduce_umla_i8_v4i32:
+; AVX10:       # %bb.0:
+; AVX10-NEXT:    vpdpbuud %xmm2, %xmm1, %xmm0
+; AVX10-NEXT:    retq
+  %a.zext = zext <16 x i8> %a to <16 x i32>
+  %b.zext = zext <16 x i8> %b to <16 x i32>
+  %mul = mul nsw <16 x i32> %a.zext, %b.zext
+  %res = call <4 x i32> @llvm.vector.partial.reduce.add.v4i32.v16i32(<4 x i32> %acc, <16 x i32> %mul)
+  ret <4 x i32> %res
+}
+
+; Test 256-bit (ymm) UMLA i8.
+define <8 x i32> @partial_reduce_umla_i8_v8i32(<8 x i32> %acc, <32 x i8> %a, <32 x i8> %b) {
+; AVX512VNNI-LABEL: partial_reduce_umla_i8_v8i32:
+; AVX512VNNI:       # %bb.0:
+; AVX512VNNI-NEXT:    vpmovzxbd {{.*#+}} zmm3 = xmm1[0],zero,zero,zero,xmm1[1],zero,zero,zero,xmm1[2],zero,zero,zero,xmm1[3],zero,zero,zero,xmm1[4],zero,zero,zero,xmm1[5],zero,zero,zero,xmm1[6],zero,zero,zero,xmm1[7],zero,zero,zero,xmm1[8],zero,zero,zero,xmm1[9],zero,zero,zero,xmm1[10],zero,zero,zero,xmm1[11],zero,zero,zero,xmm1[12],zero,zero,zero,xmm1[13],zero,zero,zero,xmm1[14],zero,zero,zero,xmm1[15],zero,zero,zero
+; AVX512VNNI-NEXT:    vextracti128 $1, %ymm1, %xmm1
+; AVX512VNNI-NEXT:    vpmovzxbd {{.*#+}} zmm1 = xmm1[0],zero,zero,zero,xmm1[1],zero,zero,zero,xmm1[2],zero,zero,zero,xmm1[3],zero,zero,zero,xmm1[4],zero,zero,zero,xmm1[5],zero,zero,zero,xmm1[6],zero,zero,zero,xmm1[7],zero,zero,zero,xmm1[8],zero,zero,zero,xmm1[9],zero,zero,zero,xmm1[10],zero,zero,zero,xmm1[11],zero,zero,zero,xmm1[12],zero,zero,zero,xmm1[13],zero,zero,zero,xmm1[14],zero,zero,zero,xmm1[15],zero,zero,zero
+; AVX512VNNI-NEXT:    vpmovzxbd {{.*#+}} zmm4 = xmm2[0],zero,zero,zero,xmm2[1],zero,zero,zero,xmm2[2],zero,zero,zero,xmm2[3],zero,zero,zero,xmm2[4],zero,zero,zero,xmm2[5],zero,zero,zero,xmm2[6],zero,zero,zero,xmm2[7],zero,zero,zero,xmm2[8],zero,zero,zero,xmm2[9],zero,zero,zero,xmm2[10],zero,zero,zero,xmm2[11],zero,zero,zero,xmm2[12],zero,zero,zero,xmm2[13],zero,zero,zero,xmm2[14],zero,zero,zero,xmm2[15],zero,zero,zero
+; AVX512VNNI-NEXT:    vpmaddwd %zmm4, %zmm3, %zmm3
+; AVX512VNNI-NEXT:    vextracti128 $1, %ymm2, %xmm2
+; AVX512VNNI-NEXT:    vpmovzxbd {{.*#+}} zmm2 = xmm2[0],zero,zero,zero,xmm2[1],zero,zero,zero,xmm2[2],zero,zero,zero,xmm2[3],zero,zero,zero,xmm2[4],zero,zero,zero,xmm2[5],zero,zero,zero,xmm2[6],zero,zero,zero,xmm2[7],zero,zero,zero,xmm2[8],zero,zero,zero,xmm2[9],zero,zero,zero,xmm2[10],zero,zero,zero,xmm2[11],zero,zero,zero,xmm2[12],zero,zero,zero,xmm2[13],zero,zero,zero,xmm2[14],zero,zero,zero,xmm2[15],zero,zero,zero
+; AVX512VNNI-NEXT:    vpmaddwd %zmm2, %zmm1, %zmm1
+; AVX512VNNI-NEXT:    vpaddd %ymm3, %ymm0, %ymm0
+; AVX512VNNI-NEXT:    vextracti64x4 $1, %zmm3, %ymm2
+; AVX512VNNI-NEXT:    vpaddd %ymm0, %ymm2, %ymm0
+; AVX512VNNI-NEXT:    vpaddd %ymm1, %ymm0, %ymm0
+; AVX512VNNI-NEXT:    vextracti64x4 $1, %zmm1, %ymm1
+; AVX512VNNI-NEXT:    vpaddd %ymm0, %ymm1, %ymm0
+; AVX512VNNI-NEXT:    retq
+;
+; AVXVNNI-LABEL: partial_reduce_umla_i8_v8i32:
+; AVXVNNI:       # %bb.0:
+; AVXVNNI-NEXT:    vextracti128 $1, %ymm1, %xmm3
+; AVXVNNI-NEXT:    vpshufd {{.*#+}} xmm4 = xmm3[2,3,2,3]
+; AVXVNNI-NEXT:    vpmovzxbd {{.*#+}} ymm4 = xmm4[0],zero,zero,zero,xmm4[1],zero,zero,zero,xmm4[2],zero,zero,zero,xmm4[3],zero,zero,zero,xmm4[4],zero,zero,zero,xmm4[5],zero,zero,zero,xmm4[6],zero,zero,zero,xmm4[7],zero,zero,zero
+; AVXVNNI-NEXT:    vpmovzxbd {{.*#+}} ymm3 = xmm3[0],zero,zero,zero,xmm3[1],zero,zero,zero,xmm3[2],zero,zero,zero,xmm3[3],zero,zero,zero,xmm3[4],zero,zero,zero,xmm3[5],zero,zero,zero,xmm3[6],zero,zero,zero,xmm3[7],zero,zero,zero
+; AVXVNNI-NEXT:    vpshufd {{.*#+}} xmm5 = xmm1[2,3,2,3]
+; AVXVNNI-NEXT:    vpmovzxbd {{.*#+}} ymm5 = xmm5[0],zero,zero,zero,xmm5[1],zero,zero,zero,xmm5[2],zero,zero,zero,xmm5[3],zero,zero,zero,xmm5[4],zero,zero,zero,xmm5[5],zero,zero,zero,xmm5[6],zero,zero,zero,xmm5[7],zero,zero,zero
+; AVXVNNI-NEXT:    vpmovzxbd {{.*#+}} ymm1 = xmm1[0],zero,zero,zero,xmm1[1],zero,zero,zero,xmm1[2],zero,zero,zero,xmm1[3],zero,zero,zero,xmm1[4],zero,zero,zero,xmm1[5],zero,zero,zero,xmm1[6],zero,zero,zero,xmm1[7],zero,zero,zero
+; AVXVNNI-NEXT:    vextracti128 $1, %ymm2, %xmm6
+; AVXVNNI-NEXT:    vpshufd {{.*#+}} xmm7 = xmm6[2,3,2,3]
+; AVXVNNI-NEXT:    vpmovzxbd {{.*#+}} ymm7 = xmm7[0],zero,zero,zero,xmm7[1],zero,zero,zero,xmm7[2],zero,zero,zero,xmm7[3],zero,zero,zero,xmm7[4],zero,zero,zero,xmm7[5],zero,zero,zero,xmm7[6],zero,zero,zero,xmm7[7],zero,zero,zero
+; AVXVNNI-NEXT:    vpmaddwd %ymm7, %ymm4, %ymm4
+; AVXVNNI-NEXT:    vpmovzxbd {{.*#+}} ymm6 = xmm6[0],zero,zero,zero,xmm6[1],zero,zero,zero,xmm6[2],zero,zero,zero,xmm6[3],zero,zero,zero,xmm6[4],zero,zero,zero,xmm6[5],zero,zero,zero,xmm6[6],zero,zero,zero,xmm6[7],zero,zero,zero
+; AVXVNNI-NEXT:    vpmaddwd %ymm6, %ymm3, %ymm3
+; AVXVNNI-NEXT:    vpshufd {{.*#+}} xmm6 = xmm2[2,3,2,3]
+; AVXVNNI-NEXT:    vpmovzxbd {{.*#+}} ymm6 = xmm6[0],zero,zero,zero,xmm6[1],zero,zero,zero,xmm6[2],zero,zero,zero,xmm6[3],zero,zero,zero,xmm6[4],zero,zero,zero,xmm6[5],zero,zero,zero,xmm6[6],zero,zero,zero,xmm6[7],zero,zero,zero
+; AVXVNNI-NEXT:    vpmovzxbd {{.*#+}} ymm2 = xmm2[0],zero,zero,zero,xmm2[1],zero,zero,zero,xmm2[2],zero,zero,zero,xmm2[3],zero,zero,zero,xmm2[4],zero,zero,zero,xmm2[5],zero,zero,zero,xmm2[6],zero,zero,zero,xmm2[7],zero,zero,zero
+; AVXVNNI-NEXT:    {vex} vpdpwssd %ymm2, %ymm1, %ymm0
+; AVXVNNI-NEXT:    {vex} vpdpwssd %ymm6, %ymm5, %ymm0
+; AVXVNNI-NEXT:    vpaddd %ymm3, %ymm0, %ymm0
+; AVXVNNI-NEXT:    vpaddd %ymm4, %ymm0, %ymm0
+; AVXVNNI-NEXT:    retq
+;
+; AVXVNNIINT8INT16-LABEL: partial_reduce_umla_i8_v8i32:
+; AVXVNNIINT8INT16:       # %bb.0:
+; AVXVNNIINT8INT16-NEXT:    vpdpbuud %ymm2, %ymm1, %ymm0
+; AVXVNNIINT8INT16-NEXT:    retq
+;
+; AVX10-LABEL: partial_reduce_umla_i8_v8i32:
+; AVX10:       # %bb.0:
+; AVX10-NEXT:    vpdpbuud %ymm2, %ymm1, %ymm0
+; AVX10-NEXT:    retq
+  %a.zext = zext <32 x i8> %a to <32 x i32>
+  %b.zext = zext <32 x i8> %b to <32 x i32>
+  %mul = mul nsw <32 x i32> %a.zext, %b.zext
+  %res = call <8 x i32> @llvm.vector.partial.reduce.add.v8i32.v32i32(<8 x i32> %acc, <32 x i32> %mul)
+  ret <8 x i32> %res
+}
+
+; Test 512-bit (zmm) UMLA i8.
+define <16 x i32> @partial_reduce_umla_i8_v16i32(<16 x i32> %acc, <64 x i8> %a, <64 x i8> %b) {
+; AVX512VNNI-LABEL: partial_reduce_umla_i8_v16i32:
+; AVX512VNNI:       # %bb.0:
+; AVX512VNNI-NEXT:    vextracti64x4 $1, %zmm1, %ymm3
+; AVX512VNNI-NEXT:    vextracti128 $1, %ymm3, %xmm4
+; AVX512VNNI-NEXT:    vpmovzxbd {{.*#+}} zmm4 = xmm4[0],zero,zero,zero,xmm4[1],zero,zero,zero,xmm4[2],zero,zero,zero,xmm4[3],zero,zero,zero,xmm4[4],zero,zero,zero,xmm4[5],zero,zero,zero,xmm4[6],zero,zero,zero,xmm4[7],zero,zero,zero,xmm4[8],zero,zero,zero,xmm4[9],zero,zero,zero,xmm4[10],zero,zero,zero,xmm4[11],zero,zero,zero,xmm4[12],zero,zero,zero,xmm4[13],zero,zero,zero,xmm4[14],zero,zero,zero,xmm4[15],zero,zero,zero
+; AVX512VNNI-NEXT:    vpmovzxbd {{.*#+}} zmm3 = xmm3[0],zero,zero,zero,xmm3[1],zero,zero,zero,xmm3[2],zero,zero,zero,xmm3[3],zero,zero,zero,xmm3[4],zero,zero,zero,xmm3[5],zero,zero,zero,xmm3[6],zero,zero,zero,xmm3[7],zero,zero,zero,xmm3[8],zero,zero,zero,xmm3[9],zero,zero,zero,xmm3[10],zero,zero,zero,xmm3[11],zero,zero,zero,xmm3[12],zero,zero,zero,xmm3[13],zero,zero,zero,xmm3[14],zero,zero,zero,xmm3[15],zero,zero,zero
+; AVX512VNNI-NEXT:    vextracti128 $1, %ymm1, %xmm5
+; AVX512VNNI-NEXT:    vpmovzxbd {{.*#+}} zmm5 = xmm5[0],zero,zero,zero,xmm5[1],zero,zero,zero,xmm5[2],zero,zero,zero,xmm5[3],zero,zero,zero,xmm5[4],zero,zero,zero,xmm5[5],zero,zero,zero,xmm5[6],zero,zero,zero,xmm5[7],zero,zero,zero,xmm5[8],zero,zero,zero,xmm5[9],zero,zero,zero,xmm5[10],zero,zero,zero,xmm5[11],zero,zero,zero,xmm5[12],zero,zero,zero,xmm5[13],zero,zero,zero,xmm5[14],zero,zero,zero,xmm5[15],zero,zero,zero
+; AVX512VNNI-NEXT:    vpmovzxbd {{.*#+}} zmm1 = xmm1[0],zero,zero,zero,xmm1[1],zero,zero,zero,xmm1[2],zero,zero,zero,xmm1[3],zero,zero,zero,xmm1[4],zero,zero,zero,xmm1[5],zero,zero,zero,xmm1[6],zero,zero,zero,xmm1[7],zero,zero,zero,xmm1[8],zero,zero,zero,xmm1[9],zero,zero,zero,xmm1[10],zero,zero,zero,xmm1[11],zero,zero,zero,xmm1[12],zero,zero,zero,xmm1[13],zero,zero,zero,xmm1[14],zero,zero,zero,xmm1[15],zero,zero,zero
+; AVX512VNNI-NEXT:    vextracti64x4 $1, %zmm2, %ymm6
+; AVX512VNNI-NEXT:    vextracti128 $1, %ymm6, %xmm7
+; AVX512VNNI-NEXT:    vpmovzxbd {{.*#+}} zmm7 = xmm7[0],zero,zero,zero,xmm7[1],zero,zero,zero,xmm7[2],zero,zero,zero,xmm7[3],zero,zero,zero,xmm7[4],zero,zero,zero,xmm7[5],zero,zero,zero,xmm7[6],zero,zero,zero,xmm7[7],zero,zero,zero,xmm7[8],zero,zero,zero,xmm7[9],zero,zero,zero,xmm7[10],zero,zero,zero,xmm7[11],zero,zero,zero,xmm7[12],zero,zero,zero,xmm7[13],zero,zero,zero,xmm7[14],zero,zero,zero,xmm7[15],zero,zero,zero
+; AVX512VNNI-NEXT:    vpmaddwd %zmm7, %zmm4, %zmm4
+; AVX512VNNI-NEXT:    vpmovzxbd {{.*#+}} zmm6 = xmm6[0],zero,zero,zero,xmm6[1],zero,zero,zero,xmm6[2],zero,zero,zero,xmm6[3],zero,zero,zero,xmm6[4],zero,zero,zero,xmm6[5],zero,zero,zero,xmm6[6],zero,zero,zero,xmm6[7],zero,zero,zero,xmm6[8],zero,zero,zero,xmm6[9],zero,zero,zero,xmm6[10],zero,zero,zero,xmm6[11],zero,zero,zero,xmm6[12],zero,zero,zero,xmm6[13],zero,zero,zero,xmm6[14],zero,zero,zero,xmm6[15],zero,zero,zero
+; AVX512VNNI-NEXT:    vpmaddwd %zmm6, %zmm3, %zmm3
+; AVX512VNNI-NEXT:    vextracti128 $1, %ymm2, %xmm6
+; AVX512VNNI-NEXT:    vpmovzxbd {{.*#+}} zmm6 = xmm6[0],zero,zero,zero,xmm6[1],zero,zero,zero,xmm6[2],zero,zero,zero,xmm6[3],zero,zero,zero,xmm6[4],zero,zero,zero,xmm6[5],zero,zero,zero,xmm6[6],zero,zero,zero,xmm6[7],zero,zero,zero,xmm6[8],zero,zero,zero,xmm6[9],zero,zero,zero,xmm6[10],zero,zero,zero,xmm6[11],zero,zero,zero,xmm6[12],zero,zero,zero,xmm6[13],zero,zero,zero,xmm6[14],zero,zero,zero,xmm6[15],zero,zero,zero
+; AVX512VNNI-NEXT:    vpmovzxbd {{.*#+}} zmm2 = xmm2[0],zero,zero,zero,xmm2[1],zero,zero,zero,xmm2[2],zero,zero,zero,xmm2[3],zero,zero,zero,xmm2[4],zero,zero,zero,xmm2[5],zero,zero,zero,xmm2[6],zero,zero,zero,xmm2[7],zero,zero,zero,xmm2[8],zero,zero,zero,xmm2[9],zero,zero,zero,xmm2[10],zero,zero,zero,xmm2[11],zero,zero,zero,xmm2[12],zero,zero,zero,xmm2[13],zero,zero,zero,xmm2[14],zero,zero,zero,xmm2[15],zero,zero,zero
+; AVX512VNNI-NEXT:    vpdpwssd %zmm2, %zmm1, %zmm0
+; AVX512VNNI-NEXT:    vpdpwssd %zmm6, %zmm5, %zmm0
+; AVX512VNNI-NEXT:    vpaddd %zmm3, %zmm0, %zmm0
+; AVX512VNNI-NEXT:    vpaddd %zmm4, %zmm0, %zmm0
+; AVX512VNNI-NEXT:    retq
+;
+; AVXVNNI-LABEL: partial_reduce_umla_i8_v16i32:
+; AVXVNNI:       # %bb.0:
+; AVXVNNI-NEXT:    vextracti128 $1, %ymm3, %xmm6
+; AVXVNNI-NEXT:    vpshufd {{.*#+}} xmm7 = xmm6[2,3,2,3]
+; AVXVNNI-NEXT:    vpmovzxbd {{.*#+}} ymm7 = xmm7[0],zero,zero,zero,xmm7[1],zero,zero,zero,xmm7[2],zero,zero,zero,xmm7[3],zero,zero,zero,xmm7[4],zero,zero,zero,xmm7[5],zero,zero,zero,xmm7[6],zero,zero,zero,xmm7[7],zero,zero,zero
+; AVXVNNI-NEXT:    vpmovzxbd {{.*#+}} ymm6 = xmm6[0],zero,zero,zero,xmm6[1],zero,zero,zero,xmm6[2],zero,zero,zero,xmm6[3],zero,zero,zero,xmm6[4],zero,zero,zero,xmm6[5],zero,zero,zero,xmm6[6],zero,zero,zero,xmm6[7],zero,zero,zero
+; AVXVNNI-NEXT:    vpshufd {{.*#+}} xmm8 = xmm3[2,3,2,3]
+; AVXVNNI-NEXT:    vpmovzxbd {{.*#+}} ymm8 = xmm8[0],zero,zero,zero,xmm8[1],zero,zero,zero,xmm8[2],zero,zero,zero,xmm8[3],zero,zero,zero,xmm8[4],zero,zero,zero,xmm8[5],zero,zero,zero,xmm8[6],zero,zero,zero,xmm8[7],zero,zero,zero
+; AVXVNNI-NEXT:    vpmovzxbd {{.*#+}} ymm3 = xmm3[0],zero,zero,zero,xmm3[1],zero,zero,zero,xmm3[2],zero,zero,zero,xmm3[3],zero,zero,zero,xmm3[4],zero,zero,zero,xmm3[5],zero,zero,zero,xmm3[6],zero,zero,zero,xmm3[7],zero,zero,zero
+; AVXVNNI-NEXT:    vextracti128 $1, %ymm2, %xmm9
+; AVXVNNI-NEXT:    vpshufd {{.*#+}} xmm10 = xmm9[2,3,2,3]
+; AVXVNNI-NEXT:    vpmovzxbd {{.*#+}} ymm10 = xmm10[0],zero,zero,zero,xmm10[1],zero,zero,zero,xmm10[2],zero,zero,zero,xmm10[3],zero,zero,zero,xmm10[4],zero,zero,zero,xmm10[5],zero,zero,zero,xmm10[6],zero,zero,zero,xmm10[7],zero,zero,zero
+; AVXVNNI-NEXT:    vpmovzxbd {{.*#+}} ymm9 = xmm9[0],zero,zero,zero,xmm9[1],zero,zero,zero,xmm9[2],zero,zero,zero,xmm9[3],zero,zero,zero,xmm9[4],zero,zero,zero,xmm9[5],zero,zero,zero,xmm9[6],zero,zero,zero,xmm9[7],zero,zero,zero
+; AVXVNNI-NEXT:    vpshufd {{.*#+}} xmm11 = xmm2[2,3,2,3]
+; AVXVNNI-NEXT:    vpmovzxbd {{.*#+}} ymm11 = xmm11[0],zero,zero,zero,xmm11[1],zero,zero,zero,xmm11[2],zero,zero,zero,xmm11[3],zero,zero,zero,xmm11[4],zero,zero,zero,xmm11[5],zero,zero,zero,xmm11[6],zero,zero,zero,xmm11[7],zero,zero,zero
+; AVXVNNI-NEXT:    vpmovzxbd {{.*#+}} ymm2 = xmm2[0],zero,zero,zero,xmm2[1],zero,zero,zero,xmm2[2],zero,zero,zero,xmm2[3],zero,zero,zero,xmm2[4],zero,zero,zero,xmm2[5],zero,zero,zero,xmm2[6],zero,zero,zero,xmm2[7],zero,zero,zero
+; AVXVNNI-NEXT:    vextracti128 $1, %ymm5, %xmm12
+; AVXVNNI-NEXT:    vpshufd {{.*#+}} xmm13 = xmm12[2,3,2,3]
+; AVXVNNI-NEXT:    vpmovzxbd {{.*#+}} ymm13 = xmm13[0],zero,zero,zero,xmm13[1],zero,zero,zero,xmm13[2],zero,zero,zero,xmm13[3],zero,zero,zero,xmm13[4],zero,zero,zero,xmm13[5],zero,zero,zero,xmm13[6],zero,zero,zero,xmm13[7],zero,zero,zero
+; AVXVNNI-NEXT:    vpmaddwd %ymm7, %ymm13, %ymm7
+; AVXVNNI-NEXT:    vpmovzxbd {{.*#+}} ymm12 = xmm12[0],zero,zero,zero,xmm12[1],zero,zero,zero,xmm12[2],zero,zero,zero,xmm12[3],zero,zero,zero,xmm12[4],zero,zero,zero,xmm12[5],zero,zero,zero,xmm12[6],zero,zero,zero,xmm12[7],zero,zero,zero
+; AVXVNNI-NEXT:    vpmaddwd %ymm6, %ymm12, %ymm6
+; AVXVNNI-NEXT:    vpshufd {{.*#+}} xmm12 = xmm5[2,3,2,3]
+; AVXVNNI-NEXT:    vpmovzxbd {{.*#+}} ymm12 = xmm12[0],zero,zero,zero,xmm12[1],zero,zero,zero,xmm12[2],zero,zero,zero,xmm12[3],zero,zero,zero,xmm12[4],zero,zero,zero,xmm12[5],zero,zero,zero,xmm12[6],zero,zero,zero,xmm12[7],zero,zero,zero
+; AVXVNNI-NEXT:    vpmovzxbd {{.*#+}} ymm5 = xmm5[0],zero,zero,zero,xmm5[1],zero,zero,zero,xmm5[2],zero,zero,zero,xmm5[3],zero,zero,zero,xmm5[4],zero,zero,zero,xmm5[5],zero,zero,zero,xmm5[6],zero,zero,zero,xmm5[7],zero,zero,zero
+; AVXVNNI-NEXT:    {vex} vpdpwssd %ymm5, %ymm3, %ymm1
+; AVXVNNI-NEXT:    {vex} vpdpwssd %ymm12, %ymm8, %ymm1
+; AVXVNNI-NEXT:    vpaddd %ymm6, %ymm1, %ymm1
+; AVXVNNI-NEXT:    vpaddd %ymm7, %ymm1, %ymm1
+; AVXVNNI-NEXT:    vextracti128 $1, %ymm4, %xmm3
+; AVXVNNI-NEXT:    vpshufd {{.*#+}} xmm5 = xmm3[2,3,2,3]
+; AVXVNNI-NEXT:    vpmovzxbd {{.*#+}} ymm5 = xmm5[0],zero,zero,zero,xmm5[1],zero,zero,zero,xmm5[2],zero,zero,zero,xmm5[3],zero,zero,zero,xmm5[4],zero,zero,zero,xmm5[5],zero,zero,zero,xmm5[6],zero,zero,zero,xmm5[7],zero,zero,zero
+; AVXVNNI-NEXT:    vpmaddwd %ymm5, %ymm10, %ymm5
+; AVXVNNI-NEXT:    vpmovzxbd {{.*#+}} ymm3 = xmm3[0],zero,zero,zero,xmm3[1],zero,zero,zero,xmm3[2],zero,zero,zero,xmm3[3],zero,zero,zero,xmm3[4],zero,zero,zero,xmm3[5],zero,zero,zero,xmm3[6],zero,zero,zero,xmm3[7],zero,zero,zero
+; AVXVNNI-NEXT:    vpmaddwd %ymm3, %ymm9, %ymm3
+; AVXVNNI-NEXT:    vpshufd {{.*#+}} xmm6 = xmm4[2,3,2,3]
+; AVXVNNI-NEXT:    vpmovzxbd {{.*#+}} ymm6 = xmm6[0],zero,zero,zero,xmm6[1],zero,zero,zero,xmm6[2],zero,zero,zero,xmm6[3],zero,zero,zero,xmm6[4],zero,zero,zero,xmm6[5],zero,zero,zero,xmm6[6],zero,zero,zero,xmm6[7],zero,zero,zero
+; AVXVNNI-NEXT:    vpmovzxbd {{.*#+}} ymm4 = xmm4[0],zero,zero,zero,xmm4[1],zero,zero,zero,xmm4[2],zero,zero,zero,xmm4[3],zero,zero,zero,xmm4[4],zero,zero,zero,xmm4[5],zero,zero,zero,xmm4[6],zero,zero,zero,xmm4[7],zero,zero,zero
+; AVXVNNI-NEXT:    {vex} vpdpwssd %ymm4, %ymm2, %ymm0
+; AVXVNNI-NEXT:    {vex} vpdpwssd %ymm6, %ymm11, %ymm0
+; AVXVNNI-NEXT:    vpaddd %ymm3, %ymm0, %ymm0
+; AVXVNNI-NEXT:    vpaddd %ymm5, %ymm0, %ymm0
+; AVXVNNI-NEXT:    retq
+;
+; AVXVNNIINT8INT16-LABEL: partial_reduce_umla_i8_v16i32:
+; AVXVNNIINT8INT16:       # %bb.0:
+; AVXVNNIINT8INT16-NEXT:    vextractf64x4 $1, %zmm2, %ymm3
+; AVXVNNIINT8INT16-NEXT:    vextractf64x4 $1, %zmm1, %ymm4
+; AVXVNNIINT8INT16-NEXT:    vextractf64x4 $1, %zmm0, %ymm5
+; AVXVNNIINT8INT16-NEXT:    vpdpbuud %ymm3, %ymm4, %ymm5
+; AVXVNNIINT8INT16-NEXT:    vpdpbuud %ymm2, %ymm1, %ymm0
+; AVXVNNIINT8INT16-NEXT:    vinsertf64x4 $1, %ymm5, %zmm0, %zmm0
+; AVXVNNIINT8INT16-NEXT:    retq
+;
+; AVX10-LABEL: partial_reduce_umla_i8_v16i32:
+; AVX10:       # %bb.0:
+; AVX10-NEXT:    vpdpbuud %zmm2, %zmm1, %zmm0
+; AVX10-NEXT:    retq
+  %a.zext = zext <64 x i8> %a to <64 x i32>
+  %b.zext = zext <64 x i8> %b to <64 x i32>
+  %mul = mul nsw <64 x i32> %a.zext, %b.zext
+  %res = call <16 x i32> @llvm.vector.partial.reduce.add.v16i32.v64i32(<16 x i32> %acc, <64 x i32> %mul)
+  ret <16 x i32> %res
+}
+
+; VNNI-INT16 i16 x i16 -> i32 SUMLA (sext x zext) tests: VPDPWSUD
+
+; Test 128-bit (xmm) SUMLA i16.
+define <4 x i32> @partial_reduce_sumla_i16_v4i32(<4 x i32> %acc, <8 x i16> %a, <8 x i16> %b) {
+; AVX512VNNI-LABEL: partial_reduce_sumla_i16_v4i32:
+; AVX512VNNI:       # %bb.0:
+; AVX512VNNI-NEXT:    vpmovsxwd %xmm1, %ymm1
+; AVX512VNNI-NEXT:    vpmovzxwd {{.*#+}} ymm2 = xmm2[0],zero,xmm2[1],zero,xmm2[2],zero,xmm2[3],zero,xmm2[4],zero,xmm2[5],zero,xmm2[6],zero,xmm2[7],zero
+; AVX512VNNI-NEXT:    vpmulld %ymm2, %ymm1, %ymm1
+; AVX512VNNI-NEXT:    vpaddd %xmm1, %xmm0, %xmm0
+; AVX512VNNI-NEXT:    vextracti128 $1, %ymm1, %xmm1
+; AVX512VNNI-NEXT:    vpaddd %xmm0, %xmm1, %xmm0
+; AVX512VNNI-NEXT:    vzeroupper
+; AVX512VNNI-NEXT:    retq
+;
+; AVXVNNI-LABEL: partial_reduce_sumla_i16_v4i32:
+; AVXVNNI:       # %bb.0:
+; AVXVNNI-NEXT:    vpmovsxwd %xmm1, %ymm1
+; AVXVNNI-NEXT:    vpmovzxwd {{.*#+}} ymm2 = xmm2[0],zero,xmm2[1],zero,xmm2[2],zero,xmm2[3],zero,xmm2[4],zero,xmm2[5],zero,xmm2[6],zero,xmm2[7],zero
+; AVXVNNI-NEXT:    vpmulld %ymm2, %ymm1, %ymm1
+; AVXVNNI-NEXT:    vpaddd %xmm1, %xmm0, %xmm0
+; AVXVNNI-NEXT:    vextracti128 $1, %ymm1, %xmm1
+; AVXVNNI-NEXT:    vpaddd %xmm0, %xmm1, %xmm0
+; AVXVNNI-NEXT:    vzeroupper
+; AVXVNNI-NEXT:    retq
+;
+; AVXVNNIINT8INT16-LABEL: partial_reduce_sumla_i16_v4i32:
+; AVXVNNIINT8INT16:       # %bb.0:
+; AVXVNNIINT8INT16-NEXT:    vpdpwsud %xmm2, %xmm1, %xmm0
+; AVXVNNIINT8INT16-NEXT:    retq
+;
+; AVX10-LABEL: partial_reduce_sumla_i16_v4i32:
+; AVX10:       # %bb.0:
+; AVX10-NEXT:    vpdpwsud %xmm2, %xmm1, %xmm0
+; AVX10-NEXT:    retq
+  %a.sext = sext <8 x i16> %a to <8 x i32>
+  %b.zext = zext <8 x i16> %b to <8 x i32>
+  %mul = mul nsw <8 x i32> %a.sext, %b.zext
+  %res = call <4 x i32> @llvm.vector.partial.reduce.add.v4i32.v8i32(<4 x i32> %acc, <8 x i32> %mul)
+  ret <4 x i32> %res
+}
+
+; Test 256-bit (ymm) SUMLA i16.
+define <8 x i32> @partial_reduce_sumla_i16_v8i32(<8 x i32> %acc, <16 x i16> %a, <16 x i16> %b) {
+; AVX512VNNI-LABEL: partial_reduce_sumla_i16_v8i32:
+; AVX512VNNI:       # %bb.0:
+; AVX512VNNI-NEXT:    vpmovsxwd %ymm1, %zmm1
+; AVX512VNNI-NEXT:    vpmovzxwd {{.*#+}} zmm2 = ymm2[0],zero,ymm2[1],zero,ymm2[2],zero,ymm2[3],zero,ymm2[4],zero,ymm2[5],zero,ymm2[6],zero,ymm2[7],zero,ymm2[8],zero,ymm2[9],zero,ymm2[10],zero,ymm2[11],zero,ymm2[12],zero,ymm2[13],zero,ymm2[14],zero,ymm2[15],zero
+; AVX512VNNI-NEXT:    vpmulld %zmm2, %zmm1, %zmm1
+; AVX512VNNI-NEXT:    vpaddd %ymm1, %ymm0, %ymm0
+; AVX512VNNI-NEXT:    vextracti64x4 $1, %zmm1, %ymm1
+; AVX512VNNI-NEXT:    vpaddd %ymm0, %ymm1, %ymm0
+; AVX512VNNI-NEXT:    retq
+;
+; AVXVNNI-LABEL: partial_reduce_sumla_i16_v8i32:
+; AVXVNNI:       # %bb.0:
+; AVXVNNI-NEXT:    vpmovsxwd %xmm1, %ymm3
+; AVXVNNI-NEXT:    vextracti128 $1, %ymm1, %xmm1
+; AVXVNNI-NEXT:    vpmovsxwd %xmm1, %ymm1
+; AVXVNNI-NEXT:    vpmovzxwd {{.*#+}} ymm4 = xmm2[0],zero,xmm2[1],zero,xmm2[2],zero,xmm2[3],zero,xmm2[4],zero,xmm2[5],zero,xmm2[6],zero,xmm2[7],zero
+; AVXVNNI-NEXT:    vpmulld %ymm4, %ymm3, %ymm3
+; AVXVNNI-NEXT:    vextracti128 $1, %ymm2, %xmm2
+; AVXVNNI-NEXT:    vpmovzxwd {{.*#+}} ymm2 = xmm2[0],zero,xmm2[1],zero,xmm2[2],zero,xmm2[3],zero,xmm2[4],zero,xmm2[5],zero,xmm2[6],zero,xmm2[7],zero
+; AVXVNNI-NEXT:    vpmulld %ymm2, %ymm1, %ymm1
+; AVXVNNI-NEXT:    vpaddd %ymm3, %ymm0, %ymm0
+; AVXVNNI-NEXT:    vpaddd %ymm1, %ymm0, %ymm0
+; AVXVNNI-NEXT:    retq
+;
+; AVXVNNIINT8INT16-LABEL: partial_reduce_sumla_i16_v8i32:
+; AVXVNNIINT8INT16:       # %bb.0:
+; AVXVNNIINT8INT16-NEXT:    vpdpwsud %ymm2, %ymm1, %ymm0
+; AVXVNNIINT8INT16-NEXT:    retq
+;
+; AVX10-LABEL: partial_reduce_sumla_i16_v8i32:
+; AVX10:       # %bb.0:
+; AVX10-NEXT:    vpdpwsud %ymm2, %ymm1, %ymm0
+; AVX10-NEXT:    retq
+  %a.sext = sext <16 x i16> %a to <16 x i32>
+  %b.zext = zext <16 x i16> %b to <16 x i32>
+  %mul = mul nsw <16 x i32> %a.sext, %b.zext
+  %res = call <8 x i32> @llvm.vector.partial.reduce.add.v8i32.v16i32(<8 x i32> %acc, <16 x i32> %mul)
+  ret <8 x i32> %res
+}
+
+; Test 512-bit (zmm) SUMLA i16.
+define <16 x i32> @partial_reduce_sumla_i16_v16i32(<16 x i32> %acc, <32 x i16> %a, <32 x i16> %b) {
+; AVX512VNNI-LABEL: partial_reduce_sumla_i16_v16i32:
+; AVX512VNNI:       # %bb.0:
+; AVX512VNNI-NEXT:    vpmovsxwd %ymm1, %zmm3
+; AVX512VNNI-NEXT:    vextracti64x4 $1, %zmm1, %ymm1
+; AVX512VNNI-NEXT:    vpmovsxwd %ymm1, %zmm1
+; AVX512VNNI-NEXT:    vpmovzxwd {{.*#+}} zmm4 = ymm2[0],zero,ymm2[1],zero,ymm2[2],zero,ymm2[3],zero,ymm2[4],zero,ymm2[5],zero,ymm2[6],zero,ymm2[7],zero,ymm2[8],zero,ymm2[9],zero,ymm2[10],zero,ymm2[11],zero,ymm2[12],zero,ymm2[13],zero,ymm2[14],zero,ymm2[15],zero
+; AVX512VNNI-NEXT:    vpmulld %zmm4, %zmm3, %zmm3
+; AVX512VNNI-NEXT:    vextracti64x4 $1, %zmm2, %ymm2
+; AVX512VNNI-NEXT:    vpmovzxwd {{.*#+}} zmm2 = ymm2[0],zero,ymm2[1],zero,ymm2[2],zero,ymm2[3],zero,ymm2[4],zero,ymm2[5],zero,ymm2[6],zero,ymm2[7],zero,ymm2[8],zero,ymm2[9],zero,ymm2[10],zero,ymm2[11],zero,ymm2[12],zero,ymm2[13],zero,ymm2[14],zero,ymm2[15],zero
+; AVX512VNNI-NEXT:    vpmulld %zmm2, %zmm1, %zmm1
+; AVX512VNNI-NEXT:    vpaddd %zmm3, %zmm0, %zmm0
+; AVX512VNNI-NEXT:    vpaddd %zmm1, %zmm0, %zmm0
+; AVX512VNNI-NEXT:    retq
+;
+; AVXVNNI-LABEL: partial_reduce_sumla_i16_v16i32:
+; AVXVNNI:       # %bb.0:
+; AVXVNNI-NEXT:    vpmovsxwd %xmm2, %ymm6
+; AVXVNNI-NEXT:    vextracti128 $1, %ymm2, %xmm2
+; AVXVNNI-NEXT:    vpmovsxwd %xmm2, %ymm2
+; AVXVNNI-NEXT:    vpmovsxwd %xmm3, %ymm7
+; AVXVNNI-NEXT:    vextracti128 $1, %ymm3, %xmm3
+; AVXVNNI-NEXT:    vpmovsxwd %xmm3, %ymm3
+; AVXVNNI-NEXT:    vpmovzxwd {{.*#+}} ymm8 = xmm4[0],zero,xmm4[1],zero,xmm4[2],zero,xmm4[3],zero,xmm4[4],zero,xmm4[5],zero,xmm4[6],zero,xmm4[7],zero
+; AVXVNNI-NEXT:    vpmulld %ymm8, %ymm6, %ymm6
+; AVXVNNI-NEXT:    vextracti128 $1, %ymm4, %xmm4
+; AVXVNNI-NEXT:    vpmovzxwd {{.*#+}} ymm4 = xmm4[0],zero,xmm4[1],zero,xmm4[2],zero,xmm4[3],zero,xmm4[4],zero,xmm4[5],zero,xmm4[6],zero,xmm4[7],zero
+; AVXVNNI-NEXT:    vpmulld %ymm4, %ymm2, %ymm2
+; AVXVNNI-NEXT:    vpmovzxwd {{.*#+}} ymm4 = xmm5[0],zero,xmm5[1],zero,xmm5[2],zero,xmm5[3],zero,xmm5[4],zero,xmm5[5],zero,xmm5[6],zero,xmm5[7],zero
+; AVXVNNI-NEXT:    vpmulld %ymm4, %ymm7, %ymm4
+; AVXVNNI-NEXT:    vextracti128 $1, %ymm5, %xmm5
+; AVXVNNI-NEXT:    vpmovzxwd {{.*#+}} ymm5 = xmm5[0],zero,xmm5[1],zero,xmm5[2],zero,xmm5[3],zero,xmm5[4],zero,xmm5[5],zero,xmm5[6],zero,xmm5[7],zero
+; AVXVNNI-NEXT:    vpmulld %ymm5, %ymm3, %ymm3
+; AVXVNNI-NEXT:    vpaddd %ymm6, %ymm0, %ymm0
+; AVXVNNI-NEXT:    vpaddd %ymm2, %ymm0, %ymm0
+; AVXVNNI-NEXT:    vpaddd %ymm4, %ymm1, %ymm1
+; AVXVNNI-NEXT:    vpaddd %ymm3, %ymm1, %ymm1
+; AVXVNNI-NEXT:    retq
+;
+; AVXVNNIINT8INT16-LABEL: partial_reduce_sumla_i16_v16i32:
+; AVXVNNIINT8INT16:       # %bb.0:
+; AVXVNNIINT8INT16-NEXT:    vextractf64x4 $1, %zmm2, %ymm3
+; AVXVNNIINT8INT16-NEXT:    vextractf64x4 $1, %zmm1, %ymm4
+; AVXVNNIINT8INT16-NEXT:    vextractf64x4 $1, %zmm0, %ymm5
+; AVXVNNIINT8INT16-NEXT:    vpdpwsud %ymm3, %ymm4, %ymm5
+; AVXVNNIINT8INT16-NEXT:    vpdpwsud %ymm2, %ymm1, %ymm0
+; AVXVNNIINT8INT16-NEXT:    vinsertf64x4 $1, %ymm5, %zmm0, %zmm0
+; AVXVNNIINT8INT16-NEXT:    retq
+;
+; AVX10-LABEL: partial_reduce_sumla_i16_v16i32:
+; AVX10:       # %bb.0:
+; AVX10-NEXT:    vpdpwsud %zmm2, %zmm1, %zmm0
+; AVX10-NEXT:    retq
+  %a.sext = sext <32 x i16> %a to <32 x i32>
+  %b.zext = zext <32 x i16> %b to <32 x i32>
+  %mul = mul nsw <32 x i32> %a.sext, %b.zext
+  %res = call <16 x i32> @llvm.vector.partial.reduce.add.v16i32.v32i32(<16 x i32> %acc, <32 x i32> %mul)
+  ret <16 x i32> %res
+}
+
+; VNNI-INT16 i16 x i16 -> i32 UMLA (zext x zext) tests: VPDPWUUD
+
+; Test 128-bit (xmm) UMLA i16.
+define <4 x i32> @partial_reduce_umla_i16_v4i32(<4 x i32> %acc, <8 x i16> %a, <8 x i16> %b) {
+; AVX512VNNI-LABEL: partial_reduce_umla_i16_v4i32:
+; AVX512VNNI:       # %bb.0:
+; AVX512VNNI-NEXT:    vpmovzxwd {{.*#+}} ymm1 = xmm1[0],zero,xmm1[1],zero,xmm1[2],zero,xmm1[3],zero,xmm1[4],zero,xmm1[5],zero,xmm1[6],zero,xmm1[7],zero
+; AVX512VNNI-NEXT:    vpmovzxwd {{.*#+}} ymm2 = xmm2[0],zero,xmm2[1],zero,xmm2[2],zero,xmm2[3],zero,xmm2[4],zero,xmm2[5],zero,xmm2[6],zero,xmm2[7],zero
+; AVX512VNNI-NEXT:    vpmulld %ymm2, %ymm1, %ymm1
+; AVX512VNNI-NEXT:    vpaddd %xmm1, %xmm0, %xmm0
+; AVX512VNNI-NEXT:    vextracti128 $1, %ymm1, %xmm1
+; AVX512VNNI-NEXT:    vpaddd %xmm0, %xmm1, %xmm0
+; AVX512VNNI-NEXT:    vzeroupper
+; AVX512VNNI-NEXT:    retq
+;
+; AVXVNNI-LABEL: partial_reduce_umla_i16_v4i32:
+; AVXVNNI:       # %bb.0:
+; AVXVNNI-NEXT:    vpmovzxwd {{.*#+}} ymm1 = xmm1[0],zero,xmm1[1],zero,xmm1[2],zero,xmm1[3],zero,xmm1[4],zero,xmm1[5],zero,xmm1[6],zero,xmm1[7],zero
+; AVXVNNI-NEXT:    vpmovzxwd {{.*#+}} ymm2 = xmm2[0],zero,xmm2[1],zero,xmm2[2],zero,xmm2[3],zero,xmm2[4],zero,xmm2[5],zero,xmm2[6],zero,xmm2[7],zero
+; AVXVNNI-NEXT:    vpmulld %ymm2, %ymm1, %ymm1
+; AVXVNNI-NEXT:    vpaddd %xmm1, %xmm0, %xmm0
+; AVXVNNI-NEXT:    vextracti128 $1, %ymm1, %xmm1
+; AVXVNNI-NEXT:    vpaddd %xmm0, %xmm1, %xmm0
+; AVXVNNI-NEXT:    vzeroupper
+; AVXVNNI-NEXT:    retq
+;
+; AVXVNNIINT8INT16-LABEL: partial_reduce_umla_i16_v4i32:
+; AVXVNNIINT8INT16:       # %bb.0:
+; AVXVNNIINT8INT16-NEXT:    vpdpwuud %xmm2, %xmm1, %xmm0
+; AVXVNNIINT8INT16-NEXT:    retq
+;
+; AVX10-LABEL: partial_reduce_umla_i16_v4i32:
+; AVX10:       # %bb.0:
+; AVX10-NEXT:    vpdpwuud %xmm2, %xmm1, %xmm0
+; AVX10-NEXT:    retq
+  %a.zext = zext <8 x i16> %a to <8 x i32>
+  %b.zext = zext <8 x i16> %b to <8 x i32>
+  %mul = mul nsw <8 x i32> %a.zext, %b.zext
+  %res = call <4 x i32> @llvm.vector.partial.reduce.add.v4i32.v8i32(<4 x i32> %acc, <8 x i32> %mul)
+  ret <4 x i32> %res
+}
+
+; Test 256-bit (ymm) UMLA i16.
+define <8 x i32> @partial_reduce_umla_i16_v8i32(<8 x i32> %acc, <16 x i16> %a, <16 x i16> %b) {
+; AVX512VNNI-LABEL: partial_reduce_umla_i16_v8i32:
+; AVX512VNNI:       # %bb.0:
+; AVX512VNNI-NEXT:    vpmovzxwd {{.*#+}} zmm1 = ymm1[0],zero,ymm1[1],zero,ymm1[2],zero,ymm1[3],zero,ymm1[4],zero,ymm1[5],zero,ymm1[6],zero,ymm1[7],zero,ymm1[8],zero,ymm1[9],zero,ymm1[10],zero,ymm1[11],zero,ymm1[12],zero,ymm1[13],zero,ymm1[14],zero,ymm1[15],zero
+; AVX512VNNI-NEXT:    vpmovzxwd {{.*#+}} zmm2 = ymm2[0],zero,ymm2[1],zero,ymm2[2],zero,ymm2[3],zero,ymm2[4],zero,ymm2[5],zero,ymm2[6],zero,ymm2[7],zero,ymm2[8],zero,ymm2[9],zero,ymm2[10],zero,ymm2[11],zero,ymm2[12],zero,ymm2[13],zero,ymm2[14],zero,ymm2[15],zero
+; AVX512VNNI-NEXT:    vpmulld %zmm2, %zmm1, %zmm1
+; AVX512VNNI-NEXT:    vpaddd %ymm1, %ymm0, %ymm0
+; AVX512VNNI-NEXT:    vextracti64x4 $1, %zmm1, %ymm1
+; AVX512VNNI-NEXT:    vpaddd %ymm0, %ymm1, %ymm0
+; AVX512VNNI-NEXT:    retq
+;
+; AVXVNNI-LABEL: partial_reduce_umla_i16_v8i32:
+; AVXVNNI:       # %bb.0:
+; AVXVNNI-NEXT:    vpmovzxwd {{.*#+}} ymm3 = xmm1[0],zero,xmm1[1],zero,xmm1[2],zero,xmm1[3],zero,xmm1[4],zero,xmm1[5],zero,xmm1[6],zero,xmm1[7],zero
+; AVXVNNI-NEXT:    vextracti128 $1, %ymm1, %xmm1
+; AVXVNNI-NEXT:    vpmovzxwd {{.*#+}} ymm1 = xmm1[0],zero,xmm1[1],zero,xmm1[2],zero,xmm1[3],zero,xmm1[4],zero,xmm1[5],zero,xmm1[6],zero,xmm1[7],zero
+; AVXVNNI-NEXT:    vpmovzxwd {{.*#+}} ymm4 = xmm2[0],zero,xmm2[1],zero,xmm2[2],zero,xmm2[3],zero,xmm2[4],zero,xmm2[5],zero,xmm2[6],zero,xmm2[7],zero
+; AVXVNNI-NEXT:    vpmulld %ymm4, %ymm3, %ymm3
+; AVXVNNI-NEXT:    vextracti128 $1, %ymm2, %xmm2
+; AVXVNNI-NEXT:    vpmovzxwd {{.*#+}} ymm2 = xmm2[0],zero,xmm2[1],zero,xmm2[2],zero,xmm2[3],zero,xmm2[4],zero,xmm2[5],zero,xmm2[6],zero,xmm2[7],zero
+; AVXVNNI-NEXT:    vpmulld %ymm2, %ymm1, %ymm1
+; AVXVNNI-NEXT:    vpaddd %ymm3, %ymm0, %ymm0
+; AVXVNNI-NEXT:    vpaddd %ymm1, %ymm0, %ymm0
+; AVXVNNI-NEXT:    retq
+;
+; AVXVNNIINT8INT16-LABEL: partial_reduce_umla_i16_v8i32:
+; AVXVNNIINT8INT16:       # %bb.0:
+; AVXVNNIINT8INT16-NEXT:    vpdpwuud %ymm2, %ymm1, %ymm0
+; AVXVNNIINT8INT16-NEXT:    retq
+;
+; AVX10-LABEL: partial_reduce_umla_i16_v8i32:
+; AVX10:       # %bb.0:
+; AVX10-NEXT:    vpdpwuud %ymm2, %ymm1, %ymm0
+; AVX10-NEXT:    retq
+  %a.zext = zext <16 x i16> %a to <16 x i32>
+  %b.zext = zext <16 x i16> %b to <16 x i32>
+  %mul = mul nsw <16 x i32> %a.zext, %b.zext
+  %res = call <8 x i32> @llvm.vector.partial.reduce.add.v8i32.v16i32(<8 x i32> %acc, <16 x i32> %mul)
+  ret <8 x i32> %res
+}
+
+; Test 512-bit (zmm) UMLA i16.
+define <16 x i32> @partial_reduce_umla_i16_v16i32(<16 x i32> %acc, <32 x i16> %a, <32 x i16> %b) {
+; AVX512VNNI-LABEL: partial_reduce_umla_i16_v16i32:
+; AVX512VNNI:       # %bb.0:
+; AVX512VNNI-NEXT:    vpmovzxwd {{.*#+}} zmm3 = ymm1[0],zero,ymm1[1],zero,ymm1[2],zero,ymm1[3],zero,ymm1[4],zero,ymm1[5],zero,ymm1[6],zero,ymm1[7],zero,ymm1[8],zero,ymm1[9],zero,ymm1[10],zero,ymm1[11],zero,ymm1[12],zero,ymm1[13],zero,ymm1[14],zero,ymm1[15],zero
+; AVX512VNNI-NEXT:    vextracti64x4 $1, %zmm1, %ymm1
+; AVX512VNNI-NEXT:    vpmovzxwd {{.*#+}} zmm1 = ymm1[0],zero,ymm1[1],zero,ymm1[2],zero,ymm1[3],zero,ymm1[4],zero,ymm1[5],zero,ymm1[6],zero,ymm1[7],zero,ymm1[8],zero,ymm1[9],zero,ymm1[10],zero,ymm1[11],zero,ymm1[12],zero,ymm1[13],zero,ymm1[14],zero,ymm1[15],zero
+; AVX512VNNI-NEXT:    vpmovzxwd {{.*#+}} zmm4 = ymm2[0],zero,ymm2[1],zero,ymm2[2],zero,ymm2[3],zero,ymm2[4],zero,ymm2[5],zero,ymm2[6],zero,ymm2[7],zero,ymm2[8],zero,ymm2[9],zero,ymm2[10],zero,ymm2[11],zero,ymm2[12],zero,ymm2[13],zero,ymm2[14],zero,ymm2[15],zero
+; AVX512VNNI-NEXT:    vpmulld %zmm4, %zmm3, %zmm3
+; AVX512VNNI-NEXT:    vextracti64x4 $1, %zmm2, %ymm2
+; AVX512VNNI-NEXT:    vpmovzxwd {{.*#+}} zmm2 = ymm2[0],zero,ymm2[1],zero,ymm2[2],zero,ymm2[3],zero,ymm2[4],zero,ymm2[5],zero,ymm2[6],zero,ymm2[7],zero,ymm2[8],zero,ymm2[9],zero,ymm2[10],zero,ymm2[11],zero,ymm2[12],zero,ymm2[13],zero,ymm2[14],zero,ymm2[15],zero
+; AVX512VNNI-NEXT:    vpmulld %zmm2, %zmm1, %zmm1
+; AVX512VNNI-NEXT:    vpaddd %zmm3, %zmm0, %zmm0
+; AVX512VNNI-NEXT:    vpaddd %zmm1, %zmm0, %zmm0
+; AVX512VNNI-NEXT:    retq
+;
+; AVXVNNI-LABEL: partial_reduce_umla_i16_v16i32:
+; AVXVNNI:       # %bb.0:
+; AVXVNNI-NEXT:    vpmovzxwd {{.*#+}} ymm6 = xmm2[0],zero,xmm2[1],zero,xmm2[2],zero,xmm2[3],zero,xmm2[4],zero,xmm2[5],zero,xmm2[6],zero,xmm2[7],zero
+; AVXVNNI-NEXT:    vextracti128 $1, %ymm2, %xmm2
+; AVXVNNI-NEXT:    vpmovzxwd {{.*#+}} ymm2 = xmm2[0],zero,xmm2[1],zero,xmm2[2],zero,xmm2[3],zero,xmm2[4],zero,xmm2[5],zero,xmm2[6],zero,xmm2[7],zero
+; AVXVNNI-NEXT:    vpmovzxwd {{.*#+}} ymm7 = xmm3[0],zero,xmm3[1],zero,xmm3[2],zero,xmm3[3],zero,xmm3[4],zero,xmm3[5],zero,xmm3[6],zero,xmm3[7],zero
+; AVXVNNI-NEXT:    vextracti128 $1, %ymm3, %xmm3
+; AVXVNNI-NEXT:    vpmovzxwd {{.*#+}} ymm3 = xmm3[0],zero,xmm3[1],zero,xmm3[2],zero,xmm3[3],zero,xmm3[4],zero,xmm3[5],zero,xmm3[6],zero,xmm3[7],zero
+; AVXVNNI-NEXT:    vpmovzxwd {{.*#+}} ymm8 = xmm4[0],zero,xmm4[1],zero,xmm4[2],zero,xmm4[3],zero,xmm4[4],zero,xmm4[5],zero,xmm4[6],zero,xmm4[7],zero
+; AVXVNNI-NEXT:    vpmulld %ymm8, %ymm6, %ymm6
+; AVXVNNI-NEXT:    vextracti128 $1, %ymm4, %xmm4
+; AVXVNNI-NEXT:    vpmovzxwd {{.*#+}} ymm4 = xmm4[0],zero,xmm4[1],zero,xmm4[2],zero,xmm4[3],zero,xmm4[4],zero,xmm4[5],zero,xmm4[6],zero,xmm4[7],zero
+; AVXVNNI-NEXT:    vpmulld %ymm4, %ymm2, %ymm2
+; AVXVNNI-NEXT:    vpmovzxwd {{.*#+}} ymm4 = xmm5[0],zero,xmm5[1],zero,xmm5[2],zero,xmm5[3],zero,xmm5[4],zero,xmm5[5],zero,xmm5[6],zero,xmm5[7],zero
+; AVXVNNI-NEXT:    vpmulld %ymm4, %ymm7, %ymm4
+; AVXVNNI-NEXT:    vextracti128 $1, %ymm5, %xmm5
+; AVXVNNI-NEXT:    vpmovzxwd {{.*#+}} ymm5 = xmm5[0],zero,xmm5[1],zero,xmm5[2],zero,xmm5[3],zero,xmm5[4],zero,xmm5[5],zero,xmm5[6],zero,xmm5[7],zero
+; AVXVNNI-NEXT:    vpmulld %ymm5, %ymm3, %ymm3
+; AVXVNNI-NEXT:    vpaddd %ymm6, %ymm0, %ymm0
+; AVXVNNI-NEXT:    vpaddd %ymm2, %ymm0, %ymm0
+; AVXVNNI-NEXT:    vpaddd %ymm4, %ymm1, %ymm1
+; AVXVNNI-NEXT:    vpaddd %ymm3, %ymm1, %ymm1
+; AVXVNNI-NEXT:    retq
+;
+; AVXVNNIINT8INT16-LABEL: partial_reduce_umla_i16_v16i32:
+; AVXVNNIINT8INT16:       # %bb.0:
+; AVXVNNIINT8INT16-NEXT:    vextractf64x4 $1, %zmm2, %ymm3
+; AVXVNNIINT8INT16-NEXT:    vextractf64x4 $1, %zmm1, %ymm4
+; AVXVNNIINT8INT16-NEXT:    vextractf64x4 $1, %zmm0, %ymm5
+; AVXVNNIINT8INT16-NEXT:    vpdpwuud %ymm3, %ymm4, %ymm5
+; AVXVNNIINT8INT16-NEXT:    vpdpwuud %ymm2, %ymm1, %ymm0
+; AVXVNNIINT8INT16-NEXT:    vinsertf64x4 $1, %ymm5, %zmm0, %zmm0
+; AVXVNNIINT8INT16-NEXT:    retq
+;
+; AVX10-LABEL: partial_reduce_umla_i16_v16i32:
+; AVX10:       # %bb.0:
+; AVX10-NEXT:    vpdpwuud %zmm2, %zmm1, %zmm0
+; AVX10-NEXT:    retq
+  %a.zext = zext <32 x i16> %a to <32 x i32>
+  %b.zext = zext <32 x i16> %b to <32 x i32>
+  %mul = mul nsw <32 x i32> %a.zext, %b.zext
+  %res = call <16 x i32> @llvm.vector.partial.reduce.add.v16i32.v32i32(<16 x i32> %acc, <32 x i32> %mul)
+  ret <16 x i32> %res
+}
+
+declare <4 x i32> @llvm.vector.partial.reduce.add.v4i32.v8i32(<4 x i32>, <8 x i32>)
+declare <8 x i32> @llvm.vector.partial.reduce.add.v8i32.v16i32(<8 x i32>, <16 x i32>)
+declare <16 x i32> @llvm.vector.partial.reduce.add.v16i32.v32i32(<16 x i32>, <32 x i32>)
+declare <4 x float> @llvm.vector.partial.reduce.fadd.v4f32.v8f32(<4 x float>, <8 x float>)
+declare <8 x float> @llvm.vector.partial.reduce.fadd.v8f32.v16f32(<8 x float>, <16 x float>)
+declare <16 x float> @llvm.vector.partial.reduce.fadd.v16f32.v32f32(<16 x float>, <32 x float>)
diff --git a/llvm/test/CodeGen/X86/vector-partial-reduce-small-vf.ll b/llvm/test/CodeGen/X86/vector-partial-reduce-small-vf.ll
new file mode 100644
index 0000000000000..ee0a4519fb5d1
--- /dev/null
+++ b/llvm/test/CodeGen/X86/vector-partial-reduce-small-vf.ll
@@ -0,0 +1,45 @@
+; NOTE: Assertions have been autogenerated by utils/update_llc_test_checks.py UTC_ARGS: --version 6
+; RUN: llc -mtriple=x86_64-unknown-linux-gnu -mattr=+avx512f,+avx512bw,+avx512vl,+avx512vnni,+avx512bf16 < %s | FileCheck %s
+
+; Unsupported small partial-reduction shapes should expand or widen safely
+; instead of crashing during type legalization.
+
+define <2 x i32> @partial_reduce_sumla_i8_v2i32(<2 x i32> %acc, <8 x i8> %a, <8 x i8> %b) {
+; CHECK-LABEL: partial_reduce_sumla_i8_v2i32:
+; CHECK:       # %bb.0:
+; CHECK-NEXT:    vpdpbusd %xmm2, %xmm1, %xmm0
+; CHECK-NEXT:    retq
+  %a.zext = zext <8 x i8> %a to <8 x i32>
+  %b.sext = sext <8 x i8> %b to <8 x i32>
+  %mul = mul nsw <8 x i32> %a.zext, %b.sext
+  %res = call <2 x i32> @llvm.vector.partial.reduce.add.v2i32.v8i32(<2 x i32> %acc, <8 x i32> %mul)
+  ret <2 x i32> %res
+}
+
+define <2 x i32> @partial_reduce_smla_i16_v2i32(<2 x i32> %acc, <4 x i16> %a, <4 x i16> %b) {
+; CHECK-LABEL: partial_reduce_smla_i16_v2i32:
+; CHECK:       # %bb.0:
+; CHECK-NEXT:    vpdpwssd %xmm2, %xmm1, %xmm0
+; CHECK-NEXT:    retq
+  %a.sext = sext <4 x i16> %a to <4 x i32>
+  %b.sext = sext <4 x i16> %b to <4 x i32>
+  %mul = mul nsw <4 x i32> %a.sext, %b.sext
+  %res = call <2 x i32> @llvm.vector.partial.reduce.add.v2i32.v4i32(<2 x i32> %acc, <4 x i32> %mul)
+  ret <2 x i32> %res
+}
+
+define <2 x float> @partial_reduce_fmla_bf16_v2f32(<2 x float> %acc, <4 x bfloat> %a, <4 x bfloat> %b) {
+; CHECK-LABEL: partial_reduce_fmla_bf16_v2f32:
+; CHECK:       # %bb.0:
+; CHECK-NEXT:    vdpbf16ps %xmm2, %xmm1, %xmm0
+; CHECK-NEXT:    retq
+  %a.ext = fpext <4 x bfloat> %a to <4 x float>
+  %b.ext = fpext <4 x bfloat> %b to <4 x float>
+  %mul = fmul <4 x float> %a.ext, %b.ext
+  %res = call <2 x float> @llvm.vector.partial.reduce.fadd.v2f32.v4f32(<2 x float> %acc, <4 x float> %mul)
+  ret <2 x float> %res
+}
+
+declare <2 x i32> @llvm.vector.partial.reduce.add.v2i32.v8i32(<2 x i32>, <8 x i32>)
+declare <2 x i32> @llvm.vector.partial.reduce.add.v2i32.v4i32(<2 x i32>, <4 x i32>)
+declare <2 x float> @llvm.vector.partial.reduce.fadd.v2f32.v4f32(<2 x float>, <4 x float>)



More information about the llvm-commits mailing list