[llvm] [X86] getMaskNode - perform pre-truncation of oversized scalar mask sources (PR #187063)
Simon Pilgrim via llvm-commits
llvm-commits at lists.llvm.org
Tue Mar 17 10:05:26 PDT 2026
https://github.com/RKSimon created https://github.com/llvm/llvm-project/pull/187063
Allows us to use getMaskNode to canonicalize predicate masks in big shift lowering
>From 4756ad03506322648f0a5dbce037c55ed02de52c Mon Sep 17 00:00:00 2001
From: Simon Pilgrim <llvm-dev at redking.me.uk>
Date: Tue, 17 Mar 2026 16:28:03 +0000
Subject: [PATCH] [X86] getMaskNode - perform pre-truncation of oversized
scalar mask sources
Allows us to use getMaskNode to canonicalize predicate masks in big shift lowering
---
llvm/lib/Target/X86/X86ISelLowering.cpp | 22 ++++++++++++----------
1 file changed, 12 insertions(+), 10 deletions(-)
diff --git a/llvm/lib/Target/X86/X86ISelLowering.cpp b/llvm/lib/Target/X86/X86ISelLowering.cpp
index 165db8b760be5..c4fdfd2ecc9d2 100644
--- a/llvm/lib/Target/X86/X86ISelLowering.cpp
+++ b/llvm/lib/Target/X86/X86ISelLowering.cpp
@@ -10939,7 +10939,7 @@ static SDValue getMaskNode(SDValue Mask, MVT MaskVT,
const SDLoc &dl) {
MVT SrcVT = Mask.getSimpleValueType();
assert(SrcVT.isScalarInteger() && "Expected scalar integer mask source!");
- assert(MaskVT.bitsLE(Mask.getSimpleValueType()) && "Unexpected mask size!");
+ assert(MaskVT.bitsLE(SrcVT) && "Unexpected mask size!");
assert(MaskVT.getVectorElementType() == MVT::i1 && "Bool vector expected!");
if (isAllOnesConstant(Mask))
@@ -10947,6 +10947,12 @@ static SDValue getMaskNode(SDValue Mask, MVT MaskVT,
if (X86::isZeroNode(Mask))
return DAG.getConstant(0, dl, MaskVT);
+ // Attempt to pre-truncate the mask source (to a minimum of i8).
+ if (SrcVT.getSizeInBits() > MaskVT.getVectorNumElements()) {
+ SrcVT = MVT::getIntegerVT(std::max((int)MaskVT.getVectorNumElements(), 8));
+ Mask = DAG.getNode(ISD::TRUNCATE, dl, SrcVT, Mask);
+ }
+
if (SrcVT == MVT::i64 && Subtarget.is32Bit()) {
assert(MaskVT == MVT::v64i1 && "Expected v64i1 mask!");
assert(Subtarget.hasBWI() && "Expected AVX512BW target!");
@@ -34508,15 +34514,13 @@ void X86TargetLowering::ReplaceNodeResults(SDNode *N,
SDValue AllBitsMask =
DAG.getNode(Opc, dl, MVT::i64,
DAG.getConstant(AllBitsVal, dl, MVT::i64), AmtLane);
- AllBitsMask = DAG.getBitcast(
- BoolVT, DAG.getZExtOrTrunc(AllBitsMask, dl, MVT::i8));
+ AllBitsMask = getMaskNode(AllBitsMask, BoolVT, Subtarget, DAG, dl);
Res = DAG.getSelect(dl, VecVT, AllBitsMask,
DAG.getAllOnesConstant(dl, VecVT), Res);
}
SDValue LaneMask = DAG.getNode(Opc, dl, MVT::i64, LaneBit, AmtLane);
- LaneMask =
- DAG.getBitcast(BoolVT, DAG.getZExtOrTrunc(LaneMask, dl, MVT::i8));
+ LaneMask = getMaskNode(LaneMask, BoolVT, Subtarget, DAG, dl);
SDValue Elt = DAG.getNode(Opc, dl, MVT::i64, EltBit, AmtMod);
Res = DAG.getSelect(dl, VecVT, LaneMask, DAG.getSplat(VecVT, dl, Elt),
Res);
@@ -34529,11 +34533,9 @@ void X86TargetLowering::ReplaceNodeResults(SDNode *N,
// ShiftAmt/64 'laneshift', and then shuffle one element along to get the
// shifted in bits from the neighbouring element. Finally use a funnel shift
// with the ShiftAmt%64 'elementshift' to get the final result.
- SDValue Mask =
- DAG.getNode(ISD::TRUNCATE, dl, MVT::i8,
- DAG.getNode(ISD::SHL, dl, MVT::i32,
- DAG.getAllOnesConstant(dl, MVT::i32), AmtLane));
- Mask = DAG.getBitcast(BoolVT, Mask);
+ SDValue Mask = DAG.getNode(ISD::SHL, dl, MVT::i32,
+ DAG.getAllOnesConstant(dl, MVT::i32), AmtLane);
+ Mask = getMaskNode(Mask, BoolVT, Subtarget, DAG, dl);
Src = DAG.getBitcast(VecVT, Src);
SmallVector<int, 8> ShufMask(NumElts);
More information about the llvm-commits
mailing list