[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