[llvm] [X86][AVX-512] Fold `zext(and(bitcast(mask), C))` --> `and(anyext(bitcast(mask)), zext(C))` (PR #220251)
via llvm-commits
llvm-commits at lists.llvm.org
Thu Sep 3 01:24:55 PDT 2026
https://github.com/VachanVY updated https://github.com/llvm/llvm-project/pull/220251
>From b42fd9c65125b54a60db6aaf1c0c11807c8d7235 Mon Sep 17 00:00:00 2001
From: Vachan V Y <vachanvy05 at gmail.com>
Date: Wed, 2 Sep 2026 20:24:53 +0530
Subject: [PATCH 1/3] [X86][AVX-512] Fold `zext(and(bitcast(mask), C))` -->
`and(anyext(bitcast(mask)), zext(C))`
---
llvm/lib/Target/X86/X86ISelLowering.cpp | 49 +++++++++++++++++++
.../test/CodeGen/X86/aext-and-trunc-avx512.ll | 24 +++++++++
2 files changed, 73 insertions(+)
diff --git a/llvm/lib/Target/X86/X86ISelLowering.cpp b/llvm/lib/Target/X86/X86ISelLowering.cpp
index a0c92a22b7e2f..8dff93dfd704d 100644
--- a/llvm/lib/Target/X86/X86ISelLowering.cpp
+++ b/llvm/lib/Target/X86/X86ISelLowering.cpp
@@ -58098,6 +58098,52 @@ static SDValue widenBuildVec(SDNode *Extend, SelectionDAG &DAG) {
return SDValue();
}
+// zext(and(bitcast(mask), C)) --> and(anyext(bitcast(mask)), zext(C)).
+static SDValue combineZextMaskLogicToGPR(SDNode *N, SelectionDAG &DAG,
+ const X86Subtarget &Subtarget) {
+ if (N->getOpcode() != ISD::ZERO_EXTEND || !Subtarget.hasAVX512())
+ return SDValue();
+
+ EVT VT = N->getValueType(0);
+ if (VT != MVT::i32 && VT != MVT::i64)
+ return SDValue();
+
+ SDValue Logic = N->getOperand(0);
+ if (Logic.getOpcode() != ISD::AND || !Logic.hasOneUse())
+ return SDValue();
+
+ EVT NarrowVT = Logic.getValueType();
+ if (NarrowVT != MVT::i8 && NarrowVT != MVT::i16)
+ return SDValue();
+
+ auto GetMaskBitcast = [](SDValue V) -> SDValue {
+ if (V.getOpcode() != ISD::BITCAST)
+ return SDValue();
+ EVT SrcVT = V.getOperand(0).getValueType();
+ if (!SrcVT.isVector() || SrcVT.getVectorElementType() != MVT::i1)
+ return SDValue();
+ return V;
+ };
+
+ SDValue MaskBC = GetMaskBitcast(Logic.getOperand(0));
+ SDValue Other = Logic.getOperand(1);
+ if (!MaskBC) {
+ MaskBC = GetMaskBitcast(Logic.getOperand(1));
+ Other = Logic.getOperand(0);
+ }
+ if (!MaskBC)
+ return SDValue();
+
+ if (!isa<ConstantSDNode>(Other))
+ return SDValue();
+
+ SDLoc DL(N);
+ SDValue WideMask = DAG.getNode(ISD::ANY_EXTEND, DL, MVT::i32, MaskBC);
+ SDValue WideOther = DAG.getNode(ISD::ZERO_EXTEND, DL, MVT::i32, Other);
+ SDValue Wide = DAG.getNode(ISD::AND, DL, MVT::i32, WideMask, WideOther);
+ return DAG.getZExtOrTrunc(Wide, DL, VT);
+}
+
static SDValue combineZext(SDNode *N, SelectionDAG &DAG,
TargetLowering::DAGCombinerInfo &DCI,
const X86Subtarget &Subtarget) {
@@ -58123,6 +58169,9 @@ static SDValue combineZext(SDNode *N, SelectionDAG &DAG,
return SDValue(N, 0);
}
+ if (SDValue V = combineZextMaskLogicToGPR(N, DAG, Subtarget))
+ return V;
+
if (SDValue NewCMov = combineToExtendCMOV(N, DAG))
return NewCMov;
diff --git a/llvm/test/CodeGen/X86/aext-and-trunc-avx512.ll b/llvm/test/CodeGen/X86/aext-and-trunc-avx512.ll
index 78f525e4637bc..9ccde84c769f5 100644
--- a/llvm/test/CodeGen/X86/aext-and-trunc-avx512.ll
+++ b/llvm/test/CodeGen/X86/aext-and-trunc-avx512.ll
@@ -39,3 +39,27 @@ define i8 @ctpop_aext_i3_v3i1(ptr %p) {
%ext = zext i3 %ct to i8
ret i8 %ext
}
+
+define i64 @pr120389(<8 x i64> %0) {
+; BW-LABEL: pr120389:
+; BW: # %bb.0:
+; BW-NEXT: vpxor %xmm1, %xmm1, %xmm1
+; BW-NEXT: vpcmpgtq %zmm0, %zmm1, %k0
+; BW-NEXT: kmovd %k0, %eax
+; BW-NEXT: andl $1, %eax
+; BW-NEXT: vzeroupper
+; BW-NEXT: retq
+;
+; DQ-LABEL: pr120389:
+; DQ: # %bb.0:
+; DQ-NEXT: vpmovq2m %zmm0, %k0
+; DQ-NEXT: kmovd %k0, %eax
+; DQ-NEXT: andl $1, %eax
+; DQ-NEXT: vzeroupper
+; DQ-NEXT: retq
+ %2 = icmp slt <8 x i64> %0, zeroinitializer
+ %3 = bitcast <8 x i1> %2 to i8
+ %4 = and i8 %3, 1
+ %5 = zext nneg i8 %4 to i64
+ ret i64 %5
+}
>From 861bafecfbe36bd8be685930e2e7012ee060f990 Mon Sep 17 00:00:00 2001
From: Vachan V Y <vachanvy05 at gmail.com>
Date: Wed, 2 Sep 2026 23:00:44 +0530
Subject: [PATCH 2/3] [X86][AVX-512] Tablegen implementation
Fold `zext(and(bitcast(mask), C))` --> `and(anyext(bitcast(mask)), zext(C))`
---
llvm/lib/Target/X86/X86ISelLowering.cpp | 49 -------------------------
llvm/lib/Target/X86/X86InstrAVX512.td | 17 +++++++++
2 files changed, 17 insertions(+), 49 deletions(-)
diff --git a/llvm/lib/Target/X86/X86ISelLowering.cpp b/llvm/lib/Target/X86/X86ISelLowering.cpp
index 8dff93dfd704d..a0c92a22b7e2f 100644
--- a/llvm/lib/Target/X86/X86ISelLowering.cpp
+++ b/llvm/lib/Target/X86/X86ISelLowering.cpp
@@ -58098,52 +58098,6 @@ static SDValue widenBuildVec(SDNode *Extend, SelectionDAG &DAG) {
return SDValue();
}
-// zext(and(bitcast(mask), C)) --> and(anyext(bitcast(mask)), zext(C)).
-static SDValue combineZextMaskLogicToGPR(SDNode *N, SelectionDAG &DAG,
- const X86Subtarget &Subtarget) {
- if (N->getOpcode() != ISD::ZERO_EXTEND || !Subtarget.hasAVX512())
- return SDValue();
-
- EVT VT = N->getValueType(0);
- if (VT != MVT::i32 && VT != MVT::i64)
- return SDValue();
-
- SDValue Logic = N->getOperand(0);
- if (Logic.getOpcode() != ISD::AND || !Logic.hasOneUse())
- return SDValue();
-
- EVT NarrowVT = Logic.getValueType();
- if (NarrowVT != MVT::i8 && NarrowVT != MVT::i16)
- return SDValue();
-
- auto GetMaskBitcast = [](SDValue V) -> SDValue {
- if (V.getOpcode() != ISD::BITCAST)
- return SDValue();
- EVT SrcVT = V.getOperand(0).getValueType();
- if (!SrcVT.isVector() || SrcVT.getVectorElementType() != MVT::i1)
- return SDValue();
- return V;
- };
-
- SDValue MaskBC = GetMaskBitcast(Logic.getOperand(0));
- SDValue Other = Logic.getOperand(1);
- if (!MaskBC) {
- MaskBC = GetMaskBitcast(Logic.getOperand(1));
- Other = Logic.getOperand(0);
- }
- if (!MaskBC)
- return SDValue();
-
- if (!isa<ConstantSDNode>(Other))
- return SDValue();
-
- SDLoc DL(N);
- SDValue WideMask = DAG.getNode(ISD::ANY_EXTEND, DL, MVT::i32, MaskBC);
- SDValue WideOther = DAG.getNode(ISD::ZERO_EXTEND, DL, MVT::i32, Other);
- SDValue Wide = DAG.getNode(ISD::AND, DL, MVT::i32, WideMask, WideOther);
- return DAG.getZExtOrTrunc(Wide, DL, VT);
-}
-
static SDValue combineZext(SDNode *N, SelectionDAG &DAG,
TargetLowering::DAGCombinerInfo &DCI,
const X86Subtarget &Subtarget) {
@@ -58169,9 +58123,6 @@ static SDValue combineZext(SDNode *N, SelectionDAG &DAG,
return SDValue(N, 0);
}
- if (SDValue V = combineZextMaskLogicToGPR(N, DAG, Subtarget))
- return V;
-
if (SDValue NewCMov = combineToExtendCMOV(N, DAG))
return NewCMov;
diff --git a/llvm/lib/Target/X86/X86InstrAVX512.td b/llvm/lib/Target/X86/X86InstrAVX512.td
index 9fc27c5aff13c..1c1f4d5b28805 100644
--- a/llvm/lib/Target/X86/X86InstrAVX512.td
+++ b/llvm/lib/Target/X86/X86InstrAVX512.td
@@ -2732,6 +2732,23 @@ def : Pat<(i32 (anyext (i8 (bitconvert (v8i1 VK8:$src))))),
def : Pat<(i64 (anyext (i8 (bitconvert (v8i1 VK8:$src))))),
(INSERT_SUBREG (IMPLICIT_DEF), (COPY_TO_REGCLASS VK8:$src, GR32), sub_32bit)>;
+// zext(and(bitcast(mask), C)) --> and(anyext(bitcast(mask)), zext(C)). PR120389.
+def zext_i32imm : SDNodeXForm<imm, [{
+ return CurDAG->getTargetConstant(N->getZExtValue(), SDLoc(N), MVT::i32);
+}]>;
+def : Pat<(i32 (zext (and (i8 (bitconvert (v8i1 VK8:$src))), imm:$c))),
+ (AND32ri (COPY_TO_REGCLASS VK8:$src, GR32), (zext_i32imm imm:$c))>;
+def : Pat<(i64 (zext (and (i8 (bitconvert (v8i1 VK8:$src))), imm:$c))),
+ (SUBREG_TO_REG
+ (AND32ri (COPY_TO_REGCLASS VK8:$src, GR32), (zext_i32imm imm:$c)),
+ sub_32bit)>;
+def : Pat<(i32 (zext (and (i16 (bitconvert (v16i1 VK16:$src))), imm:$c))),
+ (AND32ri (COPY_TO_REGCLASS VK16:$src, GR32), (zext_i32imm imm:$c))>;
+def : Pat<(i64 (zext (and (i16 (bitconvert (v16i1 VK16:$src))), imm:$c))),
+ (SUBREG_TO_REG
+ (AND32ri (COPY_TO_REGCLASS VK16:$src, GR32), (zext_i32imm imm:$c)),
+ sub_32bit)>;
+
def : Pat<(v32i1 (bitconvert (i32 GR32:$src))),
(COPY_TO_REGCLASS GR32:$src, VK32)>;
def : Pat<(i32 (bitconvert (v32i1 VK32:$src))),
>From 1a2d3c22b991dbb1feb35bad5cf60d167da47965 Mon Sep 17 00:00:00 2001
From: Vachan V Y <vachanvy05 at gmail.com>
Date: Thu, 3 Sep 2026 13:54:38 +0530
Subject: [PATCH 3/3] [X86][AVX-512] Address review comment - add tests without
opt
Fold `zext(and(bitcast(mask), C))` --> `and(anyext(bitcast(mask)), zext(C))`
---
.../test/CodeGen/X86/aext-and-trunc-avx512.ll | 98 +++++++++++++++++--
1 file changed, 88 insertions(+), 10 deletions(-)
diff --git a/llvm/test/CodeGen/X86/aext-and-trunc-avx512.ll b/llvm/test/CodeGen/X86/aext-and-trunc-avx512.ll
index 9ccde84c769f5..d7acd63bba411 100644
--- a/llvm/test/CodeGen/X86/aext-and-trunc-avx512.ll
+++ b/llvm/test/CodeGen/X86/aext-and-trunc-avx512.ll
@@ -40,26 +40,104 @@ define i8 @ctpop_aext_i3_v3i1(ptr %p) {
ret i8 %ext
}
-define i64 @pr120389(<8 x i64> %0) {
-; BW-LABEL: pr120389:
+; Fold zext(and(bitcast(mask), C)) --> and(anyext(bitcast(mask)), zext(C)). PR120389.
+
+define i32 @mask_and_zext_i8_i32(<8 x i64> %0) {
+; BW-LABEL: mask_and_zext_i8_i32:
; BW: # %bb.0:
; BW-NEXT: vpxor %xmm1, %xmm1, %xmm1
; BW-NEXT: vpcmpgtq %zmm0, %zmm1, %k0
; BW-NEXT: kmovd %k0, %eax
-; BW-NEXT: andl $1, %eax
+; BW-NEXT: andb $3, %al
+; BW-NEXT: movzbl %al, %eax
; BW-NEXT: vzeroupper
; BW-NEXT: retq
;
-; DQ-LABEL: pr120389:
+; DQ-LABEL: mask_and_zext_i8_i32:
; DQ: # %bb.0:
; DQ-NEXT: vpmovq2m %zmm0, %k0
; DQ-NEXT: kmovd %k0, %eax
-; DQ-NEXT: andl $1, %eax
+; DQ-NEXT: andb $3, %al
+; DQ-NEXT: movzbl %al, %eax
+; DQ-NEXT: vzeroupper
+; DQ-NEXT: retq
+ %cmp = icmp slt <8 x i64> %0, zeroinitializer
+ %bc = bitcast <8 x i1> %cmp to i8
+ %and = and i8 %bc, 3
+ %ext = zext i8 %and to i32
+ ret i32 %ext
+}
+
+define i64 @mask_and_zext_i8_i64(<8 x i64> %0) {
+; BW-LABEL: mask_and_zext_i8_i64:
+; BW: # %bb.0:
+; BW-NEXT: vpxor %xmm1, %xmm1, %xmm1
+; BW-NEXT: vpcmpgtq %zmm0, %zmm1, %k0
+; BW-NEXT: kmovd %k0, %eax
+; BW-NEXT: andb $5, %al
+; BW-NEXT: movzbl %al, %eax
+; BW-NEXT: vzeroupper
+; BW-NEXT: retq
+;
+; DQ-LABEL: mask_and_zext_i8_i64:
+; DQ: # %bb.0:
+; DQ-NEXT: vpmovq2m %zmm0, %k0
+; DQ-NEXT: kmovd %k0, %eax
+; DQ-NEXT: andb $5, %al
+; DQ-NEXT: movzbl %al, %eax
+; DQ-NEXT: vzeroupper
+; DQ-NEXT: retq
+ %cmp = icmp slt <8 x i64> %0, zeroinitializer
+ %bc = bitcast <8 x i1> %cmp to i8
+ %and = and i8 %bc, 5
+ %ext = zext nneg i8 %and to i64
+ ret i64 %ext
+}
+
+define i32 @mask_and_zext_i16_i32(<16 x i32> %0) {
+; BW-LABEL: mask_and_zext_i16_i32:
+; BW: # %bb.0:
+; BW-NEXT: vpxor %xmm1, %xmm1, %xmm1
+; BW-NEXT: vpcmpgtd %zmm0, %zmm1, %k0
+; BW-NEXT: kmovd %k0, %eax
+; BW-NEXT: andl $7, %eax
+; BW-NEXT: vzeroupper
+; BW-NEXT: retq
+;
+; DQ-LABEL: mask_and_zext_i16_i32:
+; DQ: # %bb.0:
+; DQ-NEXT: vpmovd2m %zmm0, %k0
+; DQ-NEXT: kmovd %k0, %eax
+; DQ-NEXT: andl $7, %eax
+; DQ-NEXT: vzeroupper
+; DQ-NEXT: retq
+ %cmp = icmp slt <16 x i32> %0, zeroinitializer
+ %bc = bitcast <16 x i1> %cmp to i16
+ %and = and i16 %bc, 7
+ %ext = zext i16 %and to i32
+ ret i32 %ext
+}
+
+define i64 @mask_and_zext_i16_i64(<16 x i32> %0) {
+; BW-LABEL: mask_and_zext_i16_i64:
+; BW: # %bb.0:
+; BW-NEXT: vpxor %xmm1, %xmm1, %xmm1
+; BW-NEXT: vpcmpgtd %zmm0, %zmm1, %k0
+; BW-NEXT: kmovd %k0, %eax
+; BW-NEXT: andl $15, %eax
+; BW-NEXT: vzeroupper
+; BW-NEXT: retq
+;
+; DQ-LABEL: mask_and_zext_i16_i64:
+; DQ: # %bb.0:
+; DQ-NEXT: vpmovd2m %zmm0, %k0
+; DQ-NEXT: kmovd %k0, %eax
+; DQ-NEXT: andl $15, %eax
; DQ-NEXT: vzeroupper
; DQ-NEXT: retq
- %2 = icmp slt <8 x i64> %0, zeroinitializer
- %3 = bitcast <8 x i1> %2 to i8
- %4 = and i8 %3, 1
- %5 = zext nneg i8 %4 to i64
- ret i64 %5
+ %cmp = icmp slt <16 x i32> %0, zeroinitializer
+ %bc = bitcast <16 x i1> %cmp to i16
+ %and = and i16 %bc, 15
+ %ext = zext nneg i16 %and to i64
+ ret i64 %ext
}
More information about the llvm-commits
mailing list