[llvm] [RISCV] Reassociate scalar accumulator chains of vector reductions (PR #206471)

via llvm-commits llvm-commits at lists.llvm.org
Mon Jun 29 05:40:28 PDT 2026


llvmorg-github-actions[bot] wrote:


<!--LLVM PR SUMMARY COMMENT-->

@llvm/pr-subscribers-backend-risc-v

Author: Pengcheng Wang (wangpc-pp)

<details>
<summary>Changes</summary>

A left-leaning scalar add chain of vector reductions, e.g.
`add(reduce(x0), add(reduce(x1), add(reduce(x2), acc)))`, is only partially
collapsed by the generic `DAGCombiner::reassociateReduction`, which folds a
single `add(vecreduce(x), vecreduce(y)) -> vecreduce(add(x, y))`. The cascade
breaks for the chains SLP produces (e.g. x264's SAD), leaving one reduction
per term, and with `Zvdot4a8i` one `vdot4a + vredsum` per term.

In this PR, we add a RISC-V combine that reassociates:
```
add(vecreduce_add(X), add(vecreduce_add(Y), Z))
    -> add(vecreduce_add(add(X, Y)), Z)
```

This collapses the whole chain into a single reduction, after which
`foldReduceOperandViaVDOT4A` and `combineVdot4aAccum` chain all dot
products into one accumulator with a single `vredsum`, and the external
accumulator folds in as the reduction start value.

Fixes #<!-- -->168047.

Assisted-by: TraeCli (AI assistant)


---

Patch is 23.06 KiB, truncated to 20.00 KiB below, full version: https://github.com/llvm/llvm-project/pull/206471.diff


4 Files Affected:

- (modified) llvm/lib/Target/RISCV/RISCVISelLowering.cpp (+48) 
- (modified) llvm/test/CodeGen/RISCV/rvv/fixed-vectors-sad.ll (+90-86) 
- (modified) llvm/test/CodeGen/RISCV/rvv/fixed-vectors-zvdot4a8i.ll (+140) 
- (modified) llvm/test/CodeGen/RISCV/rvv/zvdot4a8i-sdnode.ll (+99) 


``````````diff
diff --git a/llvm/lib/Target/RISCV/RISCVISelLowering.cpp b/llvm/lib/Target/RISCV/RISCVISelLowering.cpp
index e5ce36afb7b2f..37f75be08493f 100644
--- a/llvm/lib/Target/RISCV/RISCVISelLowering.cpp
+++ b/llvm/lib/Target/RISCV/RISCVISelLowering.cpp
@@ -16633,6 +16633,52 @@ static SDValue combineAddMulh(SDNode *N, SelectionDAG &DAG,
   return DAG.getNode(RISCVISD::MULHSU, DL, VT, X, Mulh.getOperand(1));
 }
 
+// Reassociate a scalar accumulator chain of vector reductions so the
+// reductions become adjacent and can be merged:
+//   add(vecreduce_add(X), add(vecreduce_add(Y), Z))
+//     -> add(vecreduce_add(add(X, Y)), Z)
+// The generic DAGCombiner only folds add(vecreduce(x), vecreduce(y)) ->
+// vecreduce(add(x, y)) when both operands of a single add are reductions
+// (DAGCombiner::reassociateReduction). SLP often produces a left-leaning
+// scalar chain add(reduce_i, acc) where the cascade breaks, leaving one
+// reduction (and, with Zvdot4a8i, one vredsum) per term. Applied to fixpoint
+// by the combiner worklist, this collapses the whole chain into
+// add(vecreduce_add(BigSum), acc), after which the generic fold and the
+// vdot4a accumulator chaining produce a single reduction.
+static SDValue combineReduceAddChain(SDNode *N, SelectionDAG &DAG,
+                                     const RISCVSubtarget &Subtarget) {
+  using namespace SDPatternMatch;
+  EVT VT = N->getValueType(0);
+  // Only scalar integer adds; vector reductions feed a scalar result.
+  if (!Subtarget.hasVInstructions() || !VT.isScalarInteger())
+    return SDValue();
+
+  // Match add(vecreduce_add(X), Tail) where Tail is add(vecreduce_add(Y), Z).
+  SDValue X, Y, Z;
+  if (!sd_match(N, m_Add(m_OneUse(m_UnaryOp(ISD::VECREDUCE_ADD, m_Value(X))),
+                         m_OneUse(m_Add(m_OneUse(m_UnaryOp(ISD::VECREDUCE_ADD,
+                                                           m_Value(Y))),
+                                        m_Value(Z))))))
+    return SDValue();
+
+  // The two reductions must be over the same vector type to merge them.
+  EVT SrcVT = X.getValueType();
+  if (SrcVT != Y.getValueType())
+    return SDValue();
+
+  // The add on the source vector type must be available.
+  const TargetLowering &TLI = DAG.getTargetLoweringInfo();
+  if (!TLI.isOperationLegalOrCustom(ISD::ADD, SrcVT))
+    return SDValue();
+
+  // Reassociating the wrapping reduction sum invalidates the per-step nsw/nuw
+  // facts, so build the new nodes without flags.
+  SDLoc DL(N);
+  SDValue Sum = DAG.getNode(ISD::ADD, DL, SrcVT, X, Y);
+  SDValue Red = DAG.getNode(ISD::VECREDUCE_ADD, DL, VT, Sum);
+  return DAG.getNode(ISD::ADD, DL, VT, Red, Z);
+}
+
 static SDValue performADDCombine(SDNode *N,
                                  TargetLowering::DAGCombinerInfo &DCI,
                                  const RISCVSubtarget &Subtarget) {
@@ -16647,6 +16693,8 @@ static SDValue performADDCombine(SDNode *N,
     if (SDValue V = combineShlAddIAdd(N, DAG, Subtarget))
       return V;
   }
+  if (SDValue V = combineReduceAddChain(N, DAG, Subtarget))
+    return V;
   if (SDValue V = combineBinOpToReduce(N, DAG, Subtarget))
     return V;
   if (SDValue V = combineBinOpOfExtractToReduceTree(N, DAG, Subtarget))
diff --git a/llvm/test/CodeGen/RISCV/rvv/fixed-vectors-sad.ll b/llvm/test/CodeGen/RISCV/rvv/fixed-vectors-sad.ll
index 3afa4f18e05dd..85b090ee356f2 100644
--- a/llvm/test/CodeGen/RISCV/rvv/fixed-vectors-sad.ll
+++ b/llvm/test/CodeGen/RISCV/rvv/fixed-vectors-sad.ll
@@ -141,40 +141,40 @@ entry:
 define signext i32 @sad_2block_16xi8_as_i32(ptr %a, ptr %b, i32 signext %stridea, i32 signext %strideb) {
 ; CHECK-LABEL: sad_2block_16xi8_as_i32:
 ; CHECK:       # %bb.0: # %entry
+; CHECK-NEXT:    add a4, a0, a2
 ; CHECK-NEXT:    vsetivli zero, 16, e8, m1, ta, ma
-; CHECK-NEXT:    vle8.v v8, (a0)
-; CHECK-NEXT:    vle8.v v9, (a1)
-; CHECK-NEXT:    add a0, a0, a2
-; CHECK-NEXT:    vle8.v v10, (a0)
-; CHECK-NEXT:    add a1, a1, a3
-; CHECK-NEXT:    vminu.vv v11, v8, v9
-; CHECK-NEXT:    vle8.v v12, (a1)
+; CHECK-NEXT:    vle8.v v8, (a4)
+; CHECK-NEXT:    add a5, a1, a3
+; CHECK-NEXT:    vle8.v v9, (a5)
+; CHECK-NEXT:    vminu.vv v10, v8, v9
 ; CHECK-NEXT:    vmaxu.vv v8, v8, v9
-; CHECK-NEXT:    add a0, a0, a2
-; CHECK-NEXT:    vle8.v v9, (a0)
-; CHECK-NEXT:    vsub.vv v8, v8, v11
-; CHECK-NEXT:    add a1, a1, a3
-; CHECK-NEXT:    vle8.v v11, (a1)
-; CHECK-NEXT:    vminu.vv v13, v10, v12
-; CHECK-NEXT:    vmaxu.vv v10, v10, v12
-; CHECK-NEXT:    vminu.vv v12, v9, v11
-; CHECK-NEXT:    vmaxu.vv v9, v9, v11
-; CHECK-NEXT:    vsub.vv v10, v10, v13
-; CHECK-NEXT:    vsub.vv v9, v9, v12
-; CHECK-NEXT:    vwaddu.vv v12, v10, v8
+; CHECK-NEXT:    add a4, a4, a2
+; CHECK-NEXT:    vle8.v v9, (a4)
+; CHECK-NEXT:    add a2, a4, a2
+; CHECK-NEXT:    vsub.vv v8, v8, v10
+; CHECK-NEXT:    vle8.v v10, (a2)
+; CHECK-NEXT:    add a5, a5, a3
 ; CHECK-NEXT:    vsetvli zero, zero, e16, m2, ta, ma
-; CHECK-NEXT:    vzext.vf2 v14, v9
-; CHECK-NEXT:    add a0, a0, a2
-; CHECK-NEXT:    vle8.v v16, (a0)
-; CHECK-NEXT:    vwaddu.vv v8, v14, v12
-; CHECK-NEXT:    add a1, a1, a3
-; CHECK-NEXT:    vle8.v v12, (a1)
+; CHECK-NEXT:    vzext.vf2 v12, v8
+; CHECK-NEXT:    vle8.v v8, (a5)
+; CHECK-NEXT:    add a3, a5, a3
+; CHECK-NEXT:    vle8.v v11, (a3)
 ; CHECK-NEXT:    vsetvli zero, zero, e8, m1, ta, ma
-; CHECK-NEXT:    vminu.vv v13, v16, v12
-; CHECK-NEXT:    vmaxu.vv v12, v16, v12
-; CHECK-NEXT:    vsub.vv v14, v12, v13
+; CHECK-NEXT:    vminu.vv v14, v9, v8
+; CHECK-NEXT:    vmaxu.vv v8, v9, v8
+; CHECK-NEXT:    vminu.vv v9, v10, v11
+; CHECK-NEXT:    vmaxu.vv v10, v10, v11
+; CHECK-NEXT:    vle8.v v11, (a0)
+; CHECK-NEXT:    vsub.vv v8, v8, v14
+; CHECK-NEXT:    vle8.v v14, (a1)
+; CHECK-NEXT:    vsub.vv v9, v10, v9
+; CHECK-NEXT:    vminu.vv v10, v11, v14
+; CHECK-NEXT:    vmaxu.vv v11, v11, v14
+; CHECK-NEXT:    vwaddu.vv v14, v9, v8
+; CHECK-NEXT:    vsub.vv v16, v11, v10
 ; CHECK-NEXT:    vsetvli zero, zero, e16, m2, ta, ma
-; CHECK-NEXT:    vzext.vf2 v12, v14
+; CHECK-NEXT:    vwaddu.vv v8, v14, v12
+; CHECK-NEXT:    vzext.vf2 v12, v16
 ; CHECK-NEXT:    vwaddu.wv v8, v8, v12
 ; CHECK-NEXT:    vsetvli zero, zero, e32, m4, ta, ma
 ; CHECK-NEXT:    vmv.s.x v12, zero
@@ -184,29 +184,31 @@ define signext i32 @sad_2block_16xi8_as_i32(ptr %a, ptr %b, i32 signext %stridea
 ;
 ; ZVABD-LABEL: sad_2block_16xi8_as_i32:
 ; ZVABD:       # %bb.0: # %entry
+; ZVABD-NEXT:    add a4, a0, a2
 ; ZVABD-NEXT:    vsetivli zero, 16, e8, m1, ta, ma
-; ZVABD-NEXT:    vle8.v v8, (a0)
-; ZVABD-NEXT:    vle8.v v9, (a1)
-; ZVABD-NEXT:    add a0, a0, a2
-; ZVABD-NEXT:    vle8.v v10, (a0)
+; ZVABD-NEXT:    vle8.v v8, (a4)
+; ZVABD-NEXT:    add a5, a1, a3
+; ZVABD-NEXT:    vle8.v v9, (a5)
 ; ZVABD-NEXT:    vabdu.vv v8, v8, v9
-; ZVABD-NEXT:    add a0, a0, a2
-; ZVABD-NEXT:    vle8.v v9, (a0)
-; ZVABD-NEXT:    add a1, a1, a3
-; ZVABD-NEXT:    add a4, a1, a3
+; ZVABD-NEXT:    add a4, a4, a2
+; ZVABD-NEXT:    vle8.v v9, (a4)
+; ZVABD-NEXT:    add a5, a5, a3
 ; ZVABD-NEXT:    vsetvli zero, zero, e16, m2, ta, ma
 ; ZVABD-NEXT:    vzext.vf2 v12, v8
-; ZVABD-NEXT:    vle8.v v8, (a4)
-; ZVABD-NEXT:    vle8.v v11, (a1)
+; ZVABD-NEXT:    vle8.v v8, (a5)
 ; ZVABD-NEXT:    vsetvli zero, zero, e8, m1, ta, ma
 ; ZVABD-NEXT:    vabdu.vv v8, v9, v8
-; ZVABD-NEXT:    vwabdau.vv v12, v10, v11
+; ZVABD-NEXT:    add a2, a4, a2
+; ZVABD-NEXT:    vle8.v v9, (a2)
+; ZVABD-NEXT:    add a3, a5, a3
 ; ZVABD-NEXT:    vsetvli zero, zero, e16, m2, ta, ma
 ; ZVABD-NEXT:    vzext.vf2 v14, v8
+; ZVABD-NEXT:    vle8.v v8, (a3)
+; ZVABD-NEXT:    vsetvli zero, zero, e8, m1, ta, ma
+; ZVABD-NEXT:    vwabdau.vv v14, v9, v8
+; ZVABD-NEXT:    vsetvli zero, zero, e16, m2, ta, ma
 ; ZVABD-NEXT:    vwaddu.vv v8, v14, v12
-; ZVABD-NEXT:    add a3, a4, a3
-; ZVABD-NEXT:    vle8.v v14, (a3)
-; ZVABD-NEXT:    add a0, a0, a2
+; ZVABD-NEXT:    vle8.v v14, (a1)
 ; ZVABD-NEXT:    vzext.vf2 v12, v14
 ; ZVABD-NEXT:    vle8.v v16, (a0)
 ; ZVABD-NEXT:    vzext.vf2 v14, v16
@@ -262,40 +264,40 @@ entry:
 define signext i32 @sadu_2block_16xi8_as_i32(ptr %a, ptr %b, i32 signext %stridea, i32 signext %strideb) {
 ; CHECK-LABEL: sadu_2block_16xi8_as_i32:
 ; CHECK:       # %bb.0: # %entry
+; CHECK-NEXT:    add a4, a0, a2
 ; CHECK-NEXT:    vsetivli zero, 16, e8, m1, ta, ma
-; CHECK-NEXT:    vle8.v v8, (a0)
-; CHECK-NEXT:    vle8.v v9, (a1)
-; CHECK-NEXT:    add a0, a0, a2
-; CHECK-NEXT:    vle8.v v10, (a0)
-; CHECK-NEXT:    add a1, a1, a3
-; CHECK-NEXT:    vmin.vv v11, v8, v9
-; CHECK-NEXT:    vle8.v v12, (a1)
+; CHECK-NEXT:    vle8.v v8, (a4)
+; CHECK-NEXT:    add a5, a1, a3
+; CHECK-NEXT:    vle8.v v9, (a5)
+; CHECK-NEXT:    vmin.vv v10, v8, v9
 ; CHECK-NEXT:    vmax.vv v8, v8, v9
-; CHECK-NEXT:    add a0, a0, a2
-; CHECK-NEXT:    vle8.v v9, (a0)
-; CHECK-NEXT:    vsub.vv v8, v8, v11
-; CHECK-NEXT:    add a1, a1, a3
-; CHECK-NEXT:    vle8.v v11, (a1)
-; CHECK-NEXT:    vmin.vv v13, v10, v12
-; CHECK-NEXT:    vmax.vv v10, v10, v12
-; CHECK-NEXT:    vmin.vv v12, v9, v11
-; CHECK-NEXT:    vmax.vv v9, v9, v11
-; CHECK-NEXT:    vsub.vv v10, v10, v13
-; CHECK-NEXT:    vsub.vv v9, v9, v12
-; CHECK-NEXT:    vwaddu.vv v12, v10, v8
+; CHECK-NEXT:    add a4, a4, a2
+; CHECK-NEXT:    vle8.v v9, (a4)
+; CHECK-NEXT:    add a2, a4, a2
+; CHECK-NEXT:    vsub.vv v8, v8, v10
+; CHECK-NEXT:    vle8.v v10, (a2)
+; CHECK-NEXT:    add a5, a5, a3
 ; CHECK-NEXT:    vsetvli zero, zero, e16, m2, ta, ma
-; CHECK-NEXT:    vzext.vf2 v14, v9
-; CHECK-NEXT:    add a0, a0, a2
-; CHECK-NEXT:    vle8.v v16, (a0)
-; CHECK-NEXT:    vwaddu.vv v8, v14, v12
-; CHECK-NEXT:    add a1, a1, a3
-; CHECK-NEXT:    vle8.v v12, (a1)
+; CHECK-NEXT:    vzext.vf2 v12, v8
+; CHECK-NEXT:    vle8.v v8, (a5)
+; CHECK-NEXT:    add a3, a5, a3
+; CHECK-NEXT:    vle8.v v11, (a3)
 ; CHECK-NEXT:    vsetvli zero, zero, e8, m1, ta, ma
-; CHECK-NEXT:    vmin.vv v13, v16, v12
-; CHECK-NEXT:    vmax.vv v12, v16, v12
-; CHECK-NEXT:    vsub.vv v14, v12, v13
+; CHECK-NEXT:    vmin.vv v14, v9, v8
+; CHECK-NEXT:    vmax.vv v8, v9, v8
+; CHECK-NEXT:    vmin.vv v9, v10, v11
+; CHECK-NEXT:    vmax.vv v10, v10, v11
+; CHECK-NEXT:    vle8.v v11, (a0)
+; CHECK-NEXT:    vsub.vv v8, v8, v14
+; CHECK-NEXT:    vle8.v v14, (a1)
+; CHECK-NEXT:    vsub.vv v9, v10, v9
+; CHECK-NEXT:    vmin.vv v10, v11, v14
+; CHECK-NEXT:    vmax.vv v11, v11, v14
+; CHECK-NEXT:    vwaddu.vv v14, v9, v8
+; CHECK-NEXT:    vsub.vv v16, v11, v10
 ; CHECK-NEXT:    vsetvli zero, zero, e16, m2, ta, ma
-; CHECK-NEXT:    vzext.vf2 v12, v14
+; CHECK-NEXT:    vwaddu.vv v8, v14, v12
+; CHECK-NEXT:    vzext.vf2 v12, v16
 ; CHECK-NEXT:    vwaddu.wv v8, v8, v12
 ; CHECK-NEXT:    vsetvli zero, zero, e32, m4, ta, ma
 ; CHECK-NEXT:    vmv.s.x v12, zero
@@ -305,29 +307,31 @@ define signext i32 @sadu_2block_16xi8_as_i32(ptr %a, ptr %b, i32 signext %stride
 ;
 ; ZVABD-LABEL: sadu_2block_16xi8_as_i32:
 ; ZVABD:       # %bb.0: # %entry
+; ZVABD-NEXT:    add a4, a0, a2
 ; ZVABD-NEXT:    vsetivli zero, 16, e8, m1, ta, ma
-; ZVABD-NEXT:    vle8.v v8, (a0)
-; ZVABD-NEXT:    vle8.v v9, (a1)
-; ZVABD-NEXT:    add a0, a0, a2
-; ZVABD-NEXT:    vle8.v v10, (a0)
+; ZVABD-NEXT:    vle8.v v8, (a4)
+; ZVABD-NEXT:    add a5, a1, a3
+; ZVABD-NEXT:    vle8.v v9, (a5)
 ; ZVABD-NEXT:    vabd.vv v8, v8, v9
-; ZVABD-NEXT:    add a0, a0, a2
-; ZVABD-NEXT:    vle8.v v9, (a0)
-; ZVABD-NEXT:    add a1, a1, a3
-; ZVABD-NEXT:    add a4, a1, a3
+; ZVABD-NEXT:    add a4, a4, a2
+; ZVABD-NEXT:    vle8.v v9, (a4)
+; ZVABD-NEXT:    add a5, a5, a3
 ; ZVABD-NEXT:    vsetvli zero, zero, e16, m2, ta, ma
 ; ZVABD-NEXT:    vzext.vf2 v12, v8
-; ZVABD-NEXT:    vle8.v v8, (a4)
-; ZVABD-NEXT:    vle8.v v11, (a1)
+; ZVABD-NEXT:    vle8.v v8, (a5)
 ; ZVABD-NEXT:    vsetvli zero, zero, e8, m1, ta, ma
 ; ZVABD-NEXT:    vabd.vv v8, v9, v8
-; ZVABD-NEXT:    vwabda.vv v12, v10, v11
+; ZVABD-NEXT:    add a2, a4, a2
+; ZVABD-NEXT:    vle8.v v9, (a2)
+; ZVABD-NEXT:    add a3, a5, a3
 ; ZVABD-NEXT:    vsetvli zero, zero, e16, m2, ta, ma
 ; ZVABD-NEXT:    vzext.vf2 v14, v8
+; ZVABD-NEXT:    vle8.v v8, (a3)
+; ZVABD-NEXT:    vsetvli zero, zero, e8, m1, ta, ma
+; ZVABD-NEXT:    vwabda.vv v14, v9, v8
+; ZVABD-NEXT:    vsetvli zero, zero, e16, m2, ta, ma
 ; ZVABD-NEXT:    vwaddu.vv v8, v14, v12
-; ZVABD-NEXT:    add a3, a4, a3
-; ZVABD-NEXT:    vle8.v v14, (a3)
-; ZVABD-NEXT:    add a0, a0, a2
+; ZVABD-NEXT:    vle8.v v14, (a1)
 ; ZVABD-NEXT:    vsext.vf2 v12, v14
 ; ZVABD-NEXT:    vle8.v v16, (a0)
 ; ZVABD-NEXT:    vsext.vf2 v14, v16
diff --git a/llvm/test/CodeGen/RISCV/rvv/fixed-vectors-zvdot4a8i.ll b/llvm/test/CodeGen/RISCV/rvv/fixed-vectors-zvdot4a8i.ll
index 92809b3cd6905..e49c9a7e6136e 100644
--- a/llvm/test/CodeGen/RISCV/rvv/fixed-vectors-zvdot4a8i.ll
+++ b/llvm/test/CodeGen/RISCV/rvv/fixed-vectors-zvdot4a8i.ll
@@ -1923,6 +1923,146 @@ entry:
   ret i32 %sum
 }
 
+; Two reductions added through a scalar accumulator chain should reassociate
+; so the dot products accumulate into one vector register and only a single
+; reduction remains.
+define i32 @vdot4au_vv_scalar_add_chain2(<16 x i8> %a0, <16 x i8> %b0, <16 x i8> %a1, <16 x i8> %b1, i32 %acc) {
+; NODOT-LABEL: vdot4au_vv_scalar_add_chain2:
+; NODOT:       # %bb.0: # %entry
+; NODOT-NEXT:    vsetivli zero, 16, e8, m1, ta, ma
+; NODOT-NEXT:    vwmulu.vv v12, v10, v11
+; NODOT-NEXT:    vwmulu.vv v14, v8, v9
+; NODOT-NEXT:    vsetvli zero, zero, e16, m2, ta, ma
+; NODOT-NEXT:    vwaddu.vv v8, v12, v14
+; NODOT-NEXT:    vsetvli zero, zero, e32, m4, ta, ma
+; NODOT-NEXT:    vmv.s.x v12, a0
+; NODOT-NEXT:    vredsum.vs v8, v8, v12
+; NODOT-NEXT:    vmv.x.s a0, v8
+; NODOT-NEXT:    ret
+;
+; DOT-LABEL: vdot4au_vv_scalar_add_chain2:
+; DOT:       # %bb.0: # %entry
+; DOT-NEXT:    vsetivli zero, 4, e32, m1, ta, ma
+; DOT-NEXT:    vmv.v.i v12, 0
+; DOT-NEXT:    vdot4au.vv v12, v10, v11
+; DOT-NEXT:    vdot4au.vv v12, v8, v9
+; DOT-NEXT:    vmv.s.x v8, a0
+; DOT-NEXT:    vredsum.vs v8, v12, v8
+; DOT-NEXT:    vmv.x.s a0, v8
+; DOT-NEXT:    ret
+entry:
+  %a0.zext = zext <16 x i8> %a0 to <16 x i32>
+  %b0.zext = zext <16 x i8> %b0 to <16 x i32>
+  %mul0 = mul <16 x i32> %a0.zext, %b0.zext
+  %r0 = tail call i32 @llvm.vector.reduce.add.v16i32(<16 x i32> %mul0)
+  %a1.zext = zext <16 x i8> %a1 to <16 x i32>
+  %b1.zext = zext <16 x i8> %b1 to <16 x i32>
+  %mul1 = mul <16 x i32> %a1.zext, %b1.zext
+  %r1 = tail call i32 @llvm.vector.reduce.add.v16i32(<16 x i32> %mul1)
+  %s0 = add i32 %r0, %acc
+  %s1 = add i32 %r1, %s0
+  ret i32 %s1
+}
+
+; Signed variant exercising the vdot4a.vv path.
+define i32 @vdot4a_vv_scalar_add_chain2(<16 x i8> %a0, <16 x i8> %b0, <16 x i8> %a1, <16 x i8> %b1, i32 %acc) {
+; NODOT-LABEL: vdot4a_vv_scalar_add_chain2:
+; NODOT:       # %bb.0: # %entry
+; NODOT-NEXT:    vsetivli zero, 16, e16, m2, ta, ma
+; NODOT-NEXT:    vsext.vf2 v16, v10
+; NODOT-NEXT:    vsext.vf2 v18, v8
+; NODOT-NEXT:    vsext.vf2 v20, v9
+; NODOT-NEXT:    vwmul.vv v12, v18, v20
+; NODOT-NEXT:    vsext.vf2 v8, v11
+; NODOT-NEXT:    vwmacc.vv v12, v16, v8
+; NODOT-NEXT:    vsetvli zero, zero, e32, m4, ta, ma
+; NODOT-NEXT:    vmv.s.x v8, a0
+; NODOT-NEXT:    vredsum.vs v8, v12, v8
+; NODOT-NEXT:    vmv.x.s a0, v8
+; NODOT-NEXT:    ret
+;
+; DOT-LABEL: vdot4a_vv_scalar_add_chain2:
+; DOT:       # %bb.0: # %entry
+; DOT-NEXT:    vsetivli zero, 4, e32, m1, ta, ma
+; DOT-NEXT:    vmv.v.i v12, 0
+; DOT-NEXT:    vdot4a.vv v12, v10, v11
+; DOT-NEXT:    vdot4a.vv v12, v8, v9
+; DOT-NEXT:    vmv.s.x v8, a0
+; DOT-NEXT:    vredsum.vs v8, v12, v8
+; DOT-NEXT:    vmv.x.s a0, v8
+; DOT-NEXT:    ret
+entry:
+  %a0.sext = sext <16 x i8> %a0 to <16 x i32>
+  %b0.sext = sext <16 x i8> %b0 to <16 x i32>
+  %mul0 = mul <16 x i32> %a0.sext, %b0.sext
+  %r0 = tail call i32 @llvm.vector.reduce.add.v16i32(<16 x i32> %mul0)
+  %a1.sext = sext <16 x i8> %a1 to <16 x i32>
+  %b1.sext = sext <16 x i8> %b1 to <16 x i32>
+  %mul1 = mul <16 x i32> %a1.sext, %b1.sext
+  %r1 = tail call i32 @llvm.vector.reduce.add.v16i32(<16 x i32> %mul1)
+  %s0 = add i32 %r0, %acc
+  %s1 = add i32 %r1, %s0
+  ret i32 %s1
+}
+
+; Four-term chain without an external accumulator (x264 SAD shape).
+define i32 @vdot4au_vv_chain4(<16 x i8> %a0, <16 x i8> %b0, <16 x i8> %a1, <16 x i8> %b1, <16 x i8> %a2, <16 x i8> %b2, <16 x i8> %a3, <16 x i8> %b3) {
+; NODOT-LABEL: vdot4au_vv_chain4:
+; NODOT:       # %bb.0: # %entry
+; NODOT-NEXT:    vsetivli zero, 16, e8, m1, ta, ma
+; NODOT-NEXT:    vwmulu.vv v16, v14, v15
+; NODOT-NEXT:    vwmulu.vv v18, v12, v13
+; NODOT-NEXT:    vsetvli zero, zero, e16, m2, ta, ma
+; NODOT-NEXT:    vwaddu.vv v12, v16, v18
+; NODOT-NEXT:    vsetvli zero, zero, e8, m1, ta, ma
+; NODOT-NEXT:    vwmulu.vv v16, v10, v11
+; NODOT-NEXT:    vsetvli zero, zero, e16, m2, ta, ma
+; NODOT-NEXT:    vwaddu.wv v12, v12, v16
+; NODOT-NEXT:    vsetvli zero, zero, e8, m1, ta, ma
+; NODOT-NEXT:    vwmulu.vv v10, v8, v9
+; NODOT-NEXT:    vsetvli zero, zero, e16, m2, ta, ma
+; NODOT-NEXT:    vwaddu.wv v12, v12, v10
+; NODOT-NEXT:    vsetvli zero, zero, e32, m4, ta, ma
+; NODOT-NEXT:    vmv.s.x v8, zero
+; NODOT-NEXT:    vredsum.vs v8, v12, v8
+; NODOT-NEXT:    vmv.x.s a0, v8
+; NODOT-NEXT:    ret
+;
+; DOT-LABEL: vdot4au_vv_chain4:
+; DOT:       # %bb.0: # %entry
+; DOT-NEXT:    vsetivli zero, 4, e32, m1, ta, ma
+; DOT-NEXT:    vmv.v.i v16, 0
+; DOT-NEXT:    vdot4au.vv v16, v14, v15
+; DOT-NEXT:    vdot4au.vv v16, v12, v13
+; DOT-NEXT:    vdot4au.vv v16, v10, v11
+; DOT-NEXT:    vdot4au.vv v16, v8, v9
+; DOT-NEXT:    vmv.s.x v8, zero
+; DOT-NEXT:    vredsum.vs v8, v16, v8
+; DOT-NEXT:    vmv.x.s a0, v8
+; DOT-NEXT:    ret
+entry:
+  %a0.zext = zext <16 x i8> %a0 to <16 x i32>
+  %b0.zext = zext <16 x i8> %b0 to <16 x i32>
+  %mul0 = mul <16 x i32> %a0.zext, %b0.zext
+  %r0 = tail call i32 @llvm.vector.reduce.add.v16i32(<16 x i32> %mul0)
+  %a1.zext = zext <16 x i8> %a1 to <16 x i32>
+  %b1.zext = zext <16 x i8> %b1 to <16 x i32>
+  %mul1 = mul <16 x i32> %a1.zext, %b1.zext
+  %r1 = tail call i32 @llvm.vector.reduce.add.v16i32(<16 x i32> %mul1)
+  %a2.zext = zext <16 x i8> %a2 to <16 x i32>
+  %b2.zext = zext <16 x i8> %b2 to <16 x i32>
+  %mul2 = mul <16 x i32> %a2.zext, %b2.zext
+  %r2 = tail call i32 @llvm.vector.reduce.add.v16i32(<16 x i32> %mul2)
+  %a3.zext = zext <16 x i8> %a3 to <16 x i32>
+  %b3.zext = zext <16 x i8> %b3 to <16 x i32>
+  %mul3 = mul <16 x i32> %a3.zext, %b3.zext
+  %r3 = tail call i32 @llvm.vector.reduce.add.v16i32(<16 x i32> %mul3)
+  %s0 = add i32 %r1, %r0
+  %s1 = add i32 %r2, %s0
+  %s2 = add i32 %r3, %s1
+  ret i32 %s2
+}
+
 ;; NOTE: These prefixes are unused and the list is autogenerated. Do not add tests below this line:
 ; DOT32: {{.*}}
 ; DOT64: {{.*}}
diff --git a/llvm/test/CodeGen/RISCV/rvv/zvdot4a8i-sdnode.ll b/llvm/test/CodeGen/RISCV/rvv/zvdot4a8i-sdnode.ll
index 6fe0e3c5ad57f..88dd44b6165f6 100644
--- a/llvm/test/CodeGen/RISCV/rvv/zvdot4a8i-sdnode.ll
+++ b/llvm/test/CodeGen/RISCV/rvv/zvdot4a8i-sdnode.ll
@@ -1059,6 +1059,105 @@ entry:
   ret <vscale x 2 x i32> %res
 }
 
+; Two reductions added through a scalar accumulator chain should reassociate
+; so the dot products accumulate into one vector register and only a single
+; reduction remains.
+define i32 @vdot4au_vv_scalar_add_chain2(<vscale x 16 x i8> %a0, <vscale x 16 x i8> %b0, <vscale x 16 x i8> %a1, <vscale x 16 x i8> %b1, i32 %acc) {
+; NODOT-LABEL: vdot4au_vv_scalar_add_chain2:
+; NODOT:       # %bb.0: # %entry
+; NODOT-NEXT:    vsetvli a1, zero, e8, m2, ta, ma
+; NODOT-NEXT:    vwmulu.vv v16, v12, v14
+; NODOT-NEXT:    vwmulu.vv v20, v8, v10
+; NODOT-NEXT:    vsetvli zero, zero, e16, m4, ta, ma
+; NODOT-NEXT:    vwaddu.vv v8, v16, v20
+; NODOT-NEXT:    vsetvli zero, zero, e32, m8, ta, ma
+; NODOT-NEXT:    vmv.s.x v16, a0
+; NODOT-NEXT:    vredsum.vs v8, v8, v16
+; NODOT-NEXT:    vmv.x.s a0, v8
+; NODOT-NEXT:    ret
+;
+; DOT-LABEL: vdot4au_vv_scalar_add_chain2:
+; DOT:       # %bb.0: # %entry
+; DOT-NEXT:    vsetvli a1, zero, e32, m2, ta, ma
+; DOT-NEXT:    vmv.v.i v16, 0
+; DOT-NEXT:    vdot4au.vv v16, v12, v14
+; DOT-NEXT:    vdot4au.vv v16, v8, v10
+; DOT-NEXT:    vmv.s.x v8, a0
+; DOT-NEXT:    vredsum.vs v8, v16, v8
+; DOT-NEXT:    vmv.x.s a0, v8
+; DOT-NEXT:    ret
+entry:
+  %a0.zext = zext <vscale x 16 x i8> %a0 to <vscale x 16 x i32>
+  %b0.zext = zext <vscale x 16 x i8> %b0 to <vscale x 16 x i32>
+  %mul0 = mul <vscale x 16 x i32> %a0.zext, %b0.zext
+  %r0 = tail call i32 @llvm.vector.reduce.add.v16i32(<vscale x 16 x i32> %mul0)
+  %a1.zext = zext <vscale x 16 x i8> %a1 to <vscale x 16 x i32>
+  %b1.zext = zext <vscale x 16...
[truncated]

``````````

</details>


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


More information about the llvm-commits mailing list