[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