[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