[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