[llvm] [VectorCombine] Fold insertelement chains of scalar parts to a bitcast and shuffle (PR #226224)
via llvm-commits
llvm-commits at lists.llvm.org
Sun Oct 4 01:34:13 PDT 2026
================
@@ -6063,6 +6064,118 @@ bool VectorCombine::foldInsExtVectorToShuffle(Instruction &I) {
return true;
}
+/// Try to replace a chain of insertelements of parts of the same scalar with a
+/// bitcast and a shuffle (little endian):
+/// insert (insert poison, (trunc (lshr X, 32)), 0), (trunc X), 1 -->
+/// shuffle (bitcast X to <2 x i32>), poison, <1, 0>
+bool VectorCombine::foldInsertScalarPartsToShuffle(Instruction &I) {
+ auto *VecTy = dyn_cast<FixedVectorType>(I.getType());
+ if (!VecTy)
+ return false;
+
+ // Start from the last insertelement of the chain.
+ if (I.hasOneUse() && isa<InsertElementInst>(I.user_back()))
+ return false;
+
+ Type *EltTy = VecTy->getElementType();
+ if ((!EltTy->isIntegerTy() && !EltTy->isIEEELikeFPTy()) ||
+ !DL->typeSizeEqualsStoreSize(EltTy))
+ return false;
+ unsigned EltBits = EltTy->getPrimitiveSizeInBits();
+ unsigned NumElts = VecTy->getNumElements();
+
+ Value *Src = nullptr;
+ unsigned NumSrcElts = 0;
+ SmallVector<int> Mask(NumElts, PoisonMaskElem);
+ APInt DemandedElts = APInt::getZero(NumElts);
+ InstructionCost OldCost = 0;
+ Value *Vec = &I;
+ while (auto *Ins = dyn_cast<InsertElementInst>(Vec)) {
+ if (Ins != &I && !Ins->hasOneUse())
+ return false;
+ uint64_t Idx;
+ if (!match(Ins->getOperand(2), m_ConstantInt(Idx)) || Idx >= NumElts)
+ return false;
+ Vec = Ins->getOperand(0);
+ // A later insert to the same element overrides this one.
+ if (DemandedElts[Idx])
+ continue;
+ DemandedElts.setBit(Idx);
+
+ // Match (bitcast (trunc (lshr X, ShAmt))), the bitcast and shift being
+ // optional.
+ Value *Elt = Ins->getOperand(1);
+ Value *Trunc = Elt;
+ match(Trunc, m_BitCast(m_Value(Trunc)));
+ Value *X;
+ if (!match(Trunc, m_Trunc(m_Value(X))) || !X->getType()->isIntegerTy() ||
+ Trunc->getType()->getPrimitiveSizeInBits() != EltBits)
+ return false;
+ Value *Shift = nullptr;
+ uint64_t ShAmt = 0;
+ if (match(X, m_LShr(m_Value(), m_ConstantInt(ShAmt)))) {
+ Shift = X;
+ X = cast<Instruction>(Shift)->getOperand(0);
+ }
+
+ if (!Src) {
+ unsigned SrcBits = X->getType()->getIntegerBitWidth();
+ if (SrcBits % EltBits)
+ return false;
+ Src = X;
+ NumSrcElts = SrcBits / EltBits;
+ } else if (X != Src) {
+ return false;
+ }
+ if (ShAmt % EltBits || ShAmt / EltBits >= NumSrcElts)
----------------
ParkHanbum wrote:
Nit: I think this would be slightly clearer as
uint64_t PartIndex = ShAmt / EltBits;
if (ShAmt % EltBits || PartIndex >= NumSrcElts)
return false;
Mask[Idx] =
DL->isBigEndian() ? NumSrcElts - 1 - PartIndex : PartIndex;
https://github.com/llvm/llvm-project/pull/226224
More information about the llvm-commits
mailing list