[llvm] [X86] Keep bitcasted mask zero-tests in k-registers (PR #207908)

Jaeuk Lee via llvm-commits llvm-commits at lists.llvm.org
Tue Jul 7 00:48:46 PDT 2026


https://github.com/skku970412 updated https://github.com/llvm/llvm-project/pull/207908

>From 83af4e3571163b6142eb222337b33fc26308a9d2 Mon Sep 17 00:00:00 2001
From: =?UTF-8?q?=EC=9D=B4=EC=9E=AC=EC=9A=B1?=
 <126692701+skku970412 at users.noreply.github.com>
Date: Tue, 7 Jul 2026 15:55:37 +0900
Subject: [PATCH 1/2] [X86] Keep bitcasted mask tests in k-registers

---
 llvm/lib/Target/X86/X86ISelLowering.cpp       |  55 +++++++-
 .../CodeGen/X86/avx512-mask-bitcast-test.ll   | 117 ++++++++++++++++++
 2 files changed, 166 insertions(+), 6 deletions(-)
 create mode 100644 llvm/test/CodeGen/X86/avx512-mask-bitcast-test.ll

diff --git a/llvm/lib/Target/X86/X86ISelLowering.cpp b/llvm/lib/Target/X86/X86ISelLowering.cpp
index 9209134a055c6..b5d7b1bbb89d0 100644
--- a/llvm/lib/Target/X86/X86ISelLowering.cpp
+++ b/llvm/lib/Target/X86/X86ISelLowering.cpp
@@ -24980,15 +24980,60 @@ static SDValue EmitAVX512Test(SDValue Op0, SDValue Op1, ISD::CondCode CC,
                               SDValue &X86CC) {
   assert((CC == ISD::SETEQ || CC == ISD::SETNE) && "Unsupported ISD::CondCode");
 
+  auto IsLegalTestVT = [&](MVT VT) {
+    return (Subtarget.hasAVX512() && VT == MVT::v16i1) ||
+           (Subtarget.hasDQI() && VT == MVT::v8i1) ||
+           (Subtarget.hasBWI() && (VT == MVT::v32i1 || VT == MVT::v64i1));
+  };
+
+  auto IsKTestableVT = [&](MVT VT) {
+    return (Subtarget.hasDQI() && (VT == MVT::v8i1 || VT == MVT::v16i1)) ||
+           (Subtarget.hasBWI() && (VT == MVT::v32i1 || VT == MVT::v64i1));
+  };
+
+  // Match scalar tests of a bitcasted mask:
+  //   (and (bitcast vXi1 X), Y) == 0
+  // as KTEST/KORTEST so the vXi1 mask does not have to be copied to a GPR.
+  if (isNullConstant(Op1) && Op0.getOpcode() == ISD::AND && Op0.hasOneUse()) {
+    auto MatchScalarAnd = [&](SDValue Mask, SDValue ScalarMask) -> SDValue {
+      if (Mask.getOpcode() != ISD::BITCAST || !Mask.hasOneUse() ||
+          !ScalarMask.getValueType().isScalarInteger())
+        return SDValue();
+
+      // Keep one-use immediate masks on the scalar TEST path unless we know
+      // materializing them in a k-register is profitable.
+      if (isa<ConstantSDNode>(ScalarMask))
+        return SDValue();
+
+      SDValue MaskOp = Mask.getOperand(0);
+      MVT VT = MaskOp.getSimpleValueType();
+      if (!VT.isVectorOf(MVT::i1) || !IsLegalTestVT(VT))
+        return SDValue();
+
+      EVT IntVT = EVT::getIntegerVT(*DAG.getContext(),
+                                    VT.getVectorNumElements());
+      ScalarMask = DAG.getZExtOrTrunc(ScalarMask, dl, IntVT);
+      ScalarMask = DAG.getBitcast(VT, ScalarMask);
+
+      X86::CondCode X86Cond = CC == ISD::SETEQ ? X86::COND_E : X86::COND_NE;
+      X86CC = DAG.getTargetConstant(X86Cond, dl, MVT::i8);
+      SDValue And = DAG.getNode(ISD::AND, dl, VT, MaskOp, ScalarMask);
+      return DAG.getNode(X86ISD::KORTEST, dl, MVT::i32, And, And);
+    };
+
+    if (SDValue Test = MatchScalarAnd(Op0.getOperand(0), Op0.getOperand(1)))
+      return Test;
+    if (SDValue Test = MatchScalarAnd(Op0.getOperand(1), Op0.getOperand(0)))
+      return Test;
+  }
+
   // Must be a bitcast from vXi1.
   if (Op0.getOpcode() != ISD::BITCAST)
     return SDValue();
 
   Op0 = Op0.getOperand(0);
   MVT VT = Op0.getSimpleValueType();
-  if (!(Subtarget.hasAVX512() && VT == MVT::v16i1) &&
-      !(Subtarget.hasDQI() && VT == MVT::v8i1) &&
-      !(Subtarget.hasBWI() && (VT == MVT::v32i1 || VT == MVT::v64i1)))
+  if (!IsLegalTestVT(VT))
     return SDValue();
 
   X86::CondCode X86Cond;
@@ -25002,9 +25047,7 @@ static SDValue EmitAVX512Test(SDValue Op0, SDValue Op1, ISD::CondCode CC,
 
   // If the input is an AND, we can combine it's operands into the KTEST.
   bool KTestable = false;
-  if (Subtarget.hasDQI() && (VT == MVT::v8i1 || VT == MVT::v16i1))
-    KTestable = true;
-  if (Subtarget.hasBWI() && (VT == MVT::v32i1 || VT == MVT::v64i1))
+  if (IsKTestableVT(VT))
     KTestable = true;
   if (!isNullConstant(Op1))
     KTestable = false;
diff --git a/llvm/test/CodeGen/X86/avx512-mask-bitcast-test.ll b/llvm/test/CodeGen/X86/avx512-mask-bitcast-test.ll
new file mode 100644
index 0000000000000..9cc90b97f20f7
--- /dev/null
+++ b/llvm/test/CodeGen/X86/avx512-mask-bitcast-test.ll
@@ -0,0 +1,117 @@
+; RUN: llc < %s -mtriple=x86_64-unknown-unknown -mattr=+avx512bw,+avx512vl,+avx512dq | FileCheck %s
+; RUN: llc < %s -mtriple=x86_64-unknown-unknown -mattr=+avx512f | FileCheck %s --check-prefix=AVX512F
+
+define i32 @test_v8i1_scalar_mask(<8 x i32> %a, <8 x i32> %b, i8 %mask) {
+; CHECK-LABEL: test_v8i1_scalar_mask:
+; CHECK:       # %bb.0:
+; CHECK-NEXT:    kmovd %edi, %k1
+; CHECK-NEXT:    vpcmpneqd %ymm1, %ymm0, %k0 {%k1}
+; CHECK-NEXT:    xorl %eax, %eax
+; CHECK-NEXT:    kortestb %k0, %k0
+; CHECK-NEXT:    sete %al
+; CHECK-NEXT:    vzeroupper
+; CHECK-NEXT:    retq
+  %cmp = icmp ne <8 x i32> %a, %b
+  %bc = bitcast <8 x i1> %cmp to i8
+  %and = and i8 %bc, %mask
+  %eq = icmp eq i8 %and, 0
+  %ret = zext i1 %eq to i32
+  ret i32 %ret
+}
+
+define i32 @test_v16i1_scalar_mask(<2 x i64> %a, <2 x i64> %b, i16 %mask) {
+; CHECK-LABEL: test_v16i1_scalar_mask:
+; CHECK:       # %bb.0:
+; CHECK-NEXT:    kmovd %edi, %k1
+; CHECK-NEXT:    vpcmpneqb %xmm1, %xmm0, %k0 {%k1}
+; CHECK-NEXT:    xorl %eax, %eax
+; CHECK-NEXT:    kortestw %k0, %k0
+; CHECK-NEXT:    sete %al
+; CHECK-NEXT:    retq
+  %va = bitcast <2 x i64> %a to <16 x i8>
+  %vb = bitcast <2 x i64> %b to <16 x i8>
+  %cmp = icmp ne <16 x i8> %va, %vb
+  %bc = bitcast <16 x i1> %cmp to i16
+  %and = and i16 %bc, %mask
+  %eq = icmp eq i16 %and, 0
+  %ret = zext i1 %eq to i32
+  ret i32 %ret
+}
+
+define i32 @test_v16i1_scalar_mask_ne(<2 x i64> %a, <2 x i64> %b,
+                                      i16 %mask) {
+; CHECK-LABEL: test_v16i1_scalar_mask_ne:
+; CHECK:       # %bb.0:
+; CHECK-NEXT:    kmovd %edi, %k1
+; CHECK-NEXT:    vpcmpneqb %xmm1, %xmm0, %k0 {%k1}
+; CHECK-NEXT:    xorl %eax, %eax
+; CHECK-NEXT:    kortestw %k0, %k0
+; CHECK-NEXT:    setne %al
+; CHECK-NEXT:    retq
+  %va = bitcast <2 x i64> %a to <16 x i8>
+  %vb = bitcast <2 x i64> %b to <16 x i8>
+  %cmp = icmp ne <16 x i8> %va, %vb
+  %bc = bitcast <16 x i1> %cmp to i16
+  %and = and i16 %bc, %mask
+  %ne = icmp ne i16 %and, 0
+  %ret = zext i1 %ne to i32
+  ret i32 %ret
+}
+
+define i32 @test_v32i1_scalar_mask(<32 x i16> %a, <32 x i16> %b, i32 %mask) {
+; CHECK-LABEL: test_v32i1_scalar_mask:
+; CHECK:       # %bb.0:
+; CHECK-NEXT:    kmovd %edi, %k1
+; CHECK-NEXT:    vpcmpneqw %zmm1, %zmm0, %k0 {%k1}
+; CHECK-NEXT:    xorl %eax, %eax
+; CHECK-NEXT:    kortestd %k0, %k0
+; CHECK-NEXT:    sete %al
+; CHECK-NEXT:    vzeroupper
+; CHECK-NEXT:    retq
+  %cmp = icmp ne <32 x i16> %a, %b
+  %bc = bitcast <32 x i1> %cmp to i32
+  %and = and i32 %bc, %mask
+  %eq = icmp eq i32 %and, 0
+  %ret = zext i1 %eq to i32
+  ret i32 %ret
+}
+
+; Keep constant masks on the scalar immediate-test path for now. Materializing
+; a one-use immediate mask in a k-register is not clearly better.
+define i32 @test_v16i1_constant_mask(<2 x i64> %a, <2 x i64> %b) {
+; CHECK-LABEL: test_v16i1_constant_mask:
+; CHECK:       # %bb.0:
+; CHECK-NEXT:    vpcmpneqb %xmm1, %xmm0, %k0
+; CHECK-NEXT:    kmovd %k0, %ecx
+; CHECK-NEXT:    xorl %eax, %eax
+; CHECK-NEXT:    testb $127, %cl
+; CHECK-NEXT:    sete %al
+; CHECK-NEXT:    retq
+  %va = bitcast <2 x i64> %a to <16 x i8>
+  %vb = bitcast <2 x i64> %b to <16 x i8>
+  %cmp = icmp ne <16 x i8> %va, %vb
+  %bc = bitcast <16 x i1> %cmp to i16
+  %and = and i16 %bc, 127
+  %eq = icmp eq i16 %and, 0
+  %ret = zext i1 %eq to i32
+  ret i32 %ret
+}
+
+define i32 @test_v16i1_scalar_mask_avx512f(<16 x i32> %a, <16 x i32> %b,
+                                           i16 %mask) {
+; AVX512F-LABEL: test_v16i1_scalar_mask_avx512f:
+; AVX512F:       # %bb.0:
+; AVX512F-NEXT:    kmovw %edi, %k1
+; AVX512F-NEXT:    vpcmpneqd %zmm1, %zmm0, %k0 {%k1}
+; AVX512F-NEXT:    xorl %eax, %eax
+; AVX512F-NEXT:    kortestw %k0, %k0
+; AVX512F-NEXT:    sete %al
+; AVX512F-NEXT:    vzeroupper
+; AVX512F-NEXT:    retq
+  %cmp = icmp ne <16 x i32> %a, %b
+  %bc = bitcast <16 x i1> %cmp to i16
+  %and = and i16 %bc, %mask
+  %eq = icmp eq i16 %and, 0
+  %ret = zext i1 %eq to i32
+  ret i32 %ret
+}

>From c29cb53f39ea2e13c5d08378bb157815de589f0e Mon Sep 17 00:00:00 2001
From: =?UTF-8?q?=EC=9D=B4=EC=9E=AC=EC=9A=B1?=
 <126692701+skku970412 at users.noreply.github.com>
Date: Tue, 7 Jul 2026 16:48:34 +0900
Subject: [PATCH 2/2] [X86] Format AVX512 mask test combine

---
 llvm/lib/Target/X86/X86ISelLowering.cpp | 4 ++--
 1 file changed, 2 insertions(+), 2 deletions(-)

diff --git a/llvm/lib/Target/X86/X86ISelLowering.cpp b/llvm/lib/Target/X86/X86ISelLowering.cpp
index b5d7b1bbb89d0..b110668b69ef9 100644
--- a/llvm/lib/Target/X86/X86ISelLowering.cpp
+++ b/llvm/lib/Target/X86/X86ISelLowering.cpp
@@ -25010,8 +25010,8 @@ static SDValue EmitAVX512Test(SDValue Op0, SDValue Op1, ISD::CondCode CC,
       if (!VT.isVectorOf(MVT::i1) || !IsLegalTestVT(VT))
         return SDValue();
 
-      EVT IntVT = EVT::getIntegerVT(*DAG.getContext(),
-                                    VT.getVectorNumElements());
+      EVT IntVT =
+          EVT::getIntegerVT(*DAG.getContext(), VT.getVectorNumElements());
       ScalarMask = DAG.getZExtOrTrunc(ScalarMask, dl, IntVT);
       ScalarMask = DAG.getBitcast(VT, ScalarMask);
 



More information about the llvm-commits mailing list