[llvm] [VectorCombine] Fold interleave and widen chained operations (PR #224005)
Kamlesh Kumar via llvm-commits
llvm-commits at lists.llvm.org
Tue Sep 22 01:21:10 PDT 2026
================
@@ -6032,233 +6032,242 @@ bool VectorCombine::foldInsExtVectorToShuffle(Instruction &I) {
return true;
}
-/// Fold away a matched pair of vector.deinterleave/interleave intrinsics
-/// with a chain of elementwise operations on each between the
-/// deinterleave and interleave.
-///
-/// For example:
-/// ```
-/// %d = call { <2 x i16>, <2 x i16> } @deinterleave2.v4i16(<4 x i16> %v)
-/// %f0 = extractvalue { <2 x i16>, <2 x i16> } %d, 0
-/// %f1 = extractvalue { <2 x i16>, <2 x i16> } %d, 1
-///
-/// %u0 = add <2 x i16> %f0, splat (i16 3)
-/// %u1 = add <2 x i16> %f1, splat (i16 3)
-///
-/// %r = call <4 x i16> @interleave2.v4i16(<2 x i16> %u0, <2 x i16> %u1)
-/// ```
-/// Folds to:
-/// ```
-/// %r = add <4 x i16> %v, splat (i16 3)
-/// ```
-bool VectorCombine::foldDeinterleaveInterleavePair(Instruction &I) {
- auto *Deinterleave = dyn_cast<IntrinsicInst>(&I);
- if (!Deinterleave)
- return false;
+static unsigned getNumDataOperands(const Instruction *Inst) {
+ if (auto *CB = dyn_cast<CallBase>(Inst))
+ return CB->arg_size(); // Exclude callee operand and bundles.
+ return Inst->getNumOperands();
+}
- unsigned Factor =
- getDeinterleaveIntrinsicFactor(Deinterleave->getIntrinsicID());
- if (!Factor || Deinterleave->hasOperandBundles() ||
- !Deinterleave->hasNUndroppableUses(Factor))
+/// Return true if \p Inst is an elementwise operation that can be rebuilt at a
+/// wider element count.
+static bool isSupportedElementwise(Instruction *Inst) {
+ auto *ResultTy = dyn_cast<VectorType>(Inst->getType());
+ if (!ResultTy || !isSafeToSpeculativelyExecute(Inst))
return false;
- const Intrinsic::ID ExpectedInterleaveIID =
- Intrinsic::getInterleaveIntrinsicID(Factor);
-
- // Collect one extract for each deinterleaved field.
- SmallVector<Use *, 8> CurrentUses(Factor, nullptr);
- for (Use &U : Deinterleave->uses()) {
- if (U.getUser()->isDroppable())
- continue;
-
- auto *Extract = dyn_cast<ExtractValueInst>(U.getUser());
- if (!Extract || Extract->getNumIndices() != 1)
+ if (auto *II = dyn_cast<IntrinsicInst>(Inst)) {
+ if (II->hasOperandBundles() ||
+ !isTriviallyVectorizable(II->getIntrinsicID()))
return false;
+ } else if (!isa<BinaryOperator, UnaryOperator, CastInst, CmpInst, SelectInst,
+ FreezeInst>(Inst)) {
+ return false;
+ }
- unsigned Index = *Extract->idx_begin();
- if (Index >= Factor || CurrentUses[Index])
+ // Reject operations that change the element-count.
+ // E.g., bitcast <vscale x 4 x i16> %v to <vscale x 8 x i8>
+ for (unsigned Op = 0, E = getNumDataOperands(Inst); Op != E; ++Op) {
+ auto *OperandTy = dyn_cast<VectorType>(Inst->getOperand(Op)->getType());
+ if (OperandTy &&
+ OperandTy->getElementCount() != ResultTy->getElementCount())
return false;
-
- CurrentUses[Index] = &U;
}
- using ElementwiseStep = SmallVector<Use *, 8>;
- SmallVector<ElementwiseStep, 4> Steps;
- IntrinsicInst *Interleave = nullptr;
- unsigned NumVisited = 0;
+ return true;
+}
- auto GetNumDataOperands = [](Instruction *Inst) {
- if (auto *CB = dyn_cast<CallBase>(Inst))
- return CB->arg_size(); // Exclude callee operand and bundles.
- return Inst->getNumOperands();
+static Value *getCommonSplatValue(ArrayRef<Value *> Values) {
+ auto GetSplatOrScalar = [](Value *V) {
+ return isa<VectorType>(V->getType()) ? getSplatValue(V) : V;
};
- auto IsSupportedElementwise = [&](Instruction *Inst) {
- auto *ResultTy = dyn_cast<VectorType>(Inst->getType());
- if (!ResultTy || !isSafeToSpeculativelyExecute(Inst))
- return false;
+ Value *CommonValue = GetSplatOrScalar(Values.front());
+ if (!CommonValue)
+ return nullptr;
+ for (Value *V : Values.drop_front())
+ if (GetSplatOrScalar(V) != CommonValue)
+ return nullptr;
+ return CommonValue;
+}
- if (auto *II = dyn_cast<IntrinsicInst>(Inst)) {
- if (II->hasOperandBundles() ||
- !isTriviallyVectorizable(II->getIntrinsicID()))
- return false;
- } else if (!isa<BinaryOperator, UnaryOperator, CastInst, CmpInst,
- SelectInst, FreezeInst>(Inst)) {
- return false;
- }
+/// Return the common deinterleave intrinsic if \p Members are its extracts in
+/// field order.
+static IntrinsicInst *getDeinterleaveForMembers(ArrayRef<Value *> Members,
+ unsigned Factor) {
+ IntrinsicInst *Deinterleave = nullptr;
+ for (const auto &[Index, Member] : enumerate(Members)) {
+ auto *Extract = dyn_cast<ExtractValueInst>(Member);
+ if (!Extract || Extract->getNumIndices() != 1 ||
+ *Extract->idx_begin() != Index)
+ return nullptr;
- // Reject operations that change the element-count.
- // E.g., bitcast <vscale x 4 x i16> %v to <vscale x 8 x i8>
- for (unsigned Op = 0, E = GetNumDataOperands(Inst); Op != E; ++Op) {
- auto *OperandTy = dyn_cast<VectorType>(Inst->getOperand(Op)->getType());
- if (OperandTy &&
- OperandTy->getElementCount() != ResultTy->getElementCount())
- return false;
- }
+ auto *Current = dyn_cast<IntrinsicInst>(Extract->getAggregateOperand());
+ if (!Current || Current->hasOperandBundles() ||
+ getDeinterleaveIntrinsicFactor(Current->getIntrinsicID()) != Factor ||
+ (Deinterleave && Current != Deinterleave))
+ return nullptr;
+ Deinterleave = Current;
+ }
- return true;
- };
+ if (!Deinterleave || !Deinterleave->hasNUndroppableUses(Factor))
+ return nullptr;
+ return Deinterleave;
+}
- // Traverse the Factor use chains with a breadth-first search.
- // At each level, expect every chain to perform the same operation with the
- // preceding chain value at the same operand position, until they all reach
- // the matching interleave.
- while (NumVisited + Factor <= MaxInstrsToScan) {
- NumVisited += Factor;
-
- for (Use *&CurrentUse : CurrentUses) {
- Use *NextUse = CurrentUse->getUser()->getSingleUndroppableUse();
- auto *Next =
- NextUse ? dyn_cast<Instruction>(NextUse->getUser()) : nullptr;
- if (!Next)
- return false;
+static SmallVector<Value *, 8> getMemberOperands(ArrayRef<Value *> Members,
+ unsigned OperandIndex) {
+ SmallVector<Value *, 8> Operands;
+ for (Value *Member : Members)
+ Operands.push_back(cast<Instruction>(Member)->getOperand(OperandIndex));
+ return Operands;
+}
- CurrentUse = NextUse;
- }
+/// Check whether the tree of elementwise operations each feeding \p Members
+/// can be rebuilt at the interleaved width.
+static bool canWidenOperations(ArrayRef<Value *> Members, unsigned Factor,
+ unsigned &NumScanned) {
+ assert(Members.size() == Factor && "expected one member per field");
+ if (getDeinterleaveForMembers(Members, Factor))
+ return true;
+ if (NumScanned + Factor > MaxInstrsToScan)
+ return false;
+ NumScanned += Factor;
- // Check whether every chain has reached the same interleave.
- if (auto *II = dyn_cast<IntrinsicInst>(CurrentUses.front()->getUser());
- II && II->getIntrinsicID() == ExpectedInterleaveIID) {
- if (II->hasOperandBundles())
- return false;
+ auto *FirstInst = dyn_cast<Instruction>(Members.front());
+ if (!FirstInst || !isSupportedElementwise(FirstInst) ||
+ !FirstInst->getSingleUndroppableUse())
+ return false;
- for (unsigned Index = 0; Index != Factor; ++Index)
- if (CurrentUses[Index]->getUser() != II ||
- CurrentUses[Index]->getOperandNo() != Index)
- return false;
+ for (Value *Member : Members.drop_front()) {
+ auto *Inst = dyn_cast<Instruction>(Member);
+ if (!Inst || !isSupportedElementwise(Inst) ||
+ !Inst->getSingleUndroppableUse() ||
+ !FirstInst->isSameOperationAs(Inst, Instruction::CompareCallTargets))
+ return false;
+ }
- Interleave = II;
- break;
+ for (unsigned Op = 0, E = getNumDataOperands(FirstInst); Op != E; ++Op) {
+ SmallVector<Value *, 8> Operands = getMemberOperands(Members, Op);
+ if (!isa<VectorType>(Operands.front()->getType())) {
+ if (!all_equal(Operands))
+ return false;
+ continue;
}
- auto *FirstInst = cast<Instruction>(CurrentUses.front()->getUser());
- if (!IsSupportedElementwise(FirstInst))
+ if (!getCommonSplatValue(Operands) &&
+ !canWidenOperations(Operands, Factor, NumScanned))
return false;
+ }
+ return true;
+}
- unsigned ChainOperand = CurrentUses.front()->getOperandNo();
- bool MismatchedUse = any_of(CurrentUses, [&](Use *U) {
- auto *Inst = cast<Instruction>(U->getUser());
- return Inst != FirstInst && (U->getOperandNo() != ChainOperand ||
- !FirstInst->isSameOperationAs(
- Inst, Instruction::CompareCallTargets));
- });
- if (MismatchedUse)
- return false;
-
- auto GetSplatOrScalar = [](Value *V) {
- return isa<VectorType>(V->getType()) ? getSplatValue(V) : V;
- };
-
- // Non-chain operands must be either the same scalar or splats of that
- // scalar. This intentionally rejects differing poison/undef or non-splat
- // vector operands between chains.
- for (unsigned Op = 0, E = GetNumDataOperands(FirstInst); Op != E; ++Op) {
- if (Op == ChainOperand)
- continue;
+static Value *createWideInstruction(Instruction *NarrowInst,
+ ArrayRef<Value *> NewOperands,
+ VectorType *WideResultTy,
+ IRBuilder<InstSimplifyFolder> &Builder) {
+ if (isa<BinaryOperator, UnaryOperator>(NarrowInst))
+ return Builder.CreateNAryOp(NarrowInst->getOpcode(), NewOperands);
+ if (auto *Cast = dyn_cast<CastInst>(NarrowInst))
+ return Builder.CreateCast(Cast->getOpcode(), NewOperands[0], WideResultTy);
+ if (auto *Cmp = dyn_cast<CmpInst>(NarrowInst))
+ return Builder.CreateCmp(Cmp->getPredicate(), NewOperands[0],
+ NewOperands[1]);
+ if (isa<SelectInst>(NarrowInst))
+ return Builder.CreateSelect(
+ NewOperands[0], NewOperands[1], NewOperands[2], /*Name=*/"",
+ ProfcheckDisableMetadataFixes ? nullptr : NarrowInst);
+ if (isa<FreezeInst>(NarrowInst))
+ return Builder.CreateFreeze(NewOperands[0]);
+ if (auto *II = dyn_cast<IntrinsicInst>(NarrowInst))
+ return Builder.CreateIntrinsic(WideResultTy, II->getIntrinsicID(),
+ NewOperands);
+ llvm_unreachable("Unsupported instruction");
+}
- Value *CommonValue = GetSplatOrScalar(FirstInst->getOperand(Op));
- if (!CommonValue || any_of(CurrentUses, [&](Use *U) {
- Instruction *Inst = cast<Instruction>(U->getUser());
- return Inst != FirstInst &&
- GetSplatOrScalar(Inst->getOperand(Op)) != CommonValue;
- }))
- return false;
+static Value *widenOperations(ArrayRef<Value *> Members, unsigned Factor,
+ ElementCount WideEC,
+ IRBuilder<InstSimplifyFolder> &Builder) {
+ if (auto *Deinterleave = getDeinterleaveForMembers(Members, Factor)) {
+ Value *Source = Deinterleave->getArgOperand(0);
+ assert(cast<VectorType>(Source->getType())->getElementCount() == WideEC &&
+ "deinterleave source must have the interleaved element count");
+ return Source;
+ }
+
+ auto *NarrowInst = cast<Instruction>(Members.front());
+ unsigned NumOperands = getNumDataOperands(NarrowInst);
+ SmallVector<Value *, 4> NewOperands;
+ NewOperands.reserve(NumOperands);
+ for (unsigned Op = 0; Op != NumOperands; ++Op) {
+ SmallVector<Value *, 8> Operands = getMemberOperands(Members, Op);
+ Value *NewOperand = Operands.front();
+ if (isa<VectorType>(NewOperand->getType())) {
+ if (Value *CommonValue = getCommonSplatValue(Operands)) {
+ Builder.SetCurrentDebugLocation(NarrowInst->getDebugLoc());
+ NewOperand = Builder.CreateVectorSplat(WideEC, CommonValue);
+ } else {
+ NewOperand = widenOperations(Operands, Factor, WideEC, Builder);
+ }
----------------
kamleshbhalui wrote:
I do not think it can be tested, as canwidenoperation already checks and bails.
https://github.com/llvm/llvm-project/pull/224005
More information about the llvm-commits
mailing list