diff options
Diffstat (limited to 'llvm/lib/Transforms/Vectorize/VectorCombine.cpp')
| -rw-r--r-- | llvm/lib/Transforms/Vectorize/VectorCombine.cpp | 267 |
1 files changed, 232 insertions, 35 deletions
diff --git a/llvm/lib/Transforms/Vectorize/VectorCombine.cpp b/llvm/lib/Transforms/Vectorize/VectorCombine.cpp index 787f146bdddc..d18bcd34620c 100644 --- a/llvm/lib/Transforms/Vectorize/VectorCombine.cpp +++ b/llvm/lib/Transforms/Vectorize/VectorCombine.cpp @@ -14,6 +14,7 @@ #include "llvm/Transforms/Vectorize/VectorCombine.h" #include "llvm/ADT/Statistic.h" +#include "llvm/Analysis/AssumptionCache.h" #include "llvm/Analysis/BasicAliasAnalysis.h" #include "llvm/Analysis/GlobalsModRef.h" #include "llvm/Analysis/Loads.h" @@ -50,14 +51,18 @@ static cl::opt<bool> DisableBinopExtractShuffle( "disable-binop-extract-shuffle", cl::init(false), cl::Hidden, cl::desc("Disable binop extract to shuffle transforms")); +static cl::opt<unsigned> MaxInstrsToScan( + "vector-combine-max-scan-instrs", cl::init(30), cl::Hidden, + cl::desc("Max number of instructions to scan for vector combining.")); + static const unsigned InvalidIndex = std::numeric_limits<unsigned>::max(); namespace { class VectorCombine { public: VectorCombine(Function &F, const TargetTransformInfo &TTI, - const DominatorTree &DT) - : F(F), Builder(F.getContext()), TTI(TTI), DT(DT) {} + const DominatorTree &DT, AAResults &AA, AssumptionCache &AC) + : F(F), Builder(F.getContext()), TTI(TTI), DT(DT), AA(AA), AC(AC) {} bool run(); @@ -66,6 +71,8 @@ private: IRBuilder<> Builder; const TargetTransformInfo &TTI; const DominatorTree &DT; + AAResults &AA; + AssumptionCache &AC; bool vectorizeLoadInsert(Instruction &I); ExtractElementInst *getShuffleExtract(ExtractElementInst *Ext0, @@ -83,6 +90,8 @@ private: bool foldBitcastShuf(Instruction &I); bool scalarizeBinopOrCmp(Instruction &I); bool foldExtractedCmps(Instruction &I); + bool foldSingleElementStore(Instruction &I); + bool scalarizeLoadExtract(Instruction &I); }; } // namespace @@ -192,8 +201,18 @@ bool VectorCombine::vectorizeLoadInsert(Instruction &I) { InstructionCost NewCost = TTI.getMemoryOpCost(Instruction::Load, MinVecTy, Alignment, AS); // Optionally, we are shuffling the loaded vector element(s) into place. + // For the mask set everything but element 0 to undef to prevent poison from + // propagating from the extra loaded memory. This will also optionally + // shrink/grow the vector from the loaded size to the output size. + // We assume this operation has no cost in codegen if there was no offset. + // Note that we could use freeze to avoid poison problems, but then we might + // still need a shuffle to change the vector size. + unsigned OutputNumElts = Ty->getNumElements(); + SmallVector<int, 16> Mask(OutputNumElts, UndefMaskElem); + assert(OffsetEltIndex < MinVecNumElts && "Address offset too big"); + Mask[0] = OffsetEltIndex; if (OffsetEltIndex) - NewCost += TTI.getShuffleCost(TTI::SK_PermuteSingleSrc, MinVecTy); + NewCost += TTI.getShuffleCost(TTI::SK_PermuteSingleSrc, MinVecTy, Mask); // We can aggressively convert to the vector form because the backend can // invert this transform if it does not result in a performance win. @@ -205,17 +224,6 @@ bool VectorCombine::vectorizeLoadInsert(Instruction &I) { IRBuilder<> Builder(Load); Value *CastedPtr = Builder.CreateBitCast(SrcPtr, MinVecTy->getPointerTo(AS)); Value *VecLd = Builder.CreateAlignedLoad(MinVecTy, CastedPtr, Alignment); - - // Set everything but element 0 to undef to prevent poison from propagating - // from the extra loaded memory. This will also optionally shrink/grow the - // vector from the loaded size to the output size. - // We assume this operation has no cost in codegen if there was no offset. - // Note that we could use freeze to avoid poison problems, but then we might - // still need a shuffle to change the vector size. - unsigned OutputNumElts = Ty->getNumElements(); - SmallVector<int, 16> Mask(OutputNumElts, UndefMaskElem); - assert(OffsetEltIndex < MinVecNumElts && "Address offset too big"); - Mask[0] = OffsetEltIndex; VecLd = Builder.CreateShuffleVector(VecLd, Mask); replaceValue(I, *VecLd); @@ -516,15 +524,6 @@ bool VectorCombine::foldBitcastShuf(Instruction &I) { if (!SrcTy || !DestTy || I.getOperand(0)->getType() != SrcTy) return false; - // The new shuffle must not cost more than the old shuffle. The bitcast is - // moved ahead of the shuffle, so assume that it has the same cost as before. - InstructionCost DestCost = - TTI.getShuffleCost(TargetTransformInfo::SK_PermuteSingleSrc, DestTy); - InstructionCost SrcCost = - TTI.getShuffleCost(TargetTransformInfo::SK_PermuteSingleSrc, SrcTy); - if (DestCost > SrcCost || !DestCost.isValid()) - return false; - unsigned DestNumElts = DestTy->getNumElements(); unsigned SrcNumElts = SrcTy->getNumElements(); SmallVector<int, 16> NewMask; @@ -542,6 +541,16 @@ bool VectorCombine::foldBitcastShuf(Instruction &I) { if (!widenShuffleMaskElts(ScaleFactor, Mask, NewMask)) return false; } + + // The new shuffle must not cost more than the old shuffle. The bitcast is + // moved ahead of the shuffle, so assume that it has the same cost as before. + InstructionCost DestCost = TTI.getShuffleCost( + TargetTransformInfo::SK_PermuteSingleSrc, DestTy, NewMask); + InstructionCost SrcCost = + TTI.getShuffleCost(TargetTransformInfo::SK_PermuteSingleSrc, SrcTy, Mask); + if (DestCost > SrcCost || !DestCost.isValid()) + return false; + // bitcast (shuf V, MaskC) --> shuf (bitcast V), MaskC' ++NumShufOfBitcast; Value *CastV = Builder.CreateBitCast(V, DestTy); @@ -725,8 +734,10 @@ bool VectorCombine::foldExtractedCmps(Instruction &I) { int ExpensiveIndex = ConvertToShuf == Ext0 ? Index0 : Index1; auto *CmpTy = cast<FixedVectorType>(CmpInst::makeCmpResultType(X->getType())); InstructionCost NewCost = TTI.getCmpSelInstrCost(CmpOpcode, X->getType()); - NewCost += - TTI.getShuffleCost(TargetTransformInfo::SK_PermuteSingleSrc, CmpTy); + SmallVector<int, 32> ShufMask(VecTy->getNumElements(), UndefMaskElem); + ShufMask[CheapIndex] = ExpensiveIndex; + NewCost += TTI.getShuffleCost(TargetTransformInfo::SK_PermuteSingleSrc, CmpTy, + ShufMask); NewCost += TTI.getArithmeticInstrCost(I.getOpcode(), CmpTy); NewCost += TTI.getVectorInstrCost(Ext0->getOpcode(), CmpTy, CheapIndex); @@ -752,6 +763,189 @@ bool VectorCombine::foldExtractedCmps(Instruction &I) { return true; } +// Check if memory loc modified between two instrs in the same BB +static bool isMemModifiedBetween(BasicBlock::iterator Begin, + BasicBlock::iterator End, + const MemoryLocation &Loc, AAResults &AA) { + unsigned NumScanned = 0; + return std::any_of(Begin, End, [&](const Instruction &Instr) { + return isModSet(AA.getModRefInfo(&Instr, Loc)) || + ++NumScanned > MaxInstrsToScan; + }); +} + +/// Check if it is legal to scalarize a memory access to \p VecTy at index \p +/// Idx. \p Idx must access a valid vector element. +static bool canScalarizeAccess(FixedVectorType *VecTy, Value *Idx, + Instruction *CtxI, AssumptionCache &AC) { + if (auto *C = dyn_cast<ConstantInt>(Idx)) + return C->getValue().ult(VecTy->getNumElements()); + + APInt Zero(Idx->getType()->getScalarSizeInBits(), 0); + APInt MaxElts(Idx->getType()->getScalarSizeInBits(), VecTy->getNumElements()); + ConstantRange ValidIndices(Zero, MaxElts); + ConstantRange IdxRange = computeConstantRange(Idx, true, &AC, CtxI, 0); + return ValidIndices.contains(IdxRange); +} + +/// The memory operation on a vector of \p ScalarType had alignment of +/// \p VectorAlignment. Compute the maximal, but conservatively correct, +/// alignment that will be valid for the memory operation on a single scalar +/// element of the same type with index \p Idx. +static Align computeAlignmentAfterScalarization(Align VectorAlignment, + Type *ScalarType, Value *Idx, + const DataLayout &DL) { + if (auto *C = dyn_cast<ConstantInt>(Idx)) + return commonAlignment(VectorAlignment, + C->getZExtValue() * DL.getTypeStoreSize(ScalarType)); + return commonAlignment(VectorAlignment, DL.getTypeStoreSize(ScalarType)); +} + +// Combine patterns like: +// %0 = load <4 x i32>, <4 x i32>* %a +// %1 = insertelement <4 x i32> %0, i32 %b, i32 1 +// store <4 x i32> %1, <4 x i32>* %a +// to: +// %0 = bitcast <4 x i32>* %a to i32* +// %1 = getelementptr inbounds i32, i32* %0, i64 0, i64 1 +// store i32 %b, i32* %1 +bool VectorCombine::foldSingleElementStore(Instruction &I) { + StoreInst *SI = dyn_cast<StoreInst>(&I); + if (!SI || !SI->isSimple() || + !isa<FixedVectorType>(SI->getValueOperand()->getType())) + return false; + + // TODO: Combine more complicated patterns (multiple insert) by referencing + // TargetTransformInfo. + Instruction *Source; + Value *NewElement; + Value *Idx; + if (!match(SI->getValueOperand(), + m_InsertElt(m_Instruction(Source), m_Value(NewElement), + m_Value(Idx)))) + return false; + + if (auto *Load = dyn_cast<LoadInst>(Source)) { + auto VecTy = cast<FixedVectorType>(SI->getValueOperand()->getType()); + const DataLayout &DL = I.getModule()->getDataLayout(); + Value *SrcAddr = Load->getPointerOperand()->stripPointerCasts(); + // Don't optimize for atomic/volatile load or store. Ensure memory is not + // modified between, vector type matches store size, and index is inbounds. + if (!Load->isSimple() || Load->getParent() != SI->getParent() || + !DL.typeSizeEqualsStoreSize(Load->getType()) || + !canScalarizeAccess(VecTy, Idx, Load, AC) || + SrcAddr != SI->getPointerOperand()->stripPointerCasts() || + isMemModifiedBetween(Load->getIterator(), SI->getIterator(), + MemoryLocation::get(SI), AA)) + return false; + + Value *GEP = Builder.CreateInBoundsGEP( + SI->getValueOperand()->getType(), SI->getPointerOperand(), + {ConstantInt::get(Idx->getType(), 0), Idx}); + StoreInst *NSI = Builder.CreateStore(NewElement, GEP); + NSI->copyMetadata(*SI); + Align ScalarOpAlignment = computeAlignmentAfterScalarization( + std::max(SI->getAlign(), Load->getAlign()), NewElement->getType(), Idx, + DL); + NSI->setAlignment(ScalarOpAlignment); + replaceValue(I, *NSI); + // Need erasing the store manually. + I.eraseFromParent(); + return true; + } + + return false; +} + +/// Try to scalarize vector loads feeding extractelement instructions. +bool VectorCombine::scalarizeLoadExtract(Instruction &I) { + Value *Ptr; + Value *Idx; + if (!match(&I, m_ExtractElt(m_Load(m_Value(Ptr)), m_Value(Idx)))) + return false; + + auto *LI = cast<LoadInst>(I.getOperand(0)); + const DataLayout &DL = I.getModule()->getDataLayout(); + if (LI->isVolatile() || !DL.typeSizeEqualsStoreSize(LI->getType())) + return false; + + auto *FixedVT = dyn_cast<FixedVectorType>(LI->getType()); + if (!FixedVT) + return false; + + InstructionCost OriginalCost = TTI.getMemoryOpCost( + Instruction::Load, LI->getType(), Align(LI->getAlignment()), + LI->getPointerAddressSpace()); + InstructionCost ScalarizedCost = 0; + + Instruction *LastCheckedInst = LI; + unsigned NumInstChecked = 0; + // Check if all users of the load are extracts with no memory modifications + // between the load and the extract. Compute the cost of both the original + // code and the scalarized version. + for (User *U : LI->users()) { + auto *UI = dyn_cast<ExtractElementInst>(U); + if (!UI || UI->getParent() != LI->getParent()) + return false; + + if (!isGuaranteedNotToBePoison(UI->getOperand(1), &AC, LI, &DT)) + return false; + + // Check if any instruction between the load and the extract may modify + // memory. + if (LastCheckedInst->comesBefore(UI)) { + for (Instruction &I : + make_range(std::next(LI->getIterator()), UI->getIterator())) { + // Bail out if we reached the check limit or the instruction may write + // to memory. + if (NumInstChecked == MaxInstrsToScan || I.mayWriteToMemory()) + return false; + NumInstChecked++; + } + } + + if (!LastCheckedInst) + LastCheckedInst = UI; + else if (LastCheckedInst->comesBefore(UI)) + LastCheckedInst = UI; + + if (!canScalarizeAccess(FixedVT, UI->getOperand(1), &I, AC)) + return false; + + auto *Index = dyn_cast<ConstantInt>(UI->getOperand(1)); + OriginalCost += + TTI.getVectorInstrCost(Instruction::ExtractElement, LI->getType(), + Index ? Index->getZExtValue() : -1); + ScalarizedCost += + TTI.getMemoryOpCost(Instruction::Load, FixedVT->getElementType(), + Align(1), LI->getPointerAddressSpace()); + ScalarizedCost += TTI.getAddressComputationCost(FixedVT->getElementType()); + } + + if (ScalarizedCost >= OriginalCost) + return false; + + // Replace extracts with narrow scalar loads. + for (User *U : LI->users()) { + auto *EI = cast<ExtractElementInst>(U); + Builder.SetInsertPoint(EI); + + Value *Idx = EI->getOperand(1); + Value *GEP = + Builder.CreateInBoundsGEP(FixedVT, Ptr, {Builder.getInt32(0), Idx}); + auto *NewLoad = cast<LoadInst>(Builder.CreateLoad( + FixedVT->getElementType(), GEP, EI->getName() + ".scalar")); + + Align ScalarOpAlignment = computeAlignmentAfterScalarization( + LI->getAlign(), FixedVT->getElementType(), Idx, DL); + NewLoad->setAlignment(ScalarOpAlignment); + + replaceValue(*EI, *NewLoad); + } + + return true; +} + /// This is the entry point for all transforms. Pass manager differences are /// handled in the callers of this function. bool VectorCombine::run() { @@ -767,11 +961,8 @@ bool VectorCombine::run() { // Ignore unreachable basic blocks. if (!DT.isReachableFromEntry(&BB)) continue; - // Do not delete instructions under here and invalidate the iterator. - // Walk the block forwards to enable simple iterative chains of transforms. - // TODO: It could be more efficient to remove dead instructions - // iteratively in this loop rather than waiting until the end. - for (Instruction &I : BB) { + // Use early increment range so that we can erase instructions in loop. + for (Instruction &I : make_early_inc_range(BB)) { if (isa<DbgInfoIntrinsic>(I)) continue; Builder.SetInsertPoint(&I); @@ -780,6 +971,8 @@ bool VectorCombine::run() { MadeChange |= foldBitcastShuf(I); MadeChange |= scalarizeBinopOrCmp(I); MadeChange |= foldExtractedCmps(I); + MadeChange |= scalarizeLoadExtract(I); + MadeChange |= foldSingleElementStore(I); } } @@ -802,8 +995,10 @@ public: } void getAnalysisUsage(AnalysisUsage &AU) const override { + AU.addRequired<AssumptionCacheTracker>(); AU.addRequired<DominatorTreeWrapperPass>(); AU.addRequired<TargetTransformInfoWrapperPass>(); + AU.addRequired<AAResultsWrapperPass>(); AU.setPreservesCFG(); AU.addPreserved<DominatorTreeWrapperPass>(); AU.addPreserved<GlobalsAAWrapperPass>(); @@ -815,9 +1010,11 @@ public: bool runOnFunction(Function &F) override { if (skipFunction(F)) return false; + auto &AC = getAnalysis<AssumptionCacheTracker>().getAssumptionCache(F); auto &TTI = getAnalysis<TargetTransformInfoWrapperPass>().getTTI(F); auto &DT = getAnalysis<DominatorTreeWrapperPass>().getDomTree(); - VectorCombine Combiner(F, TTI, DT); + auto &AA = getAnalysis<AAResultsWrapperPass>().getAAResults(); + VectorCombine Combiner(F, TTI, DT, AA, AC); return Combiner.run(); } }; @@ -827,6 +1024,7 @@ char VectorCombineLegacyPass::ID = 0; INITIALIZE_PASS_BEGIN(VectorCombineLegacyPass, "vector-combine", "Optimize scalar/vector ops", false, false) +INITIALIZE_PASS_DEPENDENCY(AssumptionCacheTracker) INITIALIZE_PASS_DEPENDENCY(DominatorTreeWrapperPass) INITIALIZE_PASS_END(VectorCombineLegacyPass, "vector-combine", "Optimize scalar/vector ops", false, false) @@ -836,15 +1034,14 @@ Pass *llvm::createVectorCombinePass() { PreservedAnalyses VectorCombinePass::run(Function &F, FunctionAnalysisManager &FAM) { + auto &AC = FAM.getResult<AssumptionAnalysis>(F); TargetTransformInfo &TTI = FAM.getResult<TargetIRAnalysis>(F); DominatorTree &DT = FAM.getResult<DominatorTreeAnalysis>(F); - VectorCombine Combiner(F, TTI, DT); + AAResults &AA = FAM.getResult<AAManager>(F); + VectorCombine Combiner(F, TTI, DT, AA, AC); if (!Combiner.run()) return PreservedAnalyses::all(); PreservedAnalyses PA; PA.preserveSet<CFGAnalyses>(); - PA.preserve<GlobalsAA>(); - PA.preserve<AAManager>(); - PA.preserve<BasicAA>(); return PA; } |
