release/23.x: [DAGCombiner][AArch64] Fix the multiplier when folding a partial reduction of a mask (#215187) - #215512
Open
llvmbot wants to merge 2 commits into
Open
release/23.x: [DAGCombiner][AArch64] Fix the multiplier when folding a partial reduction of a mask (#215187)#215512llvmbot wants to merge 2 commits into
llvmbot wants to merge 2 commits into
Conversation
…ction of a mask (llvm#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 llvm#204897. (cherry picked from commit 9beddc4)
Member
Author
|
@sdesmalen-arm What do you think about merging this PR to the release branch? |
|
@llvm/pr-subscribers-llvm-selectiondag @llvm/pr-subscribers-backend-aarch64 Author: llvmbot ChangesBackport 9beddc4 Requested by: @sdesmalen-arm Full diff: https://github.com/llvm/llvm-project/pull/215512.diff 2 Files Affected:
diff --git a/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp b/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp
index e772abffbadff..3976a5800ee2f 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..3353cdf0cea58 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,215 @@ 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: uadalp 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: uadalp 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: sadalp 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: sadalp 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
+}
+
|
sdesmalen-arm
approved these changes
Aug 11, 2026
sdesmalen-arm
left a comment
Contributor
There was a problem hiding this comment.
This is a bug-fix that should make it onto the release branch. It is also low risk to cherry-pick (i.e. it only affects a specific case that was broken before and is fixed by this PR)
Contributor
|
I think one of the tests might need updating? |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Backport 9beddc4
Requested by: @sdesmalen-arm