[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