[llvm] Reapply "[SelectionDAG] Recurse through mask expression trees in WidenVSELECTMask" (PR #217307)
Valeriy Savchenko via llvm-commits
llvm-commits at lists.llvm.org
Wed Aug 19 06:02:34 PDT 2026
https://github.com/SavchenkoValeriy updated https://github.com/llvm/llvm-project/pull/217307
>From d1b044e4a1bc70829280947bf90c67a76c91a590 Mon Sep 17 00:00:00 2001
From: Valeriy Savchenko <vsavchenko at apple.com>
Date: Fri, 14 Aug 2026 12:05:21 +0100
Subject: [PATCH 1/2] Reapply "[SelectionDAG] Recurse through mask expression
trees in WidenVSELECTMask (#188085)" (#191151)
This reverts commit 9f47bcdb7c8a0d1be6481df3a5ac09d2eaf02690.
---
llvm/lib/CodeGen/SelectionDAG/LegalizeTypes.h | 14 ++
.../SelectionDAG/LegalizeVectorTypes.cpp | 181 +++++++++++-----
llvm/test/CodeGen/AArch64/arm64-zip.ll | 33 ++-
.../AArch64/vselect-widen-mask-tree.ll | 200 ++++++++++++++++++
.../X86/bitcast-int-to-vector-bool-sext.ll | 65 +++---
5 files changed, 387 insertions(+), 106 deletions(-)
create mode 100644 llvm/test/CodeGen/AArch64/vselect-widen-mask-tree.ll
diff --git a/llvm/lib/CodeGen/SelectionDAG/LegalizeTypes.h b/llvm/lib/CodeGen/SelectionDAG/LegalizeTypes.h
index ab252f5db2dcf..3a1e4b40a5c0d 100644
--- a/llvm/lib/CodeGen/SelectionDAG/LegalizeTypes.h
+++ b/llvm/lib/CodeGen/SelectionDAG/LegalizeTypes.h
@@ -1186,10 +1186,24 @@ class LLVM_LIBRARY_VISIBILITY DAGTypeLegalizer {
/// By default, the vector will be widened with undefined values.
SDValue ModifyToType(SDValue InOp, EVT NVT, bool FillWithZeroes = false);
+ /// Adjust element width (sign-extend/truncate) and element count
+ /// (extract/concat) of Mask to match ToMaskVT.
+ SDValue adjustMaskToType(SDValue Mask, EVT ToMaskVT);
+
+ /// Pick an intermediate VT and adjust both operands to it, minimizing
+ /// extend/truncate overhead given the final target ToVT.
+ EVT unifyMaskTypes(SDValue &Op0, SDValue &Op1, EVT ToVT);
+
/// Return a mask of vector type MaskVT to replace InMask. Also adjust
/// MaskVT to ToMaskVT if needed with vector extension or truncation.
SDValue convertMask(SDValue InMask, EVT MaskVT, EVT ToMaskVT);
+ /// Recursively convert a mask expression tree to ToVT, walking through
+ /// mask-preserving operations down to SETCC leaves. Avoids redundant
+ /// extend/truncate chains that arise when each node is converted
+ /// independently. Returns SDValue() if the tree cannot be converted.
+ SDValue convertMaskTree(SDValue V, EVT ToVT, unsigned Depth = 0);
+
//===--------------------------------------------------------------------===//
// Generic Splitting: LegalizeTypesGeneric.cpp
//===--------------------------------------------------------------------===//
diff --git a/llvm/lib/CodeGen/SelectionDAG/LegalizeVectorTypes.cpp b/llvm/lib/CodeGen/SelectionDAG/LegalizeVectorTypes.cpp
index 3417c9734af3b..2fd528ed77f64 100644
--- a/llvm/lib/CodeGen/SelectionDAG/LegalizeVectorTypes.cpp
+++ b/llvm/lib/CodeGen/SelectionDAG/LegalizeVectorTypes.cpp
@@ -7295,9 +7295,8 @@ static inline bool isSETCCorConvertedSETCC(SDValue N) {
// to ToMaskVT if needed with vector extension or truncation.
SDValue DAGTypeLegalizer::convertMask(SDValue InMask, EVT MaskVT,
EVT ToMaskVT) {
- // Currently a SETCC or a AND/OR/XOR with two SETCCs are handled.
- // FIXME: This code seems to be too restrictive, we might consider
- // generalizing it or dropping it.
+ // Called from convertMaskTree for SETCC leaf nodes. Re-creates the SETCC with
+ // result type MaskVT, then sign-extends/truncates and pads to ToMaskVT.
assert(isSETCCorConvertedSETCC(InMask) && "Unexpected mask argument.");
// Make a new Mask node, with a legal result VT.
@@ -7314,9 +7313,14 @@ SDValue DAGTypeLegalizer::convertMask(SDValue InMask, EVT MaskVT,
Mask = DAG.getNode(InMask->getOpcode(), SDLoc(InMask), MaskVT, Ops,
InMask->getFlags());
- // If MaskVT has smaller or bigger elements than ToMaskVT, a vector sign
- // extend or truncate is needed.
+ return adjustMaskToType(Mask, ToMaskVT);
+}
+
+// Adjust element width (sign-extend/truncate) and element count
+// (extract/concat) of Mask to match ToMaskVT.
+SDValue DAGTypeLegalizer::adjustMaskToType(SDValue Mask, EVT ToMaskVT) {
LLVMContext &Ctx = *DAG.getContext();
+ EVT MaskVT = Mask.getValueType();
unsigned MaskScalarBits = MaskVT.getScalarSizeInBits();
unsigned ToMaskScalBits = ToMaskVT.getScalarSizeInBits();
if (MaskScalarBits < ToMaskScalBits) {
@@ -7351,6 +7355,125 @@ SDValue DAGTypeLegalizer::convertMask(SDValue InMask, EVT MaskVT,
return Mask;
}
+// Adjust both operands to a common intermediate mask type, picking a scalar
+// width that minimizes extend/truncate overhead given the final target ToVT.
+EVT DAGTypeLegalizer::unifyMaskTypes(SDValue &Op0, SDValue &Op1, EVT ToVT) {
+ assert(Op0.getValueType().getVectorNumElements() ==
+ Op1.getValueType().getVectorNumElements() &&
+ "unifyMaskTypes only handles scalar width differences");
+ unsigned Bits0 = Op0.getValueType().getScalarSizeInBits();
+ unsigned Bits1 = Op1.getValueType().getScalarSizeInBits();
+ unsigned NarrowBits = std::min(Bits0, Bits1);
+ unsigned WideBits = std::max(Bits0, Bits1);
+ unsigned ToBits = ToVT.getScalarSizeInBits();
+ unsigned IntBits = NarrowBits == WideBits ? NarrowBits
+ : ToBits >= WideBits ? WideBits
+ : ToBits <= NarrowBits ? NarrowBits
+ : ToBits;
+ EVT OpVT = EVT::getVectorVT(*DAG.getContext(), MVT::getIntegerVT(IntBits),
+ Op0.getValueType().getVectorNumElements());
+ Op0 = adjustMaskToType(Op0, OpVT);
+ Op1 = adjustMaskToType(Op1, OpVT);
+ return OpVT;
+}
+
+SDValue DAGTypeLegalizer::convertMaskTree(SDValue V, EVT ToVT, unsigned Depth) {
+ if (Depth >= DAG.MaxRecursionDepth)
+ return SDValue();
+
+ SDValue Result = [&]() -> SDValue {
+ unsigned Opcode = V.getOpcode();
+
+ // Base case: SETCC produces the mask at its natural type.
+ if (isSETCCOp(Opcode)) {
+ EVT MaskVT = getSetCCResultType(getSETCCOperandType(V));
+ return convertMask(V, MaskVT, MaskVT);
+ }
+
+ // Base case: all-zeros or all-ones BUILD_VECTOR. Use ToVT directly since
+ // these are invariant under sign-extend/truncate.
+ if (ISD::isBuildVectorAllZeros(V.getNode()))
+ return DAG.getConstant(0, SDLoc(V), ToVT);
+ if (ISD::isBuildVectorAllOnes(V.getNode()))
+ return DAG.getAllOnesConstant(SDLoc(V), ToVT);
+
+ SDLoc DL(V);
+
+ // Logical operations (AND/OR/XOR): try picking the best fitting width out
+ // of children's element widths.
+ if (isLogicalMaskOp(Opcode)) {
+ SDValue Op0 = convertMaskTree(V.getOperand(0), ToVT, Depth + 1);
+ if (!Op0)
+ return SDValue();
+ SDValue Op1 = convertMaskTree(V.getOperand(1), ToVT, Depth + 1);
+ if (!Op1)
+ return SDValue();
+ EVT OpVT = unifyMaskTypes(Op0, Op1, ToVT);
+ return DAG.getNode(Opcode, DL, OpVT, Op0, Op1);
+ }
+
+ // FREEZE: widen the operand and re-wrap.
+ if (Opcode == ISD::FREEZE) {
+ SDValue Inner = convertMaskTree(V.getOperand(0), ToVT, Depth + 1);
+ if (!Inner)
+ return SDValue();
+ return DAG.getNode(ISD::FREEZE, DL, Inner.getValueType(), Inner);
+ }
+
+ // Vector shuffle: try inferring the best fitting width from operands.
+ if (Opcode == ISD::VECTOR_SHUFFLE) {
+ // Bail out when the number of elements is different, we can't
+ // simply reuse shuffle mask in this case.
+ if (V.getValueType().getVectorNumElements() !=
+ ToVT.getVectorNumElements())
+ return SDValue();
+
+ auto *Shuf = cast<ShuffleVectorSDNode>(V);
+ SDValue Op0 = convertMaskTree(V.getOperand(0), ToVT, Depth + 1);
+ if (!Op0)
+ return SDValue();
+ if (V.getOperand(1).isUndef()) {
+ EVT OpVT = Op0.getValueType();
+ return DAG.getVectorShuffle(OpVT, DL, Op0, DAG.getUNDEF(OpVT),
+ Shuf->getMask());
+ }
+ SDValue Op1 = convertMaskTree(V.getOperand(1), ToVT, Depth + 1);
+ if (!Op1)
+ return SDValue();
+ EVT OpVT = unifyMaskTypes(Op0, Op1, ToVT);
+ return DAG.getVectorShuffle(OpVT, DL, Op0, Op1, Shuf->getMask());
+ }
+
+ // SELECT/VSELECT: try inferring the best fitting width from operands.
+ if (Opcode == ISD::SELECT || Opcode == ISD::VSELECT) {
+ SDValue Op1 = convertMaskTree(V.getOperand(1), ToVT, Depth + 1);
+ if (!Op1)
+ return SDValue();
+ SDValue Op2 = convertMaskTree(V.getOperand(2), ToVT, Depth + 1);
+ if (!Op2)
+ return SDValue();
+ EVT OpVT = unifyMaskTypes(Op1, Op2, ToVT);
+
+ SDValue Cond = V.getOperand(0);
+ if (Opcode == ISD::VSELECT) {
+ Cond = convertMaskTree(Cond, ToVT, Depth + 1);
+ if (!Cond)
+ return SDValue();
+ Cond = adjustMaskToType(Cond, OpVT);
+ }
+ return DAG.getNode(Opcode, DL, OpVT, Cond, Op1, Op2);
+ }
+
+ return SDValue();
+ }();
+
+ if (!Result)
+ return SDValue();
+ if (Depth == 0)
+ Result = adjustMaskToType(Result, ToVT);
+ return Result;
+}
+
// This method tries to handle some special cases for the vselect mask
// and if needed adjusting the mask vector type to match that of the VSELECT.
// Without it, many cases end up with scalarization of the SETCC, with many
@@ -7362,9 +7485,6 @@ SDValue DAGTypeLegalizer::WidenVSELECTMask(SDNode *N) {
if (N->getOpcode() != ISD::VSELECT)
return SDValue();
- if (!isSETCCOp(Cond->getOpcode()) && !isLogicalMaskOp(Cond->getOpcode()))
- return SDValue();
-
// If this is a splitted VSELECT that was previously already handled, do
// nothing.
EVT CondVT = Cond->getValueType(0);
@@ -7417,49 +7537,8 @@ SDValue DAGTypeLegalizer::WidenVSELECTMask(SDNode *N) {
if (!ToMaskVT.getScalarType().isInteger())
ToMaskVT = ToMaskVT.changeVectorElementTypeToInteger();
- SDValue Mask;
- if (isSETCCOp(Cond->getOpcode())) {
- EVT MaskVT = getSetCCResultType(getSETCCOperandType(Cond));
- Mask = convertMask(Cond, MaskVT, ToMaskVT);
- } else if (isLogicalMaskOp(Cond->getOpcode()) &&
- isSETCCOp(Cond->getOperand(0).getOpcode()) &&
- isSETCCOp(Cond->getOperand(1).getOpcode())) {
- // Cond is (AND/OR/XOR (SETCC, SETCC))
- SDValue SETCC0 = Cond->getOperand(0);
- SDValue SETCC1 = Cond->getOperand(1);
- EVT VT0 = getSetCCResultType(getSETCCOperandType(SETCC0));
- EVT VT1 = getSetCCResultType(getSETCCOperandType(SETCC1));
- unsigned ScalarBits0 = VT0.getScalarSizeInBits();
- unsigned ScalarBits1 = VT1.getScalarSizeInBits();
- unsigned ScalarBits_ToMask = ToMaskVT.getScalarSizeInBits();
- EVT MaskVT;
- // If the two SETCCs have different VTs, either extend/truncate one of
- // them to the other "towards" ToMaskVT, or truncate one and extend the
- // other to ToMaskVT.
- if (ScalarBits0 != ScalarBits1) {
- EVT NarrowVT = ((ScalarBits0 < ScalarBits1) ? VT0 : VT1);
- EVT WideVT = ((NarrowVT == VT0) ? VT1 : VT0);
- if (ScalarBits_ToMask >= WideVT.getScalarSizeInBits())
- MaskVT = WideVT;
- else if (ScalarBits_ToMask <= NarrowVT.getScalarSizeInBits())
- MaskVT = NarrowVT;
- else
- MaskVT = ToMaskVT;
- } else
- // If the two SETCCs have the same VT, don't change it.
- MaskVT = VT0;
-
- // Make new SETCCs and logical nodes.
- SETCC0 = convertMask(SETCC0, VT0, MaskVT);
- SETCC1 = convertMask(SETCC1, VT1, MaskVT);
- Cond = DAG.getNode(Cond->getOpcode(), SDLoc(Cond), MaskVT, SETCC0, SETCC1);
-
- // Convert the logical op for VSELECT if needed.
- Mask = convertMask(Cond, MaskVT, ToMaskVT);
- } else
- return SDValue();
-
- return Mask;
+ // Try to recursively widen the mask expression tree to the target type.
+ return convertMaskTree(Cond, ToMaskVT);
}
SDValue DAGTypeLegalizer::WidenVecRes_Select(SDNode *N) {
diff --git a/llvm/test/CodeGen/AArch64/arm64-zip.ll b/llvm/test/CodeGen/AArch64/arm64-zip.ll
index bf017686ecc22..7b979fdde1c78 100644
--- a/llvm/test/CodeGen/AArch64/arm64-zip.ll
+++ b/llvm/test/CodeGen/AArch64/arm64-zip.ll
@@ -1,6 +1,6 @@
; NOTE: Assertions have been autogenerated by utils/update_llc_test_checks.py
; RUN: llc < %s -mtriple=arm64-eabi -aarch64-neon-syntax=apple | FileCheck %s --check-prefixes=CHECK,CHECK-SD
-; RUN: llc < %s -mtriple=arm64-eabi -aarch64-neon-syntax=apple -global-isel | FileCheck %s --check-prefixes=CHECK,CHECK-GI
+; RUN: llc < %s -mtriple=arm64-eabi -aarch64-neon-syntax=apple -global-isel -global-isel-abort=2 2>&1 | FileCheck %s --check-prefixes=CHECK,CHECK-GI
define <8 x i8> @vzipi8(ptr %A, ptr %B) nounwind {
; CHECK-LABEL: vzipi8:
@@ -378,13 +378,10 @@ define <4 x float> @shuffle_zip1(<4 x float> %arg) {
; CHECK-SD-LABEL: shuffle_zip1:
; CHECK-SD: // %bb.0: // %bb
; CHECK-SD-NEXT: fcmgt.4s v0, v0, #0.0
-; CHECK-SD-NEXT: uzp1.8h v1, v0, v0
-; CHECK-SD-NEXT: xtn.4h v0, v0
-; CHECK-SD-NEXT: xtn.4h v1, v1
-; CHECK-SD-NEXT: zip2.4h v0, v0, v1
; CHECK-SD-NEXT: fmov.4s v1, #1.00000000
-; CHECK-SD-NEXT: zip1.4h v0, v0, v0
-; CHECK-SD-NEXT: sshll.4s v0, v0, #0
+; CHECK-SD-NEXT: uzp1.4s v2, v0, v0
+; CHECK-SD-NEXT: zip2.4s v0, v0, v2
+; CHECK-SD-NEXT: zip1.4s v0, v0, v0
; CHECK-SD-NEXT: and.16b v0, v0, v1
; CHECK-SD-NEXT: ret
;
@@ -433,14 +430,11 @@ define <4 x i32> @shuffle_zip2(<4 x i32> %arg) {
; CHECK-SD-LABEL: shuffle_zip2:
; CHECK-SD: // %bb.0: // %bb
; CHECK-SD-NEXT: cmtst.4s v0, v0, v0
-; CHECK-SD-NEXT: movi.4h v1, #1
-; CHECK-SD-NEXT: uzp1.8h v2, v0, v0
-; CHECK-SD-NEXT: xtn.4h v0, v0
-; CHECK-SD-NEXT: xtn.4h v2, v2
-; CHECK-SD-NEXT: zip2.4h v0, v0, v2
-; CHECK-SD-NEXT: zip1.4h v0, v0, v0
-; CHECK-SD-NEXT: and.8b v0, v0, v1
-; CHECK-SD-NEXT: ushll.4s v0, v0, #0
+; CHECK-SD-NEXT: movi.4s v1, #1
+; CHECK-SD-NEXT: uzp1.4s v2, v0, v0
+; CHECK-SD-NEXT: zip2.4s v0, v0, v2
+; CHECK-SD-NEXT: zip1.4s v0, v0, v0
+; CHECK-SD-NEXT: and.16b v0, v0, v1
; CHECK-SD-NEXT: ret
;
; CHECK-GI-LABEL: shuffle_zip2:
@@ -488,13 +482,10 @@ define <4 x i32> @shuffle_zip3(<4 x i32> %arg) {
; CHECK-SD-LABEL: shuffle_zip3:
; CHECK-SD: // %bb.0: // %bb
; CHECK-SD-NEXT: cmgt.4s v0, v0, #0
-; CHECK-SD-NEXT: uzp1.8h v1, v0, v0
-; CHECK-SD-NEXT: xtn.4h v0, v0
-; CHECK-SD-NEXT: xtn.4h v1, v1
-; CHECK-SD-NEXT: zip2.4h v0, v0, v1
; CHECK-SD-NEXT: movi.4s v1, #1
-; CHECK-SD-NEXT: zip1.4h v0, v0, v0
-; CHECK-SD-NEXT: ushll.4s v0, v0, #0
+; CHECK-SD-NEXT: uzp1.4s v2, v0, v0
+; CHECK-SD-NEXT: zip2.4s v0, v0, v2
+; CHECK-SD-NEXT: zip1.4s v0, v0, v0
; CHECK-SD-NEXT: and.16b v0, v0, v1
; CHECK-SD-NEXT: ret
;
diff --git a/llvm/test/CodeGen/AArch64/vselect-widen-mask-tree.ll b/llvm/test/CodeGen/AArch64/vselect-widen-mask-tree.ll
new file mode 100644
index 0000000000000..918502940920f
--- /dev/null
+++ b/llvm/test/CodeGen/AArch64/vselect-widen-mask-tree.ll
@@ -0,0 +1,200 @@
+; NOTE: Assertions have been autogenerated by utils/update_llc_test_checks.py
+; RUN: llc -mtriple=arm64-apple-ios -o - %s | FileCheck %s
+
+define <4 x i32> @freeze_or_setcc(<4 x i32> %a, <4 x i32> %b, <4 x i32> %x, <4 x i32> %y) {
+; CHECK-LABEL: freeze_or_setcc:
+; CHECK: ; %bb.0:
+; CHECK-NEXT: add.4s v4, v0, v1
+; CHECK-NEXT: cmgt.4s v0, v1, v0
+; CHECK-NEXT: cmgt.4s v4, v4, #0
+; CHECK-NEXT: orr.16b v0, v4, v0
+; CHECK-NEXT: bsl.16b v0, v2, v3
+; CHECK-NEXT: ret
+ %add = add nsw <4 x i32> %a, %b
+ %cmp1 = icmp sgt <4 x i32> %add, zeroinitializer
+ %sub = sub nsw <4 x i32> %a, %b
+ %cmp2 = icmp slt <4 x i32> %sub, zeroinitializer
+ %or = or <4 x i1> %cmp1, %cmp2
+ %fr = freeze <4 x i1> %or
+ %sel = select <4 x i1> %fr, <4 x i32> %x, <4 x i32> %y
+ ret <4 x i32> %sel
+}
+
+define <4 x i32> @select_allones_or_setcc(<4 x i32> %a, <4 x i32> %x, <4 x i32> %y, i1 %cond) {
+; CHECK-LABEL: select_allones_or_setcc:
+; CHECK: ; %bb.0:
+; CHECK-NEXT: tst w0, #0x1
+; CHECK-NEXT: cmgt.4s v0, v0, #0
+; CHECK-NEXT: csetm w8, ne
+; CHECK-NEXT: dup.4s v3, w8
+; CHECK-NEXT: orr.16b v0, v0, v3
+; CHECK-NEXT: bsl.16b v0, v1, v2
+; CHECK-NEXT: ret
+ %cmp = icmp sgt <4 x i32> %a, zeroinitializer
+ %mask = select i1 %cond, <4 x i1> <i1 true, i1 true, i1 true, i1 true>, <4 x i1> %cmp
+ %sel = select <4 x i1> %mask, <4 x i32> %x, <4 x i32> %y
+ ret <4 x i32> %sel
+}
+
+define <4 x i32> @select_setcc_or_allzeros(<4 x i32> %a, <4 x i32> %x, <4 x i32> %y, i1 %cond) {
+; CHECK-LABEL: select_setcc_or_allzeros:
+; CHECK: ; %bb.0:
+; CHECK-NEXT: tst w0, #0x1
+; CHECK-NEXT: cmgt.4s v0, v0, #0
+; CHECK-NEXT: csetm w8, ne
+; CHECK-NEXT: dup.4s v3, w8
+; CHECK-NEXT: and.16b v0, v0, v3
+; CHECK-NEXT: bsl.16b v0, v1, v2
+; CHECK-NEXT: ret
+ %cmp = icmp sgt <4 x i32> %a, zeroinitializer
+ %mask = select i1 %cond, <4 x i1> %cmp, <4 x i1> zeroinitializer
+ %sel = select <4 x i1> %mask, <4 x i32> %x, <4 x i32> %y
+ ret <4 x i32> %sel
+}
+
+define <4 x i32> @select_allzeros_or_allones(<4 x i32> %x, <4 x i32> %y, i1 %cond) {
+; CHECK-LABEL: select_allzeros_or_allones:
+; CHECK: ; %bb.0:
+; CHECK-NEXT: tst w0, #0x1
+; CHECK-NEXT: csetm w8, ne
+; CHECK-NEXT: dup.4s v2, w8
+; CHECK-NEXT: bit.16b v0, v1, v2
+; CHECK-NEXT: ret
+ %mask = select i1 %cond, <4 x i1> zeroinitializer, <4 x i1> <i1 true, i1 true, i1 true, i1 true>
+ %sel = select <4 x i1> %mask, <4 x i32> %x, <4 x i32> %y
+ ret <4 x i32> %sel
+}
+
+define <4 x i32> @vselect_of_setccs(<4 x i32> %a, <4 x i32> %b, <4 x i32> %x, <4 x i32> %y) {
+; CHECK-LABEL: vselect_of_setccs:
+; CHECK: ; %bb.0:
+; CHECK-NEXT: movi.4s v4, #100
+; CHECK-NEXT: cmgt.4s v5, v0, #0
+; CHECK-NEXT: cmeq.4s v1, v1, #0
+; CHECK-NEXT: cmgt.4s v0, v4, v0
+; CHECK-NEXT: bit.16b v0, v5, v1
+; CHECK-NEXT: bsl.16b v0, v2, v3
+; CHECK-NEXT: ret
+ %cmp1 = icmp sgt <4 x i32> %a, zeroinitializer
+ %cmp2 = icmp slt <4 x i32> %a, <i32 100, i32 100, i32 100, i32 100>
+ %cond = icmp eq <4 x i32> %b, zeroinitializer
+ %mask = select <4 x i1> %cond, <4 x i1> %cmp1, <4 x i1> %cmp2
+ %sel = select <4 x i1> %mask, <4 x i32> %x, <4 x i32> %y
+ ret <4 x i32> %sel
+}
+
+define <4 x i32> @select_scalar_cond_setccs(<4 x i32> %a, <4 x i32> %x, <4 x i32> %y, i1 %cond) {
+; CHECK-LABEL: select_scalar_cond_setccs:
+; CHECK: ; %bb.0: ; %entry
+; CHECK-NEXT: movi.4s v3, #100
+; CHECK-NEXT: tst w0, #0x1
+; CHECK-NEXT: cmgt.4s v4, v0, #0
+; CHECK-NEXT: csetm w8, ne
+; CHECK-NEXT: cmgt.4s v0, v3, v0
+; CHECK-NEXT: dup.4s v3, w8
+; CHECK-NEXT: bif.16b v0, v4, v3
+; CHECK-NEXT: bsl.16b v0, v1, v2
+; CHECK-NEXT: ret
+entry:
+ %cmp1 = icmp sgt <4 x i32> %a, zeroinitializer
+ br i1 %cond, label %then, label %else
+
+then:
+ %cmp2 = icmp slt <4 x i32> %a, <i32 100, i32 100, i32 100, i32 100>
+ br label %merge
+
+else:
+ br label %merge
+
+merge:
+ %mask = phi <4 x i1> [ %cmp2, %then ], [ %cmp1, %else ]
+ %fr = freeze <4 x i1> %mask
+ %sel = select <4 x i1> %fr, <4 x i32> %x, <4 x i32> %y
+ ret <4 x i32> %sel
+}
+
+define <3 x i64> @or_setcc_i16_i32_sel_i64(<3 x i16> %a, <3 x i16> %b, <3 x i32> %c, <3 x i32> %d, <3 x i64> %x, <3 x i64> %y) {
+; CHECK-LABEL: or_setcc_i16_i32_sel_i64:
+; CHECK: ; %bb.0:
+; CHECK-NEXT: cmgt.4h v0, v0, v1
+; CHECK-NEXT: cmgt.4s v1, v2, v3
+; CHECK-NEXT: ; kill: def $d7 killed $d7 def $q7
+; CHECK-NEXT: ; kill: def $d4 killed $d4 def $q4
+; CHECK-NEXT: ; kill: def $d5 killed $d5 def $q5
+; CHECK-NEXT: ; kill: def $d6 killed $d6 def $q6
+; CHECK-NEXT: mov.d v4[1], v5[0]
+; CHECK-NEXT: sshll.4s v0, v0, #0
+; CHECK-NEXT: orr.16b v1, v0, v1
+; CHECK-NEXT: ldp d0, d2, [sp]
+; CHECK-NEXT: mov.d v7[1], v0[0]
+; CHECK-NEXT: sshll.2d v0, v1, #0
+; CHECK-NEXT: sshll2.2d v1, v1, #0
+; CHECK-NEXT: bit.16b v2, v6, v1
+; CHECK-NEXT: bsl.16b v0, v4, v7
+; CHECK-NEXT: ; kill: def $d2 killed $d2 killed $q2
+; CHECK-NEXT: mov d1, v0[1]
+; CHECK-NEXT: ; kill: def $d0 killed $d0 killed $q0
+; CHECK-NEXT: ret
+ %cmp0 = icmp sgt <3 x i16> %a, %b
+ %cmp1 = icmp sgt <3 x i32> %c, %d
+ %or = or <3 x i1> %cmp0, %cmp1
+ %sel = select <3 x i1> %or, <3 x i64> %x, <3 x i64> %y
+ ret <3 x i64> %sel
+}
+
+define <3 x i64> @and_setcc_i32_i32_sel_i64(<3 x i32> %a, <3 x i32> %b, <3 x i32> %c, <3 x i32> %d, <3 x i64> %x, <3 x i64> %y) {
+; CHECK-LABEL: and_setcc_i32_i32_sel_i64:
+; CHECK: ; %bb.0:
+; CHECK-NEXT: cmgt.4s v2, v2, v3
+; CHECK-NEXT: cmgt.4s v0, v0, v1
+; CHECK-NEXT: ; kill: def $d7 killed $d7 def $q7
+; CHECK-NEXT: ; kill: def $d4 killed $d4 def $q4
+; CHECK-NEXT: ; kill: def $d5 killed $d5 def $q5
+; CHECK-NEXT: ; kill: def $d6 killed $d6 def $q6
+; CHECK-NEXT: mov.d v4[1], v5[0]
+; CHECK-NEXT: and.16b v1, v0, v2
+; CHECK-NEXT: ldp d0, d2, [sp]
+; CHECK-NEXT: mov.d v7[1], v0[0]
+; CHECK-NEXT: sshll.2d v0, v1, #0
+; CHECK-NEXT: sshll2.2d v1, v1, #0
+; CHECK-NEXT: bit.16b v2, v6, v1
+; CHECK-NEXT: bsl.16b v0, v4, v7
+; CHECK-NEXT: ; kill: def $d2 killed $d2 killed $q2
+; CHECK-NEXT: mov d1, v0[1]
+; CHECK-NEXT: ; kill: def $d0 killed $d0 killed $q0
+; CHECK-NEXT: ret
+ %cmp0 = icmp sgt <3 x i32> %a, %b
+ %cmp1 = icmp sgt <3 x i32> %c, %d
+ %and = and <3 x i1> %cmp0, %cmp1
+ %sel = select <3 x i1> %and, <3 x i64> %x, <3 x i64> %y
+ ret <3 x i64> %sel
+}
+
+define <3 x i64> @or_setcc_i16_i16_sel_i64(<3 x i16> %a, <3 x i16> %b, <3 x i16> %c, <3 x i16> %d, <3 x i64> %x, <3 x i64> %y) {
+; CHECK-LABEL: or_setcc_i16_i16_sel_i64:
+; CHECK: ; %bb.0:
+; CHECK-NEXT: cmgt.4h v2, v2, v3
+; CHECK-NEXT: cmgt.4h v0, v0, v1
+; CHECK-NEXT: ; kill: def $d7 killed $d7 def $q7
+; CHECK-NEXT: ; kill: def $d4 killed $d4 def $q4
+; CHECK-NEXT: ; kill: def $d5 killed $d5 def $q5
+; CHECK-NEXT: ; kill: def $d6 killed $d6 def $q6
+; CHECK-NEXT: mov.d v4[1], v5[0]
+; CHECK-NEXT: orr.8b v0, v0, v2
+; CHECK-NEXT: sshll.4s v1, v0, #0
+; CHECK-NEXT: ldp d0, d2, [sp]
+; CHECK-NEXT: mov.d v7[1], v0[0]
+; CHECK-NEXT: sshll.2d v0, v1, #0
+; CHECK-NEXT: sshll2.2d v1, v1, #0
+; CHECK-NEXT: bit.16b v2, v6, v1
+; CHECK-NEXT: bsl.16b v0, v4, v7
+; CHECK-NEXT: ; kill: def $d2 killed $d2 killed $q2
+; CHECK-NEXT: mov d1, v0[1]
+; CHECK-NEXT: ; kill: def $d0 killed $d0 killed $q0
+; CHECK-NEXT: ret
+ %cmp0 = icmp sgt <3 x i16> %a, %b
+ %cmp1 = icmp sgt <3 x i16> %c, %d
+ %or = or <3 x i1> %cmp0, %cmp1
+ %sel = select <3 x i1> %or, <3 x i64> %x, <3 x i64> %y
+ ret <3 x i64> %sel
+}
diff --git a/llvm/test/CodeGen/X86/bitcast-int-to-vector-bool-sext.ll b/llvm/test/CodeGen/X86/bitcast-int-to-vector-bool-sext.ll
index 0633daa200ef4..ed39881f65083 100644
--- a/llvm/test/CodeGen/X86/bitcast-int-to-vector-bool-sext.ll
+++ b/llvm/test/CodeGen/X86/bitcast-int-to-vector-bool-sext.ll
@@ -660,32 +660,30 @@ define <8 x i32> @PR157382(ptr %p0, ptr %p1, ptr %p2) {
; SSE2-SSSE3: # %bb.0:
; SSE2-SSSE3-NEXT: movdqu (%rdi), %xmm3
; SSE2-SSSE3-NEXT: movdqu 16(%rdi), %xmm2
-; SSE2-SSSE3-NEXT: movdqu (%rsi), %xmm0
+; SSE2-SSSE3-NEXT: movdqu (%rsi), %xmm5
; SSE2-SSSE3-NEXT: movdqu 16(%rsi), %xmm4
; SSE2-SSSE3-NEXT: movq {{.*#+}} xmm1 = mem[0],zero
-; SSE2-SSSE3-NEXT: pxor %xmm5, %xmm5
+; SSE2-SSSE3-NEXT: pxor %xmm0, %xmm0
; SSE2-SSSE3-NEXT: pxor %xmm6, %xmm6
-; SSE2-SSSE3-NEXT: pcmpgtd %xmm3, %xmm6
+; SSE2-SSSE3-NEXT: pcmpgtd %xmm2, %xmm6
; SSE2-SSSE3-NEXT: pcmpeqd %xmm7, %xmm7
; SSE2-SSSE3-NEXT: pxor %xmm7, %xmm6
; SSE2-SSSE3-NEXT: pxor %xmm8, %xmm8
-; SSE2-SSSE3-NEXT: pcmpgtd %xmm2, %xmm8
+; SSE2-SSSE3-NEXT: pcmpgtd %xmm3, %xmm8
; SSE2-SSSE3-NEXT: pxor %xmm7, %xmm8
-; SSE2-SSSE3-NEXT: pcmpgtd %xmm5, %xmm0
-; SSE2-SSSE3-NEXT: por %xmm6, %xmm0
-; SSE2-SSSE3-NEXT: pcmpgtd %xmm5, %xmm4
-; SSE2-SSSE3-NEXT: por %xmm8, %xmm4
-; SSE2-SSSE3-NEXT: packssdw %xmm4, %xmm0
+; SSE2-SSSE3-NEXT: pcmpgtd %xmm0, %xmm4
+; SSE2-SSSE3-NEXT: por %xmm6, %xmm4
+; SSE2-SSSE3-NEXT: pcmpgtd %xmm0, %xmm5
+; SSE2-SSSE3-NEXT: por %xmm8, %xmm5
; SSE2-SSSE3-NEXT: punpcklbw {{.*#+}} xmm1 = xmm1[0,0,1,1,2,2,3,3,4,4,5,5,6,6,7,7]
-; SSE2-SSSE3-NEXT: pcmpeqb %xmm5, %xmm1
+; SSE2-SSSE3-NEXT: pcmpeqb %xmm0, %xmm1
; SSE2-SSSE3-NEXT: pxor %xmm7, %xmm1
-; SSE2-SSSE3-NEXT: por %xmm0, %xmm1
+; SSE2-SSSE3-NEXT: movdqa %xmm1, %xmm0
; SSE2-SSSE3-NEXT: punpcklwd {{.*#+}} xmm0 = xmm0[0],xmm1[0],xmm0[1],xmm1[1],xmm0[2],xmm1[2],xmm0[3],xmm1[3]
-; SSE2-SSSE3-NEXT: psrad $16, %xmm0
-; SSE2-SSSE3-NEXT: pand %xmm3, %xmm0
+; SSE2-SSSE3-NEXT: por %xmm5, %xmm0
; SSE2-SSSE3-NEXT: punpckhwd {{.*#+}} xmm1 = xmm1[4,4,5,5,6,6,7,7]
-; SSE2-SSSE3-NEXT: pslld $31, %xmm1
-; SSE2-SSSE3-NEXT: psrad $31, %xmm1
+; SSE2-SSSE3-NEXT: por %xmm4, %xmm1
+; SSE2-SSSE3-NEXT: pand %xmm3, %xmm0
; SSE2-SSSE3-NEXT: pand %xmm2, %xmm1
; SSE2-SSSE3-NEXT: retq
;
@@ -693,28 +691,27 @@ define <8 x i32> @PR157382(ptr %p0, ptr %p1, ptr %p2) {
; AVX1: # %bb.0:
; AVX1-NEXT: vmovdqu (%rdi), %ymm0
; AVX1-NEXT: vmovq {{.*#+}} xmm1 = mem[0],zero
-; AVX1-NEXT: vpxor %xmm2, %xmm2, %xmm2
-; AVX1-NEXT: vpcmpgtd %xmm0, %xmm2, %xmm3
+; AVX1-NEXT: vextractf128 $1, %ymm0, %xmm2
+; AVX1-NEXT: vpxor %xmm3, %xmm3, %xmm3
+; AVX1-NEXT: vpcmpgtd %xmm2, %xmm3, %xmm2
; AVX1-NEXT: vpcmpeqd %xmm4, %xmm4, %xmm4
-; AVX1-NEXT: vpxor %xmm4, %xmm3, %xmm3
-; AVX1-NEXT: vextractf128 $1, %ymm0, %xmm5
-; AVX1-NEXT: vpcmpgtd %xmm5, %xmm2, %xmm5
+; AVX1-NEXT: vpxor %xmm4, %xmm2, %xmm2
+; AVX1-NEXT: vpcmpgtd %xmm0, %xmm3, %xmm5
; AVX1-NEXT: vpxor %xmm4, %xmm5, %xmm5
-; AVX1-NEXT: vmovdqu (%rsi), %xmm6
-; AVX1-NEXT: vmovdqu 16(%rsi), %xmm7
-; AVX1-NEXT: vpcmpgtd %xmm2, %xmm6, %xmm6
-; AVX1-NEXT: vpor %xmm6, %xmm3, %xmm3
-; AVX1-NEXT: vpcmpgtd %xmm2, %xmm7, %xmm6
-; AVX1-NEXT: vpor %xmm6, %xmm5, %xmm5
-; AVX1-NEXT: vpackssdw %xmm5, %xmm3, %xmm3
-; AVX1-NEXT: vpcmpeqb %xmm2, %xmm1, %xmm1
+; AVX1-NEXT: vinsertf128 $1, %xmm2, %ymm5, %ymm2
+; AVX1-NEXT: vmovdqu (%rsi), %xmm5
+; AVX1-NEXT: vmovdqu 16(%rsi), %xmm6
+; AVX1-NEXT: vpcmpgtd %xmm3, %xmm6, %xmm6
+; AVX1-NEXT: vpcmpgtd %xmm3, %xmm5, %xmm5
+; AVX1-NEXT: vinsertf128 $1, %xmm6, %ymm5, %ymm5
+; AVX1-NEXT: vorps %ymm5, %ymm2, %ymm2
+; AVX1-NEXT: vpcmpeqb %xmm3, %xmm1, %xmm1
; AVX1-NEXT: vpxor %xmm4, %xmm1, %xmm1
-; AVX1-NEXT: vpmovsxbw %xmm1, %xmm1
-; AVX1-NEXT: vpor %xmm1, %xmm3, %xmm1
-; AVX1-NEXT: vpmovsxwd %xmm1, %xmm2
-; AVX1-NEXT: vpshufd {{.*#+}} xmm1 = xmm1[2,3,2,3]
-; AVX1-NEXT: vpmovsxwd %xmm1, %xmm1
-; AVX1-NEXT: vinsertf128 $1, %xmm1, %ymm2, %ymm1
+; AVX1-NEXT: vpmovsxbd %xmm1, %xmm3
+; AVX1-NEXT: vpshufd {{.*#+}} xmm1 = xmm1[1,1,1,1]
+; AVX1-NEXT: vpmovsxbd %xmm1, %xmm1
+; AVX1-NEXT: vinsertf128 $1, %xmm1, %ymm3, %ymm1
+; AVX1-NEXT: vorps %ymm1, %ymm2, %ymm1
; AVX1-NEXT: vandps %ymm0, %ymm1, %ymm0
; AVX1-NEXT: retq
;
>From d0f6213d1184eb0c38dbc9cfe9f65a6fe7592c46 Mon Sep 17 00:00:00 2001
From: Valeriy Savchenko <vsavchenko at apple.com>
Date: Mon, 17 Aug 2026 17:24:38 +0100
Subject: [PATCH 2/2] [SelectionDAG] Restrict select mask type operations
---
llvm/lib/CodeGen/SelectionDAG/LegalizeTypes.h | 15 +-
.../SelectionDAG/LegalizeVectorTypes.cpp | 265 +++++++++++-------
.../CodeGen/X86/vselect-widen-mask-tree.ll | 93 ++++++
3 files changed, 277 insertions(+), 96 deletions(-)
create mode 100644 llvm/test/CodeGen/X86/vselect-widen-mask-tree.ll
diff --git a/llvm/lib/CodeGen/SelectionDAG/LegalizeTypes.h b/llvm/lib/CodeGen/SelectionDAG/LegalizeTypes.h
index 3a1e4b40a5c0d..711865cc66231 100644
--- a/llvm/lib/CodeGen/SelectionDAG/LegalizeTypes.h
+++ b/llvm/lib/CodeGen/SelectionDAG/LegalizeTypes.h
@@ -1192,7 +1192,12 @@ class LLVM_LIBRARY_VISIBILITY DAGTypeLegalizer {
/// Pick an intermediate VT and adjust both operands to it, minimizing
/// extend/truncate overhead given the final target ToVT.
- EVT unifyMaskTypes(SDValue &Op0, SDValue &Op1, EVT ToVT);
+ ///
+ /// IsOpLenient flags retain information on whether the corresponding
+ /// Op can be cheaply materialized to whatever type. If one of the
+ /// operands lenient, we force cast it to match other operand's type.
+ EVT unifyMaskTypes(SDValue &Op0, bool IsOpLenient0, SDValue &Op1,
+ bool IsOpLenient1, EVT ToVT);
/// Return a mask of vector type MaskVT to replace InMask. Also adjust
/// MaskVT to ToMaskVT if needed with vector extension or truncation.
@@ -1202,7 +1207,13 @@ class LLVM_LIBRARY_VISIBILITY DAGTypeLegalizer {
/// mask-preserving operations down to SETCC leaves. Avoids redundant
/// extend/truncate chains that arise when each node is converted
/// independently. Returns SDValue() if the tree cannot be converted.
- SDValue convertMaskTree(SDValue V, EVT ToVT, unsigned Depth = 0);
+ SDValue convertMaskTree(SDValue V, EVT ToVT);
+
+ /// Recursive implementation of convertMaskTree.
+ /// In addition to the value returns a boolean flag signifying whether the
+ /// tree can be materialized with any type.
+ std::pair<SDValue, bool> convertMaskTreeImpl(SDValue V, EVT ToVT,
+ unsigned Depth = 0);
//===--------------------------------------------------------------------===//
// Generic Splitting: LegalizeTypesGeneric.cpp
diff --git a/llvm/lib/CodeGen/SelectionDAG/LegalizeVectorTypes.cpp b/llvm/lib/CodeGen/SelectionDAG/LegalizeVectorTypes.cpp
index 2fd528ed77f64..9f8f87e986dcf 100644
--- a/llvm/lib/CodeGen/SelectionDAG/LegalizeVectorTypes.cpp
+++ b/llvm/lib/CodeGen/SelectionDAG/LegalizeVectorTypes.cpp
@@ -7357,10 +7357,37 @@ SDValue DAGTypeLegalizer::adjustMaskToType(SDValue Mask, EVT ToMaskVT) {
// Adjust both operands to a common intermediate mask type, picking a scalar
// width that minimizes extend/truncate overhead given the final target ToVT.
-EVT DAGTypeLegalizer::unifyMaskTypes(SDValue &Op0, SDValue &Op1, EVT ToVT) {
+EVT DAGTypeLegalizer::unifyMaskTypes(SDValue &Op0, bool IsOpLenient0,
+ SDValue &Op1, bool IsOpLenient1,
+ EVT ToVT) {
assert(Op0.getValueType().getVectorNumElements() ==
Op1.getValueType().getVectorNumElements() &&
"unifyMaskTypes only handles scalar width differences");
+
+ // If only one of the operands lenient type-wise, we can simply
+ // adjust its type to the other operand's type assuming that this
+ // adjustment can be folded away.
+ //
+ // NOTE: We essentially rely on the fact that further optimizations
+ // do spot redundant casts in "lenient" cases. If at some
+ // point we decide that we want to do better here, we can
+ // postpone converting lenient sub-trees right away and postpone
+ // it to the moment when we know the best fitting integer type
+ // to materialize them, and do it there.
+ if (IsOpLenient0 != IsOpLenient1) {
+ SDValue *LenientOp, *NonLenientOp;
+ if (IsOpLenient0) {
+ LenientOp = &Op0;
+ NonLenientOp = &Op1;
+ } else {
+ LenientOp = &Op1;
+ NonLenientOp = &Op0;
+ }
+ EVT OpVT = NonLenientOp->getValueType();
+ *LenientOp = adjustMaskToType(*LenientOp, OpVT);
+ return OpVT;
+ }
+
unsigned Bits0 = Op0.getValueType().getScalarSizeInBits();
unsigned Bits1 = Op1.getValueType().getScalarSizeInBits();
unsigned NarrowBits = std::min(Bits0, Bits1);
@@ -7377,101 +7404,151 @@ EVT DAGTypeLegalizer::unifyMaskTypes(SDValue &Op0, SDValue &Op1, EVT ToVT) {
return OpVT;
}
-SDValue DAGTypeLegalizer::convertMaskTree(SDValue V, EVT ToVT, unsigned Depth) {
- if (Depth >= DAG.MaxRecursionDepth)
- return SDValue();
-
- SDValue Result = [&]() -> SDValue {
- unsigned Opcode = V.getOpcode();
-
- // Base case: SETCC produces the mask at its natural type.
- if (isSETCCOp(Opcode)) {
- EVT MaskVT = getSetCCResultType(getSETCCOperandType(V));
- return convertMask(V, MaskVT, MaskVT);
- }
-
- // Base case: all-zeros or all-ones BUILD_VECTOR. Use ToVT directly since
- // these are invariant under sign-extend/truncate.
- if (ISD::isBuildVectorAllZeros(V.getNode()))
- return DAG.getConstant(0, SDLoc(V), ToVT);
- if (ISD::isBuildVectorAllOnes(V.getNode()))
- return DAG.getAllOnesConstant(SDLoc(V), ToVT);
-
- SDLoc DL(V);
-
- // Logical operations (AND/OR/XOR): try picking the best fitting width out
- // of children's element widths.
- if (isLogicalMaskOp(Opcode)) {
- SDValue Op0 = convertMaskTree(V.getOperand(0), ToVT, Depth + 1);
- if (!Op0)
- return SDValue();
- SDValue Op1 = convertMaskTree(V.getOperand(1), ToVT, Depth + 1);
- if (!Op1)
- return SDValue();
- EVT OpVT = unifyMaskTypes(Op0, Op1, ToVT);
- return DAG.getNode(Opcode, DL, OpVT, Op0, Op1);
- }
-
- // FREEZE: widen the operand and re-wrap.
- if (Opcode == ISD::FREEZE) {
- SDValue Inner = convertMaskTree(V.getOperand(0), ToVT, Depth + 1);
- if (!Inner)
- return SDValue();
- return DAG.getNode(ISD::FREEZE, DL, Inner.getValueType(), Inner);
- }
-
- // Vector shuffle: try inferring the best fitting width from operands.
- if (Opcode == ISD::VECTOR_SHUFFLE) {
- // Bail out when the number of elements is different, we can't
- // simply reuse shuffle mask in this case.
- if (V.getValueType().getVectorNumElements() !=
- ToVT.getVectorNumElements())
- return SDValue();
-
- auto *Shuf = cast<ShuffleVectorSDNode>(V);
- SDValue Op0 = convertMaskTree(V.getOperand(0), ToVT, Depth + 1);
- if (!Op0)
- return SDValue();
- if (V.getOperand(1).isUndef()) {
- EVT OpVT = Op0.getValueType();
- return DAG.getVectorShuffle(OpVT, DL, Op0, DAG.getUNDEF(OpVT),
- Shuf->getMask());
- }
- SDValue Op1 = convertMaskTree(V.getOperand(1), ToVT, Depth + 1);
- if (!Op1)
- return SDValue();
- EVT OpVT = unifyMaskTypes(Op0, Op1, ToVT);
- return DAG.getVectorShuffle(OpVT, DL, Op0, Op1, Shuf->getMask());
- }
-
- // SELECT/VSELECT: try inferring the best fitting width from operands.
- if (Opcode == ISD::SELECT || Opcode == ISD::VSELECT) {
- SDValue Op1 = convertMaskTree(V.getOperand(1), ToVT, Depth + 1);
- if (!Op1)
- return SDValue();
- SDValue Op2 = convertMaskTree(V.getOperand(2), ToVT, Depth + 1);
- if (!Op2)
- return SDValue();
- EVT OpVT = unifyMaskTypes(Op1, Op2, ToVT);
-
- SDValue Cond = V.getOperand(0);
- if (Opcode == ISD::VSELECT) {
- Cond = convertMaskTree(Cond, ToVT, Depth + 1);
- if (!Cond)
- return SDValue();
- Cond = adjustMaskToType(Cond, OpVT);
- }
- return DAG.getNode(Opcode, DL, OpVT, Cond, Op1, Op2);
+std::pair<SDValue, bool>
+DAGTypeLegalizer::convertMaskTreeImpl(SDValue V, EVT ToVT, unsigned Depth) {
+ // The main idea is to recursively traverse VSELECT's mask that needs
+ // widening to see if we can avoid unnecessary casts. The problem usually
+ // stems from a simple fact that SETCC might naturally produce results
+ // not in i1 (as we model it in LLVM IR) and we can continue using that
+ // type until we have to switch it up. Another important aspect is that
+ // "all ones" and "all zeros" constants can be materialized at any type,
+ // so we can try to utilize that to keep SETCC results at their natural
+ // types as much as possible.
+ //
+ // The algorithm traverses the mask-producing tree of operations that
+ // retain "mask-vector"-ness of the input (i.e. it remains a vector of
+ // -1s and 0s).
+ //
+ // SETCC1 SETCC2 CONST1 SETCC3 CONST2 CONST3
+ // | / | / | /
+ // | / | / | /
+ // |_____/ |______/ |______/
+ // | * choose the | * choose SETCC3 | * keep it as final type
+ // | most fitting | type | but consider it subject to
+ // | type | / change
+ // | | /
+ // | |______________/
+ // | | * choose SETCC3 type
+ // | /
+ // | /
+ // |_____________/
+ // | * choose the most fitting type
+ // | and then cast to the final desired type
+ // |
+ // VSELECT
+ //
+ if (Depth >= DAG.MaxRecursionDepth ||
+ // Bail out when encounter the vector element count mismatch.
+ // It potentially can be just an assertion, but we deliberately try to
+ // be overly conservative here.
+ V.getValueType().getVectorNumElements() != ToVT.getVectorNumElements())
+ return {};
+
+ unsigned Opcode = V.getOpcode();
+
+ // Base case: SETCC produces the mask at its natural type.
+ if (isSETCCOp(Opcode)) {
+ EVT MaskVT = getSetCCResultType(getSETCCOperandType(V));
+ return {convertMask(V, MaskVT, MaskVT), /*IsTypeLenient=*/false};
+ }
+
+ // Base case: all-zeros or all-ones BUILD_VECTOR. Type-lenient since these are
+ // invariant under sign-extend/truncate.
+ if (ISD::isBuildVectorAllZeros(V.getNode()))
+ return {DAG.getConstant(0, SDLoc(V), ToVT), /*IsTypeLenient=*/true};
+ if (ISD::isBuildVectorAllOnes(V.getNode()))
+ return {DAG.getAllOnesConstant(SDLoc(V), ToVT), /*IsTypeLenient=*/true};
+
+ SDLoc DL(V);
+
+ // Logical operations (AND/OR/XOR): try picking the best fitting width out
+ // of children's element widths.
+ if (isLogicalMaskOp(Opcode)) {
+ auto [Op0, IsLenientOp0] =
+ convertMaskTreeImpl(V.getOperand(0), ToVT, Depth + 1);
+ if (!Op0)
+ return {};
+ auto [Op1, IsLenientOp1] =
+ convertMaskTreeImpl(V.getOperand(1), ToVT, Depth + 1);
+ if (!Op1)
+ return {};
+ EVT OpVT = unifyMaskTypes(Op0, IsLenientOp0, Op1, IsLenientOp1, ToVT);
+ return {DAG.getNode(Opcode, DL, OpVT, Op0, Op1),
+ IsLenientOp0 && IsLenientOp1};
+ }
+
+ // FREEZE: widen the operand and re-wrap.
+ if (Opcode == ISD::FREEZE) {
+ auto [Inner, IsTypeLenient] =
+ convertMaskTreeImpl(V.getOperand(0), ToVT, Depth + 1);
+ if (!Inner)
+ return {};
+ return {DAG.getNode(ISD::FREEZE, DL, Inner.getValueType(), Inner),
+ IsTypeLenient};
+ }
+
+ // Vector shuffle: try inferring the best fitting width from operands.
+ if (Opcode == ISD::VECTOR_SHUFFLE) {
+ auto *Shuf = cast<ShuffleVectorSDNode>(V);
+ auto [Op0, IsLenientOp0] =
+ convertMaskTreeImpl(V.getOperand(0), ToVT, Depth + 1);
+ if (!Op0)
+ return {};
+ if (V.getOperand(1).isUndef()) {
+ EVT OpVT = Op0.getValueType();
+ return {DAG.getVectorShuffle(OpVT, DL, Op0, DAG.getUNDEF(OpVT),
+ Shuf->getMask()),
+ IsLenientOp0};
}
-
- return SDValue();
- }();
-
+ auto [Op1, IsLenientOp1] =
+ convertMaskTreeImpl(V.getOperand(1), ToVT, Depth + 1);
+ if (!Op1)
+ return {};
+ EVT OpVT = unifyMaskTypes(Op0, IsLenientOp0, Op1, IsLenientOp1, ToVT);
+ return {DAG.getVectorShuffle(OpVT, DL, Op0, Op1, Shuf->getMask()),
+ IsLenientOp0 && IsLenientOp1};
+ }
+
+ // SELECT/VSELECT: try inferring the best fitting width from operands.
+ if (Opcode == ISD::SELECT || Opcode == ISD::VSELECT) {
+ auto [Op1, IsLenientOp1] =
+ convertMaskTreeImpl(V.getOperand(1), ToVT, Depth + 1);
+ if (!Op1)
+ return {};
+ auto [Op2, IsLenientOp2] =
+ convertMaskTreeImpl(V.getOperand(2), ToVT, Depth + 1);
+ if (!Op2)
+ return {};
+ EVT OpVT = unifyMaskTypes(Op1, IsLenientOp1, Op2, IsLenientOp2, ToVT);
+
+ // We deliberately skip traversing/modifying VSELECT's mask because
+ //
+ // a. We only change bitwidth of the operands and it shouldn't affect
+ // condition on its own.
+ //
+ // b. This VSELECT's mask can be widened in an independent traversal
+ // if needed.
+ SDValue Cond = V.getOperand(0);
+ return {DAG.getNode(Opcode, DL, OpVT, Cond, Op1, Op2),
+ IsLenientOp1 && IsLenientOp2};
+ }
+
+ return {};
+}
+
+SDValue DAGTypeLegalizer::convertMaskTree(SDValue V, EVT ToVT) {
+ // In general, we are converting from <N x i1> into <M x iW>.
+ // This would mean that during the tree traversal we need to pay
+ // attention to both bitwidth and element count, which can be error-prone.
+ //
+ // Instead, we split the task in two, we first widen the type of the tree
+ // and then change the element count.
+ EVT MaskTreeVT = ToVT.changeVectorElementCount(
+ *DAG.getContext(), V.getValueType().getVectorElementCount());
+ auto [Result, _] = convertMaskTreeImpl(V, MaskTreeVT);
if (!Result)
- return SDValue();
- if (Depth == 0)
- Result = adjustMaskToType(Result, ToVT);
- return Result;
+ return Result;
+ return adjustMaskToType(Result, ToVT);
}
// This method tries to handle some special cases for the vselect mask
diff --git a/llvm/test/CodeGen/X86/vselect-widen-mask-tree.ll b/llvm/test/CodeGen/X86/vselect-widen-mask-tree.ll
new file mode 100644
index 0000000000000..a83cb0aee00d8
--- /dev/null
+++ b/llvm/test/CodeGen/X86/vselect-widen-mask-tree.ll
@@ -0,0 +1,93 @@
+; NOTE: Assertions have been autogenerated by utils/update_llc_test_checks.py
+; RUN: llc -mtriple=x86_64-unknown-linux-gnu -mattr=+sse2 < %s | FileCheck %s
+
+define <2 x i32> @sel_allones_setcc(<2 x i32> %a, <2 x i32> %x, <2 x i32> %y, i1 %c) {
+; CHECK-LABEL: sel_allones_setcc:
+; CHECK: # %bb.0:
+; CHECK-NEXT: movdqa %xmm0, %xmm3
+; CHECK-NEXT: pcmpeqd %xmm0, %xmm0
+; CHECK-NEXT: testb $1, %dil
+; CHECK-NEXT: jne .LBB0_2
+; CHECK-NEXT: # %bb.1:
+; CHECK-NEXT: pxor %xmm0, %xmm0
+; CHECK-NEXT: pcmpgtd %xmm0, %xmm3
+; CHECK-NEXT: movdqa %xmm3, %xmm0
+; CHECK-NEXT: .LBB0_2:
+; CHECK-NEXT: pand %xmm0, %xmm1
+; CHECK-NEXT: pandn %xmm2, %xmm0
+; CHECK-NEXT: por %xmm1, %xmm0
+; CHECK-NEXT: retq
+ %cmp = icmp sgt <2 x i32> %a, zeroinitializer
+ %mask = select i1 %c, <2 x i1> <i1 true, i1 true>, <2 x i1> %cmp
+ %sel = select <2 x i1> %mask, <2 x i32> %x, <2 x i32> %y
+ ret <2 x i32> %sel
+}
+
+define <2 x i32> @sel_setcc_allzeros(<2 x i32> %a, <2 x i32> %x, <2 x i32> %y, i1 %c) {
+; CHECK-LABEL: sel_setcc_allzeros:
+; CHECK: # %bb.0:
+; CHECK-NEXT: movdqa %xmm0, %xmm3
+; CHECK-NEXT: pxor %xmm0, %xmm0
+; CHECK-NEXT: testb $1, %dil
+; CHECK-NEXT: je .LBB1_2
+; CHECK-NEXT: # %bb.1:
+; CHECK-NEXT: pcmpgtd %xmm0, %xmm3
+; CHECK-NEXT: movdqa %xmm3, %xmm0
+; CHECK-NEXT: .LBB1_2:
+; CHECK-NEXT: pand %xmm0, %xmm1
+; CHECK-NEXT: pandn %xmm2, %xmm0
+; CHECK-NEXT: por %xmm1, %xmm0
+; CHECK-NEXT: retq
+ %cmp = icmp sgt <2 x i32> %a, zeroinitializer
+ %mask = select i1 %c, <2 x i1> %cmp, <2 x i1> zeroinitializer
+ %sel = select <2 x i1> %mask, <2 x i32> %x, <2 x i32> %y
+ ret <2 x i32> %sel
+}
+
+define <4 x i8> @sel_allzeros_setcc_v4i8(<4 x i32> %a, <4 x i8> %x, <4 x i8> %y, i1 %c) {
+; CHECK-LABEL: sel_allzeros_setcc_v4i8:
+; CHECK: # %bb.0:
+; CHECK-NEXT: movdqa %xmm0, %xmm3
+; CHECK-NEXT: pxor %xmm0, %xmm0
+; CHECK-NEXT: testb $1, %dil
+; CHECK-NEXT: jne .LBB2_2
+; CHECK-NEXT: # %bb.1:
+; CHECK-NEXT: pcmpgtd %xmm0, %xmm3
+; CHECK-NEXT: movdqa %xmm3, %xmm0
+; CHECK-NEXT: .LBB2_2:
+; CHECK-NEXT: packssdw %xmm0, %xmm0
+; CHECK-NEXT: packsswb %xmm0, %xmm0
+; CHECK-NEXT: pand %xmm0, %xmm1
+; CHECK-NEXT: pandn %xmm2, %xmm0
+; CHECK-NEXT: por %xmm1, %xmm0
+; CHECK-NEXT: retq
+ %cmp = icmp sgt <4 x i32> %a, zeroinitializer
+ %mask = select i1 %c, <4 x i1> zeroinitializer, <4 x i1> %cmp
+ %sel = select <4 x i1> %mask, <4 x i8> %x, <4 x i8> %y
+ ret <4 x i8> %sel
+}
+
+define <2 x i32> @sel_or_setcc_allones(<2 x i32> %a, <2 x i32> %b, <2 x i32> %x, <2 x i32> %y, i1 %c) {
+; CHECK-LABEL: sel_or_setcc_allones:
+; CHECK: # %bb.0:
+; CHECK-NEXT: movdqa %xmm0, %xmm4
+; CHECK-NEXT: pcmpeqd %xmm0, %xmm0
+; CHECK-NEXT: testb $1, %dil
+; CHECK-NEXT: jne .LBB3_2
+; CHECK-NEXT: # %bb.1:
+; CHECK-NEXT: pxor %xmm0, %xmm0
+; CHECK-NEXT: pcmpgtd %xmm0, %xmm4
+; CHECK-NEXT: pcmpgtd %xmm1, %xmm0
+; CHECK-NEXT: por %xmm4, %xmm0
+; CHECK-NEXT: .LBB3_2:
+; CHECK-NEXT: pand %xmm0, %xmm2
+; CHECK-NEXT: pandn %xmm3, %xmm0
+; CHECK-NEXT: por %xmm2, %xmm0
+; CHECK-NEXT: retq
+ %cmp0 = icmp sgt <2 x i32> %a, zeroinitializer
+ %cmp1 = icmp slt <2 x i32> %b, zeroinitializer
+ %or = or <2 x i1> %cmp0, %cmp1
+ %mask = select i1 %c, <2 x i1> <i1 true, i1 true>, <2 x i1> %or
+ %sel = select <2 x i1> %mask, <2 x i32> %x, <2 x i32> %y
+ ret <2 x i32> %sel
+}
More information about the llvm-commits
mailing list