[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