[llvm] [AArch64] Keep v16i8 -> v2i64 partial_reduce fixed-length for VL > 128 (PR #204938)
Matthew Blewitt via llvm-commits
llvm-commits at lists.llvm.org
Mon Jun 22 03:00:32 PDT 2026
https://github.com/mble updated https://github.com/llvm/llvm-project/pull/204938
>From 113956bd87fc9525c5bce8d1b8d18decaae36015 Mon Sep 17 00:00:00 2001
From: Matt Blewitt <mble at planetscale.com>
Date: Sat, 20 Jun 2026 10:28:18 -0700
Subject: [PATCH] [AArch64] Fix v16i8 -> v2i64 partial_reduce for VL > 128
A fixed-length llvm.vector.partial.reduce.add with a <2 x i64>
accumulator and i8 inputs was lowered, on any +sve target, through a
scalable partial reduction whose nxv4i32 dot was folded into nxv2i64 and
then truncated back to the low 128 bits via convertFromScalableVector.
The fold (UADDW{B,T}/SADDW{B,T} or splitting the scalable dot) spreads
the accumulated sums across all VL/64 lanes, so the trailing extract
keeps only the low two and silently drops the rest for any runtime
vector length greater than 128 bits. On a 256-bit machine (e.g. Neoverse
V1) exactly half the result is lost.
The dot product itself is VL-independent: the fixed-length input is
zero-padded into the scalable container, so its meaningful sums occupy
the low i32 lanes regardless of VL. Fix the fold to convert the dot back
to a fixed-length v4i32 before splitting, so the split, widen and
accumulate all happen in fixed-length and no high lanes are dropped.
Restrict the UADDW{B,T}/SADDW{B,T} path to genuinely scalable results,
which is where that scalable fold is correct.
This also corrects the streaming-SVE/SME path, which previously took the
same buggy scalable fold.
This is the i64 sibling of the v16i8 -> v2i32 case fixed in #177119
(issue #176954); the fixed-length support was introduced in #142032.
---
.../Target/AArch64/AArch64ISelLowering.cpp | 48 ++++--
.../sve-fixed-length-partial-reduce.ll | 149 ++++++++++++------
2 files changed, 134 insertions(+), 63 deletions(-)
diff --git a/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp b/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
index 67ef911117eff..fe38ef553645f 100644
--- a/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
+++ b/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
@@ -33449,6 +33449,7 @@ AArch64TargetLowering::LowerPARTIAL_REDUCE_MLA(SDValue Op,
ResultVT.isFixedLengthVector() &&
useSVEForFixedLengthVectorVT(ResultVT, /*OverrideNEON=*/true);
+ SDValue OrigAcc = Acc;
if (ConvertToScalable) {
ResultVT = getContainerForFixedLengthVector(DAG, ResultVT);
OpVT = getContainerForFixedLengthVector(DAG, LHS.getValueType());
@@ -33468,29 +33469,42 @@ AArch64TargetLowering::LowerPARTIAL_REDUCE_MLA(SDValue Op,
SDValue DotNode = DAG.getNode(Op.getOpcode(), DL, DotVT,
DAG.getConstant(0, DL, DotVT), LHS, RHS);
- SDValue Res;
bool IsUnsigned = Op.getOpcode() == ISD::PARTIAL_REDUCE_UMLA;
- if (Subtarget->hasSVE2() || Subtarget->isStreamingSVEAvailable()) {
+
+ // UADDW{B,T}/SADDW{B,T} fold the dot in the scalable domain, spreading the
+ // sums across all VL/64 lanes. That is only valid for a genuinely scalable
+ // result; a fixed-length result must convert from the scalable container
+ // before splitting (below), else the trailing extract drops the high lanes
+ // for any VL > the fixed width.
+ if (OrigResultVT.isScalableVector() &&
+ (Subtarget->hasSVE2() || Subtarget->isStreamingSVEAvailable())) {
unsigned LoOpcode = IsUnsigned ? AArch64ISD::UADDWB : AArch64ISD::SADDWB;
unsigned HiOpcode = IsUnsigned ? AArch64ISD::UADDWT : AArch64ISD::SADDWT;
SDValue Lo = DAG.getNode(LoOpcode, DL, ResultVT, Acc, DotNode);
- Res = DAG.getNode(HiOpcode, DL, ResultVT, Lo, DotNode);
- } else {
- // Fold (nx)v4i32 into (nx)v2i64
- auto [DotNodeLo, DotNodeHi] = DAG.SplitVector(DotNode, DL);
- if (IsUnsigned) {
- DotNodeLo = DAG.getZExtOrTrunc(DotNodeLo, DL, ResultVT);
- DotNodeHi = DAG.getZExtOrTrunc(DotNodeHi, DL, ResultVT);
- } else {
- DotNodeLo = DAG.getSExtOrTrunc(DotNodeLo, DL, ResultVT);
- DotNodeHi = DAG.getSExtOrTrunc(DotNodeHi, DL, ResultVT);
- }
- auto Lo = DAG.getNode(ISD::ADD, DL, ResultVT, Acc, DotNodeLo);
- Res = DAG.getNode(ISD::ADD, DL, ResultVT, Lo, DotNodeHi);
+ return DAG.getNode(HiOpcode, DL, ResultVT, Lo, DotNode);
}
- return ConvertToScalable ? convertFromScalableVector(DAG, OrigResultVT, Res)
- : Res;
+ // Fold the (nx)v4i32 dot into the (nx)v2i64 result. For a fixed-length
+ // result, convert from the scalable container before splitting: the sums sit
+ // in the low i32 lanes regardless of VL, so splitting after the extract would
+ // drop a 128-bit result's high lanes for VL > 128. See PR #177119 / issue
+ // #176954.
+ SDValue FoldDot = DotNode;
+ if (ConvertToScalable) {
+ EVT FixedDotVT = EVT::getVectorVT(*DAG.getContext(), MVT::i32,
+ OrigResultVT.getVectorNumElements() * 2);
+ FoldDot = convertFromScalableVector(DAG, FixedDotVT, DotNode);
+ }
+ auto [DotNodeLo, DotNodeHi] = DAG.SplitVector(FoldDot, DL);
+ if (IsUnsigned) {
+ DotNodeLo = DAG.getZExtOrTrunc(DotNodeLo, DL, OrigResultVT);
+ DotNodeHi = DAG.getZExtOrTrunc(DotNodeHi, DL, OrigResultVT);
+ } else {
+ DotNodeLo = DAG.getSExtOrTrunc(DotNodeLo, DL, OrigResultVT);
+ DotNodeHi = DAG.getSExtOrTrunc(DotNodeHi, DL, OrigResultVT);
+ }
+ SDValue Lo = DAG.getNode(ISD::ADD, DL, OrigResultVT, OrigAcc, DotNodeLo);
+ return DAG.getNode(ISD::ADD, DL, OrigResultVT, Lo, DotNodeHi);
}
SDValue
diff --git a/llvm/test/CodeGen/AArch64/sve-fixed-length-partial-reduce.ll b/llvm/test/CodeGen/AArch64/sve-fixed-length-partial-reduce.ll
index ae7aa9b35f62a..4a3c213e60cb3 100644
--- a/llvm/test/CodeGen/AArch64/sve-fixed-length-partial-reduce.ll
+++ b/llvm/test/CodeGen/AArch64/sve-fixed-length-partial-reduce.ll
@@ -461,12 +461,9 @@ define <2 x i64> @four_way_i8_i64_vl128_usdot(ptr %accptr, ptr %uptr, ptr %sptr)
; SVE-NEXT: ldr q1, [x1]
; SVE-NEXT: ldr q2, [x2]
; SVE-NEXT: usdot z0.s, z1.b, z2.b
-; SVE-NEXT: ldr q2, [x0]
-; SVE-NEXT: sunpklo z1.d, z0.s
-; SVE-NEXT: sunpkhi z0.d, z0.s
-; SVE-NEXT: add z1.d, z2.d, z1.d
-; SVE-NEXT: add z0.d, z1.d, z0.d
-; SVE-NEXT: // kill: def $q0 killed $q0 killed $z0
+; SVE-NEXT: ldr q1, [x0]
+; SVE-NEXT: saddw v1.2d, v1.2d, v0.2s
+; SVE-NEXT: saddw2 v0.2d, v1.2d, v0.4s
; SVE-NEXT: ret
;
; SME-LABEL: four_way_i8_i64_vl128_usdot:
@@ -475,9 +472,12 @@ define <2 x i64> @four_way_i8_i64_vl128_usdot(ptr %accptr, ptr %uptr, ptr %sptr)
; SME-NEXT: ldr q1, [x1]
; SME-NEXT: ldr q2, [x2]
; SME-NEXT: usdot z0.s, z1.b, z2.b
-; SME-NEXT: ldr q1, [x0]
-; SME-NEXT: saddwb z1.d, z1.d, z0.s
-; SME-NEXT: saddwt z0.d, z1.d, z0.s
+; SME-NEXT: ldr q2, [x0]
+; SME-NEXT: sunpklo z1.d, z0.s
+; SME-NEXT: ext z0.b, z0.b, z0.b, #8
+; SME-NEXT: sunpklo z0.d, z0.s
+; SME-NEXT: add z1.d, z2.d, z1.d
+; SME-NEXT: add z0.d, z1.d, z0.d
; SME-NEXT: ret
%acc = load <2 x i64>, ptr %accptr
%u = load <16 x i8>, ptr %uptr
@@ -831,12 +831,9 @@ define <2 x i64> @eight_way_i8_i64_vl128(ptr %accptr, ptr %uptr, ptr %sptr) {
; SVE-NEXT: ldr q1, [x1]
; SVE-NEXT: ldr q2, [x2]
; SVE-NEXT: udot z0.s, z2.b, z1.b
-; SVE-NEXT: ldr q2, [x0]
-; SVE-NEXT: uunpklo z1.d, z0.s
-; SVE-NEXT: uunpkhi z0.d, z0.s
-; SVE-NEXT: add z1.d, z2.d, z1.d
-; SVE-NEXT: add z0.d, z1.d, z0.d
-; SVE-NEXT: // kill: def $q0 killed $q0 killed $z0
+; SVE-NEXT: ldr q1, [x0]
+; SVE-NEXT: uaddw v1.2d, v1.2d, v0.2s
+; SVE-NEXT: uaddw2 v0.2d, v1.2d, v0.4s
; SVE-NEXT: ret
;
; SME-LABEL: eight_way_i8_i64_vl128:
@@ -845,9 +842,64 @@ define <2 x i64> @eight_way_i8_i64_vl128(ptr %accptr, ptr %uptr, ptr %sptr) {
; SME-NEXT: ldr q1, [x1]
; SME-NEXT: ldr q2, [x2]
; SME-NEXT: udot z0.s, z2.b, z1.b
-; SME-NEXT: ldr q1, [x0]
-; SME-NEXT: uaddwb z1.d, z1.d, z0.s
-; SME-NEXT: uaddwt z0.d, z1.d, z0.s
+; SME-NEXT: ldr q2, [x0]
+; SME-NEXT: uunpklo z1.d, z0.s
+; SME-NEXT: ext z0.b, z0.b, z0.b, #8
+; SME-NEXT: uunpklo z0.d, z0.s
+; SME-NEXT: add z1.d, z2.d, z1.d
+; SME-NEXT: add z0.d, z1.d, z0.d
+; SME-NEXT: ret
+ %acc = load <2 x i64>, ptr %accptr
+ %u = load <16 x i8>, ptr %uptr
+ %s = load <16 x i8>, ptr %sptr
+ %u.wide = zext <16 x i8> %u to <16 x i64>
+ %s.wide = zext <16 x i8> %s to <16 x i64>
+ %mult = mul nuw nsw <16 x i64> %s.wide, %u.wide
+ %partial.reduce = tail call <2 x i64> @llvm.vector.partial.reduce.add(<2 x i64> %acc, <16 x i64> %mult)
+ ret <2 x i64> %partial.reduce
+}
+
+; Regression test for the off-diagonal case: a 128-bit (<2 x i64>) result at
+; VL=256, where the fixed result width is smaller than the SVE vector length.
+; Before the fix the SVE run lowered via a scalable partial reduction whose
+; sums were spread across 4 d-lanes and then truncated to the low 2, dropping
+; half the count at runtime. The fix keeps it on the VL-independent NEON path.
+define <2 x i64> @eight_way_i8_i64_vl256(ptr %accptr, ptr %uptr, ptr %sptr) vscale_range(2,2) {
+;
+; NEON-LABEL: eight_way_i8_i64_vl256:
+; NEON: // %bb.0:
+; NEON-NEXT: movi v0.2d, #0000000000000000
+; NEON-NEXT: ldr q1, [x1]
+; NEON-NEXT: ldr q2, [x2]
+; NEON-NEXT: udot v0.4s, v2.16b, v1.16b
+; NEON-NEXT: ldr q1, [x0]
+; NEON-NEXT: uaddw v1.2d, v1.2d, v0.2s
+; NEON-NEXT: uaddw2 v0.2d, v1.2d, v0.4s
+; NEON-NEXT: ret
+;
+; SVE-LABEL: eight_way_i8_i64_vl256:
+; SVE: // %bb.0:
+; SVE-NEXT: movi v0.2d, #0000000000000000
+; SVE-NEXT: ldr q1, [x1]
+; SVE-NEXT: ldr q2, [x2]
+; SVE-NEXT: udot z0.s, z2.b, z1.b
+; SVE-NEXT: ldr q1, [x0]
+; SVE-NEXT: uaddw v1.2d, v1.2d, v0.2s
+; SVE-NEXT: uaddw2 v0.2d, v1.2d, v0.4s
+; SVE-NEXT: ret
+;
+; SME-LABEL: eight_way_i8_i64_vl256:
+; SME: // %bb.0:
+; SME-NEXT: mov z0.s, #0 // =0x0
+; SME-NEXT: ldr q1, [x1]
+; SME-NEXT: ldr q2, [x2]
+; SME-NEXT: udot z0.s, z2.b, z1.b
+; SME-NEXT: ldr q2, [x0]
+; SME-NEXT: uunpklo z1.d, z0.s
+; SME-NEXT: ext z0.b, z0.b, z0.b, #8
+; SME-NEXT: uunpklo z0.d, z0.s
+; SME-NEXT: add z1.d, z2.d, z1.d
+; SME-NEXT: add z0.d, z1.d, z0.d
; SME-NEXT: ret
%acc = load <2 x i64>, ptr %accptr
%u = load <16 x i8>, ptr %uptr
@@ -878,38 +930,38 @@ define <4 x i64> @four_way_i8_i64_vl128_double_width(ptr %accptr, ptr %uptr, ptr
;
; SVE-LABEL: four_way_i8_i64_vl128_double_width:
; SVE: // %bb.0:
-; SVE-NEXT: movi v0.2d, #0000000000000000
; SVE-NEXT: movi v1.2d, #0000000000000000
+; SVE-NEXT: movi v0.2d, #0000000000000000
; SVE-NEXT: ldp q3, q2, [x1]
; SVE-NEXT: ldp q5, q4, [x2]
-; SVE-NEXT: udot z1.s, z5.b, z3.b
-; SVE-NEXT: udot z0.s, z4.b, z2.b
-; SVE-NEXT: ldp q5, q4, [x0]
-; SVE-NEXT: uunpklo z2.d, z1.s
-; SVE-NEXT: uunpklo z3.d, z0.s
-; SVE-NEXT: uunpkhi z1.d, z1.s
-; SVE-NEXT: uunpkhi z6.d, z0.s
-; SVE-NEXT: add z0.d, z5.d, z2.d
-; SVE-NEXT: add z2.d, z4.d, z3.d
-; SVE-NEXT: add z0.d, z0.d, z1.d
-; SVE-NEXT: add z1.d, z2.d, z6.d
-; SVE-NEXT: // kill: def $q0 killed $q0 killed $z0
-; SVE-NEXT: // kill: def $q1 killed $q1 killed $z1
+; SVE-NEXT: udot z0.s, z5.b, z3.b
+; SVE-NEXT: udot z1.s, z4.b, z2.b
+; SVE-NEXT: ldp q3, q2, [x0]
+; SVE-NEXT: uaddw v3.2d, v3.2d, v0.2s
+; SVE-NEXT: uaddw v2.2d, v2.2d, v1.2s
+; SVE-NEXT: uaddw2 v0.2d, v3.2d, v0.4s
+; SVE-NEXT: uaddw2 v1.2d, v2.2d, v1.4s
; SVE-NEXT: ret
;
; SME-LABEL: four_way_i8_i64_vl128_double_width:
; SME: // %bb.0:
-; SME-NEXT: mov z1.s, #0 // =0x0
; SME-NEXT: mov z0.s, #0 // =0x0
-; SME-NEXT: ldp q3, q2, [x1]
-; SME-NEXT: ldp q5, q4, [x2]
-; SME-NEXT: udot z0.s, z5.b, z3.b
-; SME-NEXT: udot z1.s, z4.b, z2.b
-; SME-NEXT: ldp q3, q2, [x0]
-; SME-NEXT: uaddwb z3.d, z3.d, z0.s
-; SME-NEXT: uaddwb z2.d, z2.d, z1.s
-; SME-NEXT: uaddwt z0.d, z3.d, z0.s
-; SME-NEXT: uaddwt z1.d, z2.d, z1.s
+; SME-NEXT: mov z1.s, #0 // =0x0
+; SME-NEXT: ldp q2, q5, [x2]
+; SME-NEXT: ldp q3, q4, [x1]
+; SME-NEXT: udot z0.s, z2.b, z3.b
+; SME-NEXT: udot z1.s, z5.b, z4.b
+; SME-NEXT: ldp q5, q4, [x0]
+; SME-NEXT: uunpklo z2.d, z0.s
+; SME-NEXT: ext z0.b, z0.b, z0.b, #8
+; SME-NEXT: uunpklo z3.d, z1.s
+; SME-NEXT: ext z1.b, z1.b, z1.b, #8
+; SME-NEXT: uunpklo z0.d, z0.s
+; SME-NEXT: uunpklo z1.d, z1.s
+; SME-NEXT: add z2.d, z5.d, z2.d
+; SME-NEXT: add z3.d, z4.d, z3.d
+; SME-NEXT: add z0.d, z2.d, z0.d
+; SME-NEXT: add z1.d, z3.d, z1.d
; SME-NEXT: ret
%acc = load <4 x i64>, ptr %accptr
%u = load <32 x i8>, ptr %uptr
@@ -945,7 +997,8 @@ define <4 x i64> @four_way_i8_i64_vl256(ptr %accptr, ptr %uptr, ptr %sptr) vscal
; SVE-NEXT: udot z0.s, z2.b, z1.b
; SVE-NEXT: ldr z2, [x0]
; SVE-NEXT: uunpklo z1.d, z0.s
-; SVE-NEXT: uunpkhi z0.d, z0.s
+; SVE-NEXT: ext z0.b, z0.b, z0.b, #16
+; SVE-NEXT: uunpklo z0.d, z0.s
; SVE-NEXT: add z1.d, z2.d, z1.d
; SVE-NEXT: add z0.d, z1.d, z0.d
; SVE-NEXT: movprfx z1, z0
@@ -960,9 +1013,13 @@ define <4 x i64> @four_way_i8_i64_vl256(ptr %accptr, ptr %uptr, ptr %sptr) vscal
; SME-NEXT: ldr z1, [x2]
; SME-NEXT: mov z2.s, #0 // =0x0
; SME-NEXT: udot z2.s, z1.b, z0.b
-; SME-NEXT: ldr z0, [x0]
-; SME-NEXT: uaddwb z0.d, z0.d, z2.s
-; SME-NEXT: uaddwt z0.d, z0.d, z2.s
+; SME-NEXT: uunpklo z0.d, z2.s
+; SME-NEXT: movprfx z1, z2
+; SME-NEXT: ext z1.b, z1.b, z2.b, #16
+; SME-NEXT: ldr z2, [x0]
+; SME-NEXT: uunpklo z1.d, z1.s
+; SME-NEXT: add z0.d, z2.d, z0.d
+; SME-NEXT: add z0.d, z0.d, z1.d
; SME-NEXT: movprfx z1, z0
; SME-NEXT: ext z1.b, z1.b, z0.b, #16
; SME-NEXT: ret
More information about the llvm-commits
mailing list