https://github.com/dyung updated https://github.com/llvm/llvm-project/pull/215512
>From 211d021e1f7d2f71448e26776e283790808139e0 Mon Sep 17 00:00:00 2001 From: Adam Scott <[email protected]> Date: Tue, 11 Aug 2026 06:16:27 -0400 Subject: [PATCH] [DAGCombiner][AArch64] Fix the multiplier when folding a partial reduction of a mask (#215187) foldPartialReduceAdd synthesises the reduction's multiplier at the operand's type. At i1 a splat of 1 is all ones, which sign extends to -1, so the signed forms compute (-1) x (-1) = +1 per lane and sum to +n where sum(sext(mask)) must be -n. ```llvm %cmp = icmp eq <16 x i8> %a, %b %sext = sext <16 x i1> %cmp to <16 x i32> %r = call <4 x i32> @llvm.vector.partial.reduce.add(<4 x i32> %acc, <16 x i32> %sext) ``` -mattr=+neon sums to -n: ``` cmeq v1.16b, v1.16b, v2.16b sshll v2.8h, v1.8b, #0 sshll2 v1.8h, v1.16b, #0 saddw v0.4s, v0.4s, v2.4h saddw2 v0.4s, v0.4s, v2.8h saddw v0.4s, v0.4s, v1.4h saddw2 v0.4s, v0.4s, v1.8h ``` -mattr=+neon,+dotprod sums to +n: ``` movi v3.2d, #0xffffffffffffffff cmeq v1.16b, v1.16b, v2.16b sdot v0.4s, v1.16b, v3.16b ``` This patch extends i1 masks to the promoted type before the multiplier is built. That gives `movi v3.16b, #1` and the dot product path agrees with the expansion. Part of #204897. (cherry picked from commit 9beddc4f2724c9b263b31447b7c66075dfcc8ff1) This cherry-pick also contains an additional test fix for the release branch. --- llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp | 12 +- .../neon-partial-reduce-dot-product.ll | 216 ++++++++++++++++++ 2 files changed, 227 insertions(+), 1 deletion(-) diff --git a/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp b/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp index 3ae55f69013d6..5fd1b10d97f35 100644 --- a/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp +++ b/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp @@ -14113,11 +14113,21 @@ SDValue DAGCombiner::foldPartialReduceAdd(SDNode *N) { SDValue UnextOp1 = Op1.getOperand(0); EVT UnextOp1VT = UnextOp1.getValueType(); auto *Context = DAG.getContext(); + EVT PromOp1VT = TLI.getTypeToTransformTo(*Context, UnextOp1VT); if (!TLI.isPartialReduceMLALegalOrCustom( NewOpcode, TLI.getTypeToTransformTo(*Context, N->getValueType(0)), - TLI.getTypeToTransformTo(*Context, UnextOp1VT))) + PromOp1VT)) return SDValue(); + // The multiplier below is built at the operand type, where a splat of 1 in i1 + // sign extends to -1. Extend i1 masks to the promoted type first. + if (Op1IsSigned && UnextOp1VT.getVectorElementType() == MVT::i1) { + if (PromOp1VT == UnextOp1VT) + return SDValue(); + UnextOp1VT = PromOp1VT; + UnextOp1 = DAG.getNode(ISD::SIGN_EXTEND, DL, UnextOp1VT, UnextOp1); + } + SDValue Constant = N->getOpcode() == ISD::PARTIAL_REDUCE_FMLA ? DAG.getConstantFP(1, DL, UnextOp1VT) : DAG.getConstant(1, DL, UnextOp1VT); diff --git a/llvm/test/CodeGen/AArch64/neon-partial-reduce-dot-product.ll b/llvm/test/CodeGen/AArch64/neon-partial-reduce-dot-product.ll index b5801f8f48057..a3c780cad1147 100644 --- a/llvm/test/CodeGen/AArch64/neon-partial-reduce-dot-product.ll +++ b/llvm/test/CodeGen/AArch64/neon-partial-reduce-dot-product.ll @@ -1650,3 +1650,219 @@ entry: %partial.reduce = tail call <2 x i32> @llvm.vector.partial.reduce.add.v2i32.v16i32(<2 x i32> %acc, <16 x i32> %mult) ret <2 x i32> %partial.reduce } + +define <4 x i32> @partial_reduce_zext_cmp_i8tov4i32(<4 x i32> %acc, <16 x i8> %a, <16 x i8> %b) { +; CHECK-NODOT-LABEL: partial_reduce_zext_cmp_i8tov4i32: +; CHECK-NODOT: // %bb.0: +; CHECK-NODOT-NEXT: cmeq v1.16b, v1.16b, v2.16b +; CHECK-NODOT-NEXT: movi v3.4s, #1 +; CHECK-NODOT-NEXT: ushll2 v2.8h, v1.16b, #0 +; CHECK-NODOT-NEXT: ushll v1.8h, v1.8b, #0 +; CHECK-NODOT-NEXT: ushll v4.4s, v2.4h, #0 +; CHECK-NODOT-NEXT: ushll2 v5.4s, v1.8h, #0 +; CHECK-NODOT-NEXT: ushll v1.4s, v1.4h, #0 +; CHECK-NODOT-NEXT: ushll2 v2.4s, v2.8h, #0 +; CHECK-NODOT-NEXT: and v4.16b, v4.16b, v3.16b +; CHECK-NODOT-NEXT: and v5.16b, v5.16b, v3.16b +; CHECK-NODOT-NEXT: and v1.16b, v1.16b, v3.16b +; CHECK-NODOT-NEXT: and v2.16b, v2.16b, v3.16b +; CHECK-NODOT-NEXT: add v0.4s, v0.4s, v1.4s +; CHECK-NODOT-NEXT: add v1.4s, v5.4s, v4.4s +; CHECK-NODOT-NEXT: add v0.4s, v0.4s, v1.4s +; CHECK-NODOT-NEXT: add v0.4s, v0.4s, v2.4s +; CHECK-NODOT-NEXT: ret +; +; CHECK-DOT-LABEL: partial_reduce_zext_cmp_i8tov4i32: +; CHECK-DOT: // %bb.0: +; CHECK-DOT-NEXT: movi v3.16b, #1 +; CHECK-DOT-NEXT: cmeq v1.16b, v1.16b, v2.16b +; CHECK-DOT-NEXT: and v1.16b, v1.16b, v3.16b +; CHECK-DOT-NEXT: udot v0.4s, v1.16b, v3.16b +; CHECK-DOT-NEXT: ret +; +; CHECK-DOT-I8MM-LABEL: partial_reduce_zext_cmp_i8tov4i32: +; CHECK-DOT-I8MM: // %bb.0: +; CHECK-DOT-I8MM-NEXT: movi v3.16b, #1 +; CHECK-DOT-I8MM-NEXT: cmeq v1.16b, v1.16b, v2.16b +; CHECK-DOT-I8MM-NEXT: and v1.16b, v1.16b, v3.16b +; CHECK-DOT-I8MM-NEXT: udot v0.4s, v1.16b, v3.16b +; CHECK-DOT-I8MM-NEXT: ret + %cmp = icmp eq <16 x i8> %a, %b + %ext = zext <16 x i1> %cmp to <16 x i32> + %partial.reduce = tail call <4 x i32> @llvm.vector.partial.reduce.add(<4 x i32> %acc, <16 x i32> %ext) + ret <4 x i32> %partial.reduce +} + +define <4 x i32> @partial_reduce_sext_cmp_i8tov4i32(<4 x i32> %acc, <16 x i8> %a, <16 x i8> %b) { +; CHECK-NODOT-LABEL: partial_reduce_sext_cmp_i8tov4i32: +; CHECK-NODOT: // %bb.0: +; CHECK-NODOT-NEXT: cmeq v1.16b, v1.16b, v2.16b +; CHECK-NODOT-NEXT: sshll v2.8h, v1.8b, #0 +; CHECK-NODOT-NEXT: sshll2 v1.8h, v1.16b, #0 +; CHECK-NODOT-NEXT: saddw v0.4s, v0.4s, v2.4h +; CHECK-NODOT-NEXT: saddw2 v0.4s, v0.4s, v2.8h +; CHECK-NODOT-NEXT: saddw v0.4s, v0.4s, v1.4h +; CHECK-NODOT-NEXT: saddw2 v0.4s, v0.4s, v1.8h +; CHECK-NODOT-NEXT: ret +; +; CHECK-DOT-LABEL: partial_reduce_sext_cmp_i8tov4i32: +; CHECK-DOT: // %bb.0: +; CHECK-DOT-NEXT: movi v3.16b, #1 +; CHECK-DOT-NEXT: cmeq v1.16b, v1.16b, v2.16b +; CHECK-DOT-NEXT: sdot v0.4s, v1.16b, v3.16b +; CHECK-DOT-NEXT: ret +; +; CHECK-DOT-I8MM-LABEL: partial_reduce_sext_cmp_i8tov4i32: +; CHECK-DOT-I8MM: // %bb.0: +; CHECK-DOT-I8MM-NEXT: movi v3.16b, #1 +; CHECK-DOT-I8MM-NEXT: cmeq v1.16b, v1.16b, v2.16b +; CHECK-DOT-I8MM-NEXT: sdot v0.4s, v1.16b, v3.16b +; CHECK-DOT-I8MM-NEXT: ret + %cmp = icmp eq <16 x i8> %a, %b + %ext = sext <16 x i1> %cmp to <16 x i32> + %partial.reduce = tail call <4 x i32> @llvm.vector.partial.reduce.add(<4 x i32> %acc, <16 x i32> %ext) + ret <4 x i32> %partial.reduce +} + +define <2 x i64> @partial_reduce_zext_cmp_i8tov2i64(<2 x i64> %acc, <16 x i8> %a, <16 x i8> %b) { +; CHECK-NODOT-LABEL: partial_reduce_zext_cmp_i8tov2i64: +; CHECK-NODOT: // %bb.0: +; CHECK-NODOT-NEXT: cmeq v1.16b, v1.16b, v2.16b +; CHECK-NODOT-NEXT: mov w8, #1 // =0x1 +; CHECK-NODOT-NEXT: dup v5.2d, x8 +; CHECK-NODOT-NEXT: ushll2 v2.8h, v1.16b, #0 +; CHECK-NODOT-NEXT: ushll v1.8h, v1.8b, #0 +; CHECK-NODOT-NEXT: ushll v3.4s, v2.4h, #0 +; CHECK-NODOT-NEXT: ushll2 v4.4s, v1.8h, #0 +; CHECK-NODOT-NEXT: ushll v1.4s, v1.4h, #0 +; CHECK-NODOT-NEXT: ushll2 v2.4s, v2.8h, #0 +; CHECK-NODOT-NEXT: ushll v6.2d, v3.2s, #0 +; CHECK-NODOT-NEXT: ushll2 v7.2d, v4.4s, #0 +; CHECK-NODOT-NEXT: ushll v4.2d, v4.2s, #0 +; CHECK-NODOT-NEXT: ushll2 v16.2d, v1.4s, #0 +; CHECK-NODOT-NEXT: ushll v1.2d, v1.2s, #0 +; CHECK-NODOT-NEXT: ushll2 v3.2d, v3.4s, #0 +; CHECK-NODOT-NEXT: ushll2 v17.2d, v2.4s, #0 +; CHECK-NODOT-NEXT: ushll v2.2d, v2.2s, #0 +; CHECK-NODOT-NEXT: and v6.16b, v6.16b, v5.16b +; CHECK-NODOT-NEXT: and v7.16b, v7.16b, v5.16b +; CHECK-NODOT-NEXT: and v4.16b, v4.16b, v5.16b +; CHECK-NODOT-NEXT: and v16.16b, v16.16b, v5.16b +; CHECK-NODOT-NEXT: and v1.16b, v1.16b, v5.16b +; CHECK-NODOT-NEXT: and v3.16b, v3.16b, v5.16b +; CHECK-NODOT-NEXT: and v2.16b, v2.16b, v5.16b +; CHECK-NODOT-NEXT: add v0.2d, v0.2d, v1.2d +; CHECK-NODOT-NEXT: add v1.2d, v16.2d, v4.2d +; CHECK-NODOT-NEXT: add v4.2d, v7.2d, v6.2d +; CHECK-NODOT-NEXT: and v6.16b, v17.16b, v5.16b +; CHECK-NODOT-NEXT: add v0.2d, v0.2d, v1.2d +; CHECK-NODOT-NEXT: add v1.2d, v4.2d, v3.2d +; CHECK-NODOT-NEXT: add v0.2d, v0.2d, v1.2d +; CHECK-NODOT-NEXT: add v1.2d, v2.2d, v6.2d +; CHECK-NODOT-NEXT: add v0.2d, v0.2d, v1.2d +; CHECK-NODOT-NEXT: ret +; +; CHECK-DOT-LABEL: partial_reduce_zext_cmp_i8tov2i64: +; CHECK-DOT: // %bb.0: +; CHECK-DOT-NEXT: movi v3.16b, #1 +; CHECK-DOT-NEXT: cmeq v1.16b, v1.16b, v2.16b +; CHECK-DOT-NEXT: movi v2.2d, #0000000000000000 +; CHECK-DOT-NEXT: and v1.16b, v1.16b, v3.16b +; CHECK-DOT-NEXT: udot v2.4s, v1.16b, v3.16b +; CHECK-DOT-NEXT: uaddw v0.2d, v0.2d, v2.2s +; CHECK-DOT-NEXT: uaddw2 v0.2d, v0.2d, v2.4s +; CHECK-DOT-NEXT: ret +; +; CHECK-DOT-I8MM-LABEL: partial_reduce_zext_cmp_i8tov2i64: +; CHECK-DOT-I8MM: // %bb.0: +; CHECK-DOT-I8MM-NEXT: movi v3.16b, #1 +; CHECK-DOT-I8MM-NEXT: cmeq v1.16b, v1.16b, v2.16b +; CHECK-DOT-I8MM-NEXT: movi v2.2d, #0000000000000000 +; CHECK-DOT-I8MM-NEXT: and v1.16b, v1.16b, v3.16b +; CHECK-DOT-I8MM-NEXT: udot v2.4s, v1.16b, v3.16b +; CHECK-DOT-I8MM-NEXT: uaddw v0.2d, v0.2d, v2.2s +; CHECK-DOT-I8MM-NEXT: uaddw2 v0.2d, v0.2d, v2.4s +; CHECK-DOT-I8MM-NEXT: ret + %cmp = icmp eq <16 x i8> %a, %b + %ext = zext <16 x i1> %cmp to <16 x i64> + %partial.reduce = tail call <2 x i64> @llvm.vector.partial.reduce.add(<2 x i64> %acc, <16 x i64> %ext) + ret <2 x i64> %partial.reduce +} + +define <2 x i64> @partial_reduce_sext_cmp_i8tov2i64(<2 x i64> %acc, <16 x i8> %a, <16 x i8> %b) { +; CHECK-NODOT-LABEL: partial_reduce_sext_cmp_i8tov2i64: +; CHECK-NODOT: // %bb.0: +; CHECK-NODOT-NEXT: cmeq v1.16b, v1.16b, v2.16b +; CHECK-NODOT-NEXT: sshll v2.8h, v1.8b, #0 +; CHECK-NODOT-NEXT: sshll2 v1.8h, v1.16b, #0 +; CHECK-NODOT-NEXT: sshll v3.4s, v2.4h, #0 +; CHECK-NODOT-NEXT: sshll2 v2.4s, v2.8h, #0 +; CHECK-NODOT-NEXT: saddw v0.2d, v0.2d, v3.2s +; CHECK-NODOT-NEXT: saddw2 v0.2d, v0.2d, v3.4s +; CHECK-NODOT-NEXT: sshll v3.4s, v1.4h, #0 +; CHECK-NODOT-NEXT: sshll2 v1.4s, v1.8h, #0 +; CHECK-NODOT-NEXT: saddw v0.2d, v0.2d, v2.2s +; CHECK-NODOT-NEXT: saddw2 v0.2d, v0.2d, v2.4s +; CHECK-NODOT-NEXT: saddw v0.2d, v0.2d, v3.2s +; CHECK-NODOT-NEXT: saddw2 v0.2d, v0.2d, v3.4s +; CHECK-NODOT-NEXT: saddw v0.2d, v0.2d, v1.2s +; CHECK-NODOT-NEXT: saddw2 v0.2d, v0.2d, v1.4s +; CHECK-NODOT-NEXT: ret +; +; CHECK-DOT-LABEL: partial_reduce_sext_cmp_i8tov2i64: +; CHECK-DOT: // %bb.0: +; CHECK-DOT-NEXT: movi v3.16b, #1 +; CHECK-DOT-NEXT: movi v4.2d, #0000000000000000 +; CHECK-DOT-NEXT: cmeq v1.16b, v1.16b, v2.16b +; CHECK-DOT-NEXT: sdot v4.4s, v1.16b, v3.16b +; CHECK-DOT-NEXT: saddw v0.2d, v0.2d, v4.2s +; CHECK-DOT-NEXT: saddw2 v0.2d, v0.2d, v4.4s +; CHECK-DOT-NEXT: ret +; +; CHECK-DOT-I8MM-LABEL: partial_reduce_sext_cmp_i8tov2i64: +; CHECK-DOT-I8MM: // %bb.0: +; CHECK-DOT-I8MM-NEXT: movi v3.16b, #1 +; CHECK-DOT-I8MM-NEXT: movi v4.2d, #0000000000000000 +; CHECK-DOT-I8MM-NEXT: cmeq v1.16b, v1.16b, v2.16b +; CHECK-DOT-I8MM-NEXT: sdot v4.4s, v1.16b, v3.16b +; CHECK-DOT-I8MM-NEXT: saddw v0.2d, v0.2d, v4.2s +; CHECK-DOT-I8MM-NEXT: saddw2 v0.2d, v0.2d, v4.4s +; CHECK-DOT-I8MM-NEXT: ret + %cmp = icmp eq <16 x i8> %a, %b + %ext = sext <16 x i1> %cmp to <16 x i64> + %partial.reduce = tail call <2 x i64> @llvm.vector.partial.reduce.add(<2 x i64> %acc, <16 x i64> %ext) + ret <2 x i64> %partial.reduce +} + +define <2 x i64> @partial_reduce_zext_cmp_i32tov2i64(<2 x i64> %acc, <4 x i32> %a, <4 x i32> %b) { +; CHECK-COMMON-LABEL: partial_reduce_zext_cmp_i32tov2i64: +; CHECK-COMMON: // %bb.0: +; CHECK-COMMON-NEXT: cmeq v1.4s, v1.4s, v2.4s +; CHECK-COMMON-NEXT: mov w8, #1 // =0x1 +; CHECK-COMMON-NEXT: dup v2.2d, x8 +; CHECK-COMMON-NEXT: ushll v3.2d, v1.2s, #0 +; CHECK-COMMON-NEXT: ushll2 v1.2d, v1.4s, #0 +; CHECK-COMMON-NEXT: and v3.16b, v3.16b, v2.16b +; CHECK-COMMON-NEXT: and v1.16b, v1.16b, v2.16b +; CHECK-COMMON-NEXT: add v0.2d, v0.2d, v3.2d +; CHECK-COMMON-NEXT: add v0.2d, v0.2d, v1.2d +; CHECK-COMMON-NEXT: ret + %cmp = icmp eq <4 x i32> %a, %b + %ext = zext <4 x i1> %cmp to <4 x i64> + %partial.reduce = tail call <2 x i64> @llvm.vector.partial.reduce.add(<2 x i64> %acc, <4 x i64> %ext) + ret <2 x i64> %partial.reduce +} + +define <2 x i64> @partial_reduce_sext_cmp_i32tov2i64(<2 x i64> %acc, <4 x i32> %a, <4 x i32> %b) { +; CHECK-COMMON-LABEL: partial_reduce_sext_cmp_i32tov2i64: +; CHECK-COMMON: // %bb.0: +; CHECK-COMMON-NEXT: cmeq v1.4s, v1.4s, v2.4s +; CHECK-COMMON-NEXT: saddw v0.2d, v0.2d, v1.2s +; CHECK-COMMON-NEXT: saddw2 v0.2d, v0.2d, v1.4s +; CHECK-COMMON-NEXT: ret + %cmp = icmp eq <4 x i32> %a, %b + %ext = sext <4 x i1> %cmp to <4 x i64> + %partial.reduce = tail call <2 x i64> @llvm.vector.partial.reduce.add(<2 x i64> %acc, <4 x i64> %ext) + ret <2 x i64> %partial.reduce +} + _______________________________________________ llvm-branch-commits mailing list [email protected] https://lists.llvm.org/cgi-bin/mailman/listinfo/llvm-branch-commits
