diff options
Diffstat (limited to 'lib/Transforms/InstCombine/InstCombineCalls.cpp')
| -rw-r--r-- | lib/Transforms/InstCombine/InstCombineCalls.cpp | 1150 |
1 files changed, 581 insertions, 569 deletions
diff --git a/lib/Transforms/InstCombine/InstCombineCalls.cpp b/lib/Transforms/InstCombine/InstCombineCalls.cpp index aeb25d530d71..4b3333affa72 100644 --- a/lib/Transforms/InstCombine/InstCombineCalls.cpp +++ b/lib/Transforms/InstCombine/InstCombineCalls.cpp @@ -1,19 +1,19 @@ //===- InstCombineCalls.cpp -----------------------------------------------===// // -// The LLVM Compiler Infrastructure -// -// This file is distributed under the University of Illinois Open Source -// License. See LICENSE.TXT for details. +// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions. +// See https://llvm.org/LICENSE.txt for license information. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception // //===----------------------------------------------------------------------===// // -// This file implements the visitCall and visitInvoke functions. +// This file implements the visitCall, visitInvoke, and visitCallBr functions. // //===----------------------------------------------------------------------===// #include "InstCombineInternal.h" #include "llvm/ADT/APFloat.h" #include "llvm/ADT/APInt.h" +#include "llvm/ADT/APSInt.h" #include "llvm/ADT/ArrayRef.h" #include "llvm/ADT/None.h" #include "llvm/ADT/Optional.h" @@ -23,12 +23,12 @@ #include "llvm/ADT/Twine.h" #include "llvm/Analysis/AssumptionCache.h" #include "llvm/Analysis/InstructionSimplify.h" +#include "llvm/Analysis/Loads.h" #include "llvm/Analysis/MemoryBuiltins.h" -#include "llvm/Transforms/Utils/Local.h" #include "llvm/Analysis/ValueTracking.h" +#include "llvm/Analysis/VectorUtils.h" #include "llvm/IR/Attributes.h" #include "llvm/IR/BasicBlock.h" -#include "llvm/IR/CallSite.h" #include "llvm/IR/Constant.h" #include "llvm/IR/Constants.h" #include "llvm/IR/DataLayout.h" @@ -58,6 +58,7 @@ #include "llvm/Support/MathExtras.h" #include "llvm/Support/raw_ostream.h" #include "llvm/Transforms/InstCombine/InstCombineWorklist.h" +#include "llvm/Transforms/Utils/Local.h" #include "llvm/Transforms/Utils/SimplifyLibCalls.h" #include <algorithm> #include <cassert> @@ -121,6 +122,15 @@ Instruction *InstCombiner::SimplifyAnyMemTransfer(AnyMemTransferInst *MI) { return MI; } + // If we have a store to a location which is known constant, we can conclude + // that the store must be storing the constant value (else the memory + // wouldn't be constant), and this must be a noop. + if (AA->pointsToConstantMemory(MI->getDest())) { + // Set the size of the copy to 0, it will be deleted on the next iteration. + MI->setLength(Constant::getNullValue(MI->getLength()->getType())); + return MI; + } + // If MemCpyInst length is 1/2/4/8 bytes then replace memcpy with // load/store. ConstantInt *MemOpLength = dyn_cast<ConstantInt>(MI->getLength()); @@ -173,7 +183,7 @@ Instruction *InstCombiner::SimplifyAnyMemTransfer(AnyMemTransferInst *MI) { Value *Src = Builder.CreateBitCast(MI->getArgOperand(1), NewSrcPtrTy); Value *Dest = Builder.CreateBitCast(MI->getArgOperand(0), NewDstPtrTy); - LoadInst *L = Builder.CreateLoad(Src); + LoadInst *L = Builder.CreateLoad(IntType, Src); // Alignment from the mem intrinsic will be better, so use it. L->setAlignment(CopySrcAlign); if (CopyMD) @@ -219,6 +229,15 @@ Instruction *InstCombiner::SimplifyAnyMemSet(AnyMemSetInst *MI) { return MI; } + // If we have a store to a location which is known constant, we can conclude + // that the store must be storing the constant value (else the memory + // wouldn't be constant), and this must be a noop. + if (AA->pointsToConstantMemory(MI->getDest())) { + // Set the size of the copy to 0, it will be deleted on the next iteration. + MI->setLength(Constant::getNullValue(MI->getLength()->getType())); + return MI; + } + // Extract the length and alignment and fill if they are constant. ConstantInt *LenC = dyn_cast<ConstantInt>(MI->getLength()); ConstantInt *FillC = dyn_cast<ConstantInt>(MI->getValue()); @@ -523,7 +542,8 @@ static Value *simplifyX86varShift(const IntrinsicInst &II, return Builder.CreateAShr(Vec, ShiftVec); } -static Value *simplifyX86pack(IntrinsicInst &II, bool IsSigned) { +static Value *simplifyX86pack(IntrinsicInst &II, + InstCombiner::BuilderTy &Builder, bool IsSigned) { Value *Arg0 = II.getArgOperand(0); Value *Arg1 = II.getArgOperand(1); Type *ResTy = II.getType(); @@ -534,167 +554,58 @@ static Value *simplifyX86pack(IntrinsicInst &II, bool IsSigned) { Type *ArgTy = Arg0->getType(); unsigned NumLanes = ResTy->getPrimitiveSizeInBits() / 128; - unsigned NumDstElts = ResTy->getVectorNumElements(); unsigned NumSrcElts = ArgTy->getVectorNumElements(); - assert(NumDstElts == (2 * NumSrcElts) && "Unexpected packing types"); + assert(ResTy->getVectorNumElements() == (2 * NumSrcElts) && + "Unexpected packing types"); - unsigned NumDstEltsPerLane = NumDstElts / NumLanes; unsigned NumSrcEltsPerLane = NumSrcElts / NumLanes; unsigned DstScalarSizeInBits = ResTy->getScalarSizeInBits(); - assert(ArgTy->getScalarSizeInBits() == (2 * DstScalarSizeInBits) && + unsigned SrcScalarSizeInBits = ArgTy->getScalarSizeInBits(); + assert(SrcScalarSizeInBits == (2 * DstScalarSizeInBits) && "Unexpected packing types"); // Constant folding. - auto *Cst0 = dyn_cast<Constant>(Arg0); - auto *Cst1 = dyn_cast<Constant>(Arg1); - if (!Cst0 || !Cst1) + if (!isa<Constant>(Arg0) || !isa<Constant>(Arg1)) return nullptr; - SmallVector<Constant *, 32> Vals; - for (unsigned Lane = 0; Lane != NumLanes; ++Lane) { - for (unsigned Elt = 0; Elt != NumDstEltsPerLane; ++Elt) { - unsigned SrcIdx = Lane * NumSrcEltsPerLane + Elt % NumSrcEltsPerLane; - auto *Cst = (Elt >= NumSrcEltsPerLane) ? Cst1 : Cst0; - auto *COp = Cst->getAggregateElement(SrcIdx); - if (COp && isa<UndefValue>(COp)) { - Vals.push_back(UndefValue::get(ResTy->getScalarType())); - continue; - } - - auto *CInt = dyn_cast_or_null<ConstantInt>(COp); - if (!CInt) - return nullptr; - - APInt Val = CInt->getValue(); - assert(Val.getBitWidth() == ArgTy->getScalarSizeInBits() && - "Unexpected constant bitwidth"); - - if (IsSigned) { - // PACKSS: Truncate signed value with signed saturation. - // Source values less than dst minint are saturated to minint. - // Source values greater than dst maxint are saturated to maxint. - if (Val.isSignedIntN(DstScalarSizeInBits)) - Val = Val.trunc(DstScalarSizeInBits); - else if (Val.isNegative()) - Val = APInt::getSignedMinValue(DstScalarSizeInBits); - else - Val = APInt::getSignedMaxValue(DstScalarSizeInBits); - } else { - // PACKUS: Truncate signed value with unsigned saturation. - // Source values less than zero are saturated to zero. - // Source values greater than dst maxuint are saturated to maxuint. - if (Val.isIntN(DstScalarSizeInBits)) - Val = Val.trunc(DstScalarSizeInBits); - else if (Val.isNegative()) - Val = APInt::getNullValue(DstScalarSizeInBits); - else - Val = APInt::getAllOnesValue(DstScalarSizeInBits); - } - - Vals.push_back(ConstantInt::get(ResTy->getScalarType(), Val)); - } - } - - return ConstantVector::get(Vals); -} - -// Replace X86-specific intrinsics with generic floor-ceil where applicable. -static Value *simplifyX86round(IntrinsicInst &II, - InstCombiner::BuilderTy &Builder) { - ConstantInt *Arg = nullptr; - Intrinsic::ID IntrinsicID = II.getIntrinsicID(); - - if (IntrinsicID == Intrinsic::x86_sse41_round_ss || - IntrinsicID == Intrinsic::x86_sse41_round_sd) - Arg = dyn_cast<ConstantInt>(II.getArgOperand(2)); - else if (IntrinsicID == Intrinsic::x86_avx512_mask_rndscale_ss || - IntrinsicID == Intrinsic::x86_avx512_mask_rndscale_sd) - Arg = dyn_cast<ConstantInt>(II.getArgOperand(4)); - else - Arg = dyn_cast<ConstantInt>(II.getArgOperand(1)); - if (!Arg) - return nullptr; - unsigned RoundControl = Arg->getZExtValue(); - - Arg = nullptr; - unsigned SAE = 0; - if (IntrinsicID == Intrinsic::x86_avx512_mask_rndscale_ps_512 || - IntrinsicID == Intrinsic::x86_avx512_mask_rndscale_pd_512) - Arg = dyn_cast<ConstantInt>(II.getArgOperand(4)); - else if (IntrinsicID == Intrinsic::x86_avx512_mask_rndscale_ss || - IntrinsicID == Intrinsic::x86_avx512_mask_rndscale_sd) - Arg = dyn_cast<ConstantInt>(II.getArgOperand(5)); - else - SAE = 4; - if (!SAE) { - if (!Arg) - return nullptr; - SAE = Arg->getZExtValue(); + // Clamp Values - signed/unsigned both use signed clamp values, but they + // differ on the min/max values. + APInt MinValue, MaxValue; + if (IsSigned) { + // PACKSS: Truncate signed value with signed saturation. + // Source values less than dst minint are saturated to minint. + // Source values greater than dst maxint are saturated to maxint. + MinValue = + APInt::getSignedMinValue(DstScalarSizeInBits).sext(SrcScalarSizeInBits); + MaxValue = + APInt::getSignedMaxValue(DstScalarSizeInBits).sext(SrcScalarSizeInBits); + } else { + // PACKUS: Truncate signed value with unsigned saturation. + // Source values less than zero are saturated to zero. + // Source values greater than dst maxuint are saturated to maxuint. + MinValue = APInt::getNullValue(SrcScalarSizeInBits); + MaxValue = APInt::getLowBitsSet(SrcScalarSizeInBits, DstScalarSizeInBits); } - if (SAE != 4 || (RoundControl != 2 /*ceil*/ && RoundControl != 1 /*floor*/)) - return nullptr; + auto *MinC = Constant::getIntegerValue(ArgTy, MinValue); + auto *MaxC = Constant::getIntegerValue(ArgTy, MaxValue); + Arg0 = Builder.CreateSelect(Builder.CreateICmpSLT(Arg0, MinC), MinC, Arg0); + Arg1 = Builder.CreateSelect(Builder.CreateICmpSLT(Arg1, MinC), MinC, Arg1); + Arg0 = Builder.CreateSelect(Builder.CreateICmpSGT(Arg0, MaxC), MaxC, Arg0); + Arg1 = Builder.CreateSelect(Builder.CreateICmpSGT(Arg1, MaxC), MaxC, Arg1); - Value *Src, *Dst, *Mask; - bool IsScalar = false; - if (IntrinsicID == Intrinsic::x86_sse41_round_ss || - IntrinsicID == Intrinsic::x86_sse41_round_sd || - IntrinsicID == Intrinsic::x86_avx512_mask_rndscale_ss || - IntrinsicID == Intrinsic::x86_avx512_mask_rndscale_sd) { - IsScalar = true; - if (IntrinsicID == Intrinsic::x86_avx512_mask_rndscale_ss || - IntrinsicID == Intrinsic::x86_avx512_mask_rndscale_sd) { - Mask = II.getArgOperand(3); - Value *Zero = Constant::getNullValue(Mask->getType()); - Mask = Builder.CreateAnd(Mask, 1); - Mask = Builder.CreateICmp(ICmpInst::ICMP_NE, Mask, Zero); - Dst = II.getArgOperand(2); - } else - Dst = II.getArgOperand(0); - Src = Builder.CreateExtractElement(II.getArgOperand(1), (uint64_t)0); - } else { - Src = II.getArgOperand(0); - if (IntrinsicID == Intrinsic::x86_avx512_mask_rndscale_ps_128 || - IntrinsicID == Intrinsic::x86_avx512_mask_rndscale_ps_256 || - IntrinsicID == Intrinsic::x86_avx512_mask_rndscale_ps_512 || - IntrinsicID == Intrinsic::x86_avx512_mask_rndscale_pd_128 || - IntrinsicID == Intrinsic::x86_avx512_mask_rndscale_pd_256 || - IntrinsicID == Intrinsic::x86_avx512_mask_rndscale_pd_512) { - Dst = II.getArgOperand(2); - Mask = II.getArgOperand(3); - } else { - Dst = Src; - Mask = ConstantInt::getAllOnesValue( - Builder.getIntNTy(Src->getType()->getVectorNumElements())); - } + // Shuffle clamped args together at the lane level. + SmallVector<unsigned, 32> PackMask; + for (unsigned Lane = 0; Lane != NumLanes; ++Lane) { + for (unsigned Elt = 0; Elt != NumSrcEltsPerLane; ++Elt) + PackMask.push_back(Elt + (Lane * NumSrcEltsPerLane)); + for (unsigned Elt = 0; Elt != NumSrcEltsPerLane; ++Elt) + PackMask.push_back(Elt + (Lane * NumSrcEltsPerLane) + NumSrcElts); } + auto *Shuffle = Builder.CreateShuffleVector(Arg0, Arg1, PackMask); - Intrinsic::ID ID = (RoundControl == 2) ? Intrinsic::ceil : Intrinsic::floor; - Value *Res = Builder.CreateUnaryIntrinsic(ID, Src, &II); - if (!IsScalar) { - if (auto *C = dyn_cast<Constant>(Mask)) - if (C->isAllOnesValue()) - return Res; - auto *MaskTy = VectorType::get( - Builder.getInt1Ty(), cast<IntegerType>(Mask->getType())->getBitWidth()); - Mask = Builder.CreateBitCast(Mask, MaskTy); - unsigned Width = Src->getType()->getVectorNumElements(); - if (MaskTy->getVectorNumElements() > Width) { - uint32_t Indices[4]; - for (unsigned i = 0; i != Width; ++i) - Indices[i] = i; - Mask = Builder.CreateShuffleVector(Mask, Mask, - makeArrayRef(Indices, Width)); - } - return Builder.CreateSelect(Mask, Res, Dst); - } - if (IntrinsicID == Intrinsic::x86_avx512_mask_rndscale_ss || - IntrinsicID == Intrinsic::x86_avx512_mask_rndscale_sd) { - Dst = Builder.CreateExtractElement(Dst, (uint64_t)0); - Res = Builder.CreateSelect(Mask, Res, Dst); - Dst = II.getArgOperand(0); - } - return Builder.CreateInsertElement(Dst, Res, (uint64_t)0); + // Truncate to dst size. + return Builder.CreateTrunc(Shuffle, ResTy); } static Value *simplifyX86movmsk(const IntrinsicInst &II, @@ -711,43 +622,44 @@ static Value *simplifyX86movmsk(const IntrinsicInst &II, if (!ArgTy->isVectorTy()) return nullptr; - if (auto *C = dyn_cast<Constant>(Arg)) { - // Extract signbits of the vector input and pack into integer result. - APInt Result(ResTy->getPrimitiveSizeInBits(), 0); - for (unsigned I = 0, E = ArgTy->getVectorNumElements(); I != E; ++I) { - auto *COp = C->getAggregateElement(I); - if (!COp) - return nullptr; - if (isa<UndefValue>(COp)) - continue; + // Expand MOVMSK to compare/bitcast/zext: + // e.g. PMOVMSKB(v16i8 x): + // %cmp = icmp slt <16 x i8> %x, zeroinitializer + // %int = bitcast <16 x i1> %cmp to i16 + // %res = zext i16 %int to i32 + unsigned NumElts = ArgTy->getVectorNumElements(); + Type *IntegerVecTy = VectorType::getInteger(cast<VectorType>(ArgTy)); + Type *IntegerTy = Builder.getIntNTy(NumElts); - auto *CInt = dyn_cast<ConstantInt>(COp); - auto *CFp = dyn_cast<ConstantFP>(COp); - if (!CInt && !CFp) - return nullptr; + Value *Res = Builder.CreateBitCast(Arg, IntegerVecTy); + Res = Builder.CreateICmpSLT(Res, Constant::getNullValue(IntegerVecTy)); + Res = Builder.CreateBitCast(Res, IntegerTy); + Res = Builder.CreateZExtOrTrunc(Res, ResTy); + return Res; +} - if ((CInt && CInt->isNegative()) || (CFp && CFp->isNegative())) - Result.setBit(I); - } - return Constant::getIntegerValue(ResTy, Result); - } +static Value *simplifyX86addcarry(const IntrinsicInst &II, + InstCombiner::BuilderTy &Builder) { + Value *CarryIn = II.getArgOperand(0); + Value *Op1 = II.getArgOperand(1); + Value *Op2 = II.getArgOperand(2); + Type *RetTy = II.getType(); + Type *OpTy = Op1->getType(); + assert(RetTy->getStructElementType(0)->isIntegerTy(8) && + RetTy->getStructElementType(1) == OpTy && OpTy == Op2->getType() && + "Unexpected types for x86 addcarry"); - // Look for a sign-extended boolean source vector as the argument to this - // movmsk. If the argument is bitcast, look through that, but make sure the - // source of that bitcast is still a vector with the same number of elements. - // TODO: We can also convert a bitcast with wider elements, but that requires - // duplicating the bool source sign bits to match the number of elements - // expected by the movmsk call. - Arg = peekThroughBitcast(Arg); - Value *X; - if (Arg->getType()->isVectorTy() && - Arg->getType()->getVectorNumElements() == ArgTy->getVectorNumElements() && - match(Arg, m_SExt(m_Value(X))) && X->getType()->isIntOrIntVectorTy(1)) { - // call iM movmsk(sext <N x i1> X) --> zext (bitcast <N x i1> X to iN) to iM - unsigned NumElts = X->getType()->getVectorNumElements(); - Type *ScalarTy = Type::getIntNTy(Arg->getContext(), NumElts); - Value *BC = Builder.CreateBitCast(X, ScalarTy); - return Builder.CreateZExtOrTrunc(BC, ResTy); + // If carry-in is zero, this is just an unsigned add with overflow. + if (match(CarryIn, m_ZeroInt())) { + Value *UAdd = Builder.CreateIntrinsic(Intrinsic::uadd_with_overflow, OpTy, + { Op1, Op2 }); + // The types have to be adjusted to match the x86 call types. + Value *UAddResult = Builder.CreateExtractValue(UAdd, 0); + Value *UAddOV = Builder.CreateZExt(Builder.CreateExtractValue(UAdd, 1), + Builder.getInt8Ty()); + Value *Res = UndefValue::get(RetTy); + Res = Builder.CreateInsertValue(Res, UAddOV, 0); + return Builder.CreateInsertValue(Res, UAddResult, 1); } return nullptr; @@ -892,7 +804,7 @@ static Value *simplifyX86extrq(IntrinsicInst &II, Value *Op0, if (II.getIntrinsicID() == Intrinsic::x86_sse4a_extrq) { Value *Args[] = {Op0, CILength, CIIndex}; Module *M = II.getModule(); - Value *F = Intrinsic::getDeclaration(M, Intrinsic::x86_sse4a_extrqi); + Function *F = Intrinsic::getDeclaration(M, Intrinsic::x86_sse4a_extrqi); return Builder.CreateCall(F, Args); } } @@ -993,7 +905,7 @@ static Value *simplifyX86insertq(IntrinsicInst &II, Value *Op0, Value *Op1, Value *Args[] = {Op0, Op1, CILength, CIIndex}; Module *M = II.getModule(); - Value *F = Intrinsic::getDeclaration(M, Intrinsic::x86_sse4a_insertqi); + Function *F = Intrinsic::getDeclaration(M, Intrinsic::x86_sse4a_insertqi); return Builder.CreateCall(F, Args); } @@ -1134,82 +1046,42 @@ static Value *simplifyX86vpermv(const IntrinsicInst &II, return Builder.CreateShuffleVector(V1, V2, ShuffleMask); } -/// Decode XOP integer vector comparison intrinsics. -static Value *simplifyX86vpcom(const IntrinsicInst &II, - InstCombiner::BuilderTy &Builder, - bool IsSigned) { - if (auto *CInt = dyn_cast<ConstantInt>(II.getArgOperand(2))) { - uint64_t Imm = CInt->getZExtValue() & 0x7; - VectorType *VecTy = cast<VectorType>(II.getType()); - CmpInst::Predicate Pred = ICmpInst::BAD_ICMP_PREDICATE; - - switch (Imm) { - case 0x0: - Pred = IsSigned ? ICmpInst::ICMP_SLT : ICmpInst::ICMP_ULT; - break; - case 0x1: - Pred = IsSigned ? ICmpInst::ICMP_SLE : ICmpInst::ICMP_ULE; - break; - case 0x2: - Pred = IsSigned ? ICmpInst::ICMP_SGT : ICmpInst::ICMP_UGT; - break; - case 0x3: - Pred = IsSigned ? ICmpInst::ICMP_SGE : ICmpInst::ICMP_UGE; - break; - case 0x4: - Pred = ICmpInst::ICMP_EQ; break; - case 0x5: - Pred = ICmpInst::ICMP_NE; break; - case 0x6: - return ConstantInt::getSigned(VecTy, 0); // FALSE - case 0x7: - return ConstantInt::getSigned(VecTy, -1); // TRUE - } - - if (Value *Cmp = Builder.CreateICmp(Pred, II.getArgOperand(0), - II.getArgOperand(1))) - return Builder.CreateSExtOrTrunc(Cmp, VecTy); - } - return nullptr; -} - -static bool maskIsAllOneOrUndef(Value *Mask) { - auto *ConstMask = dyn_cast<Constant>(Mask); - if (!ConstMask) - return false; - if (ConstMask->isAllOnesValue() || isa<UndefValue>(ConstMask)) - return true; - for (unsigned I = 0, E = ConstMask->getType()->getVectorNumElements(); I != E; - ++I) { - if (auto *MaskElt = ConstMask->getAggregateElement(I)) - if (MaskElt->isAllOnesValue() || isa<UndefValue>(MaskElt)) - continue; - return false; - } - return true; -} +// TODO, Obvious Missing Transforms: +// * Narrow width by halfs excluding zero/undef lanes +Value *InstCombiner::simplifyMaskedLoad(IntrinsicInst &II) { + Value *LoadPtr = II.getArgOperand(0); + unsigned Alignment = cast<ConstantInt>(II.getArgOperand(1))->getZExtValue(); -static Value *simplifyMaskedLoad(const IntrinsicInst &II, - InstCombiner::BuilderTy &Builder) { // If the mask is all ones or undefs, this is a plain vector load of the 1st // argument. - if (maskIsAllOneOrUndef(II.getArgOperand(2))) { - Value *LoadPtr = II.getArgOperand(0); - unsigned Alignment = cast<ConstantInt>(II.getArgOperand(1))->getZExtValue(); - return Builder.CreateAlignedLoad(LoadPtr, Alignment, "unmaskedload"); + if (maskIsAllOneOrUndef(II.getArgOperand(2))) + return Builder.CreateAlignedLoad(II.getType(), LoadPtr, Alignment, + "unmaskedload"); + + // If we can unconditionally load from this address, replace with a + // load/select idiom. TODO: use DT for context sensitive query + if (isDereferenceableAndAlignedPointer(LoadPtr, II.getType(), Alignment, + II.getModule()->getDataLayout(), + &II, nullptr)) { + Value *LI = Builder.CreateAlignedLoad(II.getType(), LoadPtr, Alignment, + "unmaskedload"); + return Builder.CreateSelect(II.getArgOperand(2), LI, II.getArgOperand(3)); } return nullptr; } -static Instruction *simplifyMaskedStore(IntrinsicInst &II, InstCombiner &IC) { +// TODO, Obvious Missing Transforms: +// * Single constant active lane -> store +// * Narrow width by halfs excluding zero/undef lanes +Instruction *InstCombiner::simplifyMaskedStore(IntrinsicInst &II) { auto *ConstMask = dyn_cast<Constant>(II.getArgOperand(3)); if (!ConstMask) return nullptr; // If the mask is all zeros, this instruction does nothing. if (ConstMask->isNullValue()) - return IC.eraseInstFromFunction(II); + return eraseInstFromFunction(II); // If the mask is all ones, this is a plain vector store of the 1st argument. if (ConstMask->isAllOnesValue()) { @@ -1218,14 +1090,57 @@ static Instruction *simplifyMaskedStore(IntrinsicInst &II, InstCombiner &IC) { return new StoreInst(II.getArgOperand(0), StorePtr, false, Alignment); } + // Use masked off lanes to simplify operands via SimplifyDemandedVectorElts + APInt DemandedElts = possiblyDemandedEltsInMask(ConstMask); + APInt UndefElts(DemandedElts.getBitWidth(), 0); + if (Value *V = SimplifyDemandedVectorElts(II.getOperand(0), + DemandedElts, UndefElts)) { + II.setOperand(0, V); + return &II; + } + return nullptr; } -static Instruction *simplifyMaskedGather(IntrinsicInst &II, InstCombiner &IC) { - // If the mask is all zeros, return the "passthru" argument of the gather. - auto *ConstMask = dyn_cast<Constant>(II.getArgOperand(2)); - if (ConstMask && ConstMask->isNullValue()) - return IC.replaceInstUsesWith(II, II.getArgOperand(3)); +// TODO, Obvious Missing Transforms: +// * Single constant active lane load -> load +// * Dereferenceable address & few lanes -> scalarize speculative load/selects +// * Adjacent vector addresses -> masked.load +// * Narrow width by halfs excluding zero/undef lanes +// * Vector splat address w/known mask -> scalar load +// * Vector incrementing address -> vector masked load +Instruction *InstCombiner::simplifyMaskedGather(IntrinsicInst &II) { + return nullptr; +} + +// TODO, Obvious Missing Transforms: +// * Single constant active lane -> store +// * Adjacent vector addresses -> masked.store +// * Narrow store width by halfs excluding zero/undef lanes +// * Vector splat address w/known mask -> scalar store +// * Vector incrementing address -> vector masked store +Instruction *InstCombiner::simplifyMaskedScatter(IntrinsicInst &II) { + auto *ConstMask = dyn_cast<Constant>(II.getArgOperand(3)); + if (!ConstMask) + return nullptr; + + // If the mask is all zeros, a scatter does nothing. + if (ConstMask->isNullValue()) + return eraseInstFromFunction(II); + + // Use masked off lanes to simplify operands via SimplifyDemandedVectorElts + APInt DemandedElts = possiblyDemandedEltsInMask(ConstMask); + APInt UndefElts(DemandedElts.getBitWidth(), 0); + if (Value *V = SimplifyDemandedVectorElts(II.getOperand(0), + DemandedElts, UndefElts)) { + II.setOperand(0, V); + return &II; + } + if (Value *V = SimplifyDemandedVectorElts(II.getOperand(1), + DemandedElts, UndefElts)) { + II.setOperand(1, V); + return &II; + } return nullptr; } @@ -1264,25 +1179,41 @@ static Instruction *simplifyInvariantGroupIntrinsic(IntrinsicInst &II, return cast<Instruction>(Result); } -static Instruction *simplifyMaskedScatter(IntrinsicInst &II, InstCombiner &IC) { - // If the mask is all zeros, a scatter does nothing. - auto *ConstMask = dyn_cast<Constant>(II.getArgOperand(3)); - if (ConstMask && ConstMask->isNullValue()) - return IC.eraseInstFromFunction(II); - - return nullptr; -} - static Instruction *foldCttzCtlz(IntrinsicInst &II, InstCombiner &IC) { assert((II.getIntrinsicID() == Intrinsic::cttz || II.getIntrinsicID() == Intrinsic::ctlz) && "Expected cttz or ctlz intrinsic"); + bool IsTZ = II.getIntrinsicID() == Intrinsic::cttz; Value *Op0 = II.getArgOperand(0); + Value *X; + // ctlz(bitreverse(x)) -> cttz(x) + // cttz(bitreverse(x)) -> ctlz(x) + if (match(Op0, m_BitReverse(m_Value(X)))) { + Intrinsic::ID ID = IsTZ ? Intrinsic::ctlz : Intrinsic::cttz; + Function *F = Intrinsic::getDeclaration(II.getModule(), ID, II.getType()); + return CallInst::Create(F, {X, II.getArgOperand(1)}); + } + + if (IsTZ) { + // cttz(-x) -> cttz(x) + if (match(Op0, m_Neg(m_Value(X)))) { + II.setOperand(0, X); + return &II; + } + + // cttz(abs(x)) -> cttz(x) + // cttz(nabs(x)) -> cttz(x) + Value *Y; + SelectPatternFlavor SPF = matchSelectPattern(Op0, X, Y).Flavor; + if (SPF == SPF_ABS || SPF == SPF_NABS) { + II.setOperand(0, X); + return &II; + } + } KnownBits Known = IC.computeKnownBits(Op0, 0, &II); // Create a mask for bits above (ctlz) or below (cttz) the first known one. - bool IsTZ = II.getIntrinsicID() == Intrinsic::cttz; unsigned PossibleZeros = IsTZ ? Known.countMaxTrailingZeros() : Known.countMaxLeadingZeros(); unsigned DefiniteZeros = IsTZ ? Known.countMinTrailingZeros() @@ -1328,6 +1259,14 @@ static Instruction *foldCtpop(IntrinsicInst &II, InstCombiner &IC) { assert(II.getIntrinsicID() == Intrinsic::ctpop && "Expected ctpop intrinsic"); Value *Op0 = II.getArgOperand(0); + Value *X; + // ctpop(bitreverse(x)) -> ctpop(x) + // ctpop(bswap(x)) -> ctpop(x) + if (match(Op0, m_BitReverse(m_Value(X))) || match(Op0, m_BSwap(m_Value(X)))) { + II.setOperand(0, X); + return &II; + } + // FIXME: Try to simplify vectors of integers. auto *IT = dyn_cast<IntegerType>(Op0->getType()); if (!IT) @@ -1513,7 +1452,7 @@ static Value *simplifyNeonVld1(const IntrinsicInst &II, auto *BCastInst = Builder.CreateBitCast(II.getArgOperand(0), PointerType::get(II.getType(), 0)); - return Builder.CreateAlignedLoad(BCastInst, Alignment); + return Builder.CreateAlignedLoad(II.getType(), BCastInst, Alignment); } // Returns true iff the 2 intrinsics have the same operands, limiting the @@ -1827,8 +1766,18 @@ static Instruction *canonicalizeConstantArg0ToArg1(CallInst &Call) { return nullptr; } +Instruction *InstCombiner::foldIntrinsicWithOverflowCommon(IntrinsicInst *II) { + WithOverflowInst *WO = cast<WithOverflowInst>(II); + Value *OperationResult = nullptr; + Constant *OverflowResult = nullptr; + if (OptimizeOverflowCheck(WO->getBinaryOp(), WO->isSigned(), WO->getLHS(), + WO->getRHS(), *WO, OperationResult, OverflowResult)) + return CreateOverflowTuple(WO, OperationResult, OverflowResult); + return nullptr; +} + /// CallInst simplification. This mostly only handles folding of intrinsic -/// instructions. For normal calls, it allows visitCallSite to do the heavy +/// instructions. For normal calls, it allows visitCallBase to do the heavy /// lifting. Instruction *InstCombiner::visitCallInst(CallInst &CI) { if (Value *V = SimplifyCall(&CI, SQ.getWithInstruction(&CI))) @@ -1845,10 +1794,10 @@ Instruction *InstCombiner::visitCallInst(CallInst &CI) { } IntrinsicInst *II = dyn_cast<IntrinsicInst>(&CI); - if (!II) return visitCallSite(&CI); + if (!II) return visitCallBase(CI); - // Intrinsics cannot occur in an invoke, so handle them here instead of in - // visitCallSite. + // Intrinsics cannot occur in an invoke or a callbr, so handle them here + // instead of in visitCallBase. if (auto *MI = dyn_cast<AnyMemIntrinsic>(II)) { bool Changed = false; @@ -1908,6 +1857,18 @@ Instruction *InstCombiner::visitCallInst(CallInst &CI) { if (Changed) return II; } + // For vector result intrinsics, use the generic demanded vector support. + if (II->getType()->isVectorTy()) { + auto VWidth = II->getType()->getVectorNumElements(); + APInt UndefElts(VWidth, 0); + APInt AllOnesEltMask(APInt::getAllOnesValue(VWidth)); + if (Value *V = SimplifyDemandedVectorElts(II, AllOnesEltMask, UndefElts)) { + if (V != II) + return replaceInstUsesWith(*II, V); + return II; + } + } + if (Instruction *I = SimplifyNVVMIntrinsic(II, *this)) return I; @@ -1918,12 +1879,12 @@ Instruction *InstCombiner::visitCallInst(CallInst &CI) { return SimplifyDemandedVectorElts(Op, DemandedElts, UndefElts); }; - switch (II->getIntrinsicID()) { + Intrinsic::ID IID = II->getIntrinsicID(); + switch (IID) { default: break; case Intrinsic::objectsize: - if (ConstantInt *N = - lowerObjectSizeCall(II, DL, &TLI, /*MustSucceed=*/false)) - return replaceInstUsesWith(CI, N); + if (Value *V = lowerObjectSizeCall(II, DL, &TLI, /*MustSucceed=*/false)) + return replaceInstUsesWith(CI, V); return nullptr; case Intrinsic::bswap: { Value *IIOperand = II->getArgOperand(0); @@ -1940,15 +1901,15 @@ Instruction *InstCombiner::visitCallInst(CallInst &CI) { break; } case Intrinsic::masked_load: - if (Value *SimplifiedMaskedOp = simplifyMaskedLoad(*II, Builder)) + if (Value *SimplifiedMaskedOp = simplifyMaskedLoad(*II)) return replaceInstUsesWith(CI, SimplifiedMaskedOp); break; case Intrinsic::masked_store: - return simplifyMaskedStore(*II, *this); + return simplifyMaskedStore(*II); case Intrinsic::masked_gather: - return simplifyMaskedGather(*II, *this); + return simplifyMaskedGather(*II); case Intrinsic::masked_scatter: - return simplifyMaskedScatter(*II, *this); + return simplifyMaskedScatter(*II); case Intrinsic::launder_invariant_group: case Intrinsic::strip_invariant_group: if (auto *SkippedBarrier = simplifyInvariantGroupIntrinsic(*II, *this)) @@ -1982,33 +1943,62 @@ Instruction *InstCombiner::visitCallInst(CallInst &CI) { case Intrinsic::fshl: case Intrinsic::fshr: { - const APInt *SA; - if (match(II->getArgOperand(2), m_APInt(SA))) { - Value *Op0 = II->getArgOperand(0), *Op1 = II->getArgOperand(1); - unsigned BitWidth = SA->getBitWidth(); - uint64_t ShiftAmt = SA->urem(BitWidth); - assert(ShiftAmt != 0 && "SimplifyCall should have handled zero shift"); - // Normalize to funnel shift left. - if (II->getIntrinsicID() == Intrinsic::fshr) - ShiftAmt = BitWidth - ShiftAmt; + Value *Op0 = II->getArgOperand(0), *Op1 = II->getArgOperand(1); + Type *Ty = II->getType(); + unsigned BitWidth = Ty->getScalarSizeInBits(); + Constant *ShAmtC; + if (match(II->getArgOperand(2), m_Constant(ShAmtC)) && + !isa<ConstantExpr>(ShAmtC) && !ShAmtC->containsConstantExpression()) { + // Canonicalize a shift amount constant operand to modulo the bit-width. + Constant *WidthC = ConstantInt::get(Ty, BitWidth); + Constant *ModuloC = ConstantExpr::getURem(ShAmtC, WidthC); + if (ModuloC != ShAmtC) { + II->setArgOperand(2, ModuloC); + return II; + } + assert(ConstantExpr::getICmp(ICmpInst::ICMP_UGT, WidthC, ShAmtC) == + ConstantInt::getTrue(CmpInst::makeCmpResultType(Ty)) && + "Shift amount expected to be modulo bitwidth"); + + // Canonicalize funnel shift right by constant to funnel shift left. This + // is not entirely arbitrary. For historical reasons, the backend may + // recognize rotate left patterns but miss rotate right patterns. + if (IID == Intrinsic::fshr) { + // fshr X, Y, C --> fshl X, Y, (BitWidth - C) + Constant *LeftShiftC = ConstantExpr::getSub(WidthC, ShAmtC); + Module *Mod = II->getModule(); + Function *Fshl = Intrinsic::getDeclaration(Mod, Intrinsic::fshl, Ty); + return CallInst::Create(Fshl, { Op0, Op1, LeftShiftC }); + } + assert(IID == Intrinsic::fshl && + "All funnel shifts by simple constants should go left"); + + // fshl(X, 0, C) --> shl X, C + // fshl(X, undef, C) --> shl X, C + if (match(Op1, m_ZeroInt()) || match(Op1, m_Undef())) + return BinaryOperator::CreateShl(Op0, ShAmtC); - // fshl(X, 0, C) -> shl X, C - // fshl(X, undef, C) -> shl X, C - if (match(Op1, m_Zero()) || match(Op1, m_Undef())) - return BinaryOperator::CreateShl( - Op0, ConstantInt::get(II->getType(), ShiftAmt)); + // fshl(0, X, C) --> lshr X, (BW-C) + // fshl(undef, X, C) --> lshr X, (BW-C) + if (match(Op0, m_ZeroInt()) || match(Op0, m_Undef())) + return BinaryOperator::CreateLShr(Op1, + ConstantExpr::getSub(WidthC, ShAmtC)); - // fshl(0, X, C) -> lshr X, (BW-C) - // fshl(undef, X, C) -> lshr X, (BW-C) - if (match(Op0, m_Zero()) || match(Op0, m_Undef())) - return BinaryOperator::CreateLShr( - Op1, ConstantInt::get(II->getType(), BitWidth - ShiftAmt)); + // fshl i16 X, X, 8 --> bswap i16 X (reduce to more-specific form) + if (Op0 == Op1 && BitWidth == 16 && match(ShAmtC, m_SpecificInt(8))) { + Module *Mod = II->getModule(); + Function *Bswap = Intrinsic::getDeclaration(Mod, Intrinsic::bswap, Ty); + return CallInst::Create(Bswap, { Op0 }); + } } + // Left or right might be masked. + if (SimplifyDemandedInstructionBits(*II)) + return &CI; + // The shift amount (operand 2) of a funnel shift is modulo the bitwidth, // so only the low bits of the shift amount are demanded if the bitwidth is // a power-of-2. - unsigned BitWidth = II->getType()->getScalarSizeInBits(); if (!isPowerOf2_32(BitWidth)) break; APInt Op2Demanded = APInt::getLowBitsSet(BitWidth, Log2_32_Ceil(BitWidth)); @@ -2018,7 +2008,34 @@ Instruction *InstCombiner::visitCallInst(CallInst &CI) { break; } case Intrinsic::uadd_with_overflow: - case Intrinsic::sadd_with_overflow: + case Intrinsic::sadd_with_overflow: { + if (Instruction *I = canonicalizeConstantArg0ToArg1(CI)) + return I; + if (Instruction *I = foldIntrinsicWithOverflowCommon(II)) + return I; + + // Given 2 constant operands whose sum does not overflow: + // uaddo (X +nuw C0), C1 -> uaddo X, C0 + C1 + // saddo (X +nsw C0), C1 -> saddo X, C0 + C1 + Value *X; + const APInt *C0, *C1; + Value *Arg0 = II->getArgOperand(0); + Value *Arg1 = II->getArgOperand(1); + bool IsSigned = IID == Intrinsic::sadd_with_overflow; + bool HasNWAdd = IsSigned ? match(Arg0, m_NSWAdd(m_Value(X), m_APInt(C0))) + : match(Arg0, m_NUWAdd(m_Value(X), m_APInt(C0))); + if (HasNWAdd && match(Arg1, m_APInt(C1))) { + bool Overflow; + APInt NewC = + IsSigned ? C1->sadd_ov(*C0, Overflow) : C1->uadd_ov(*C0, Overflow); + if (!Overflow) + return replaceInstUsesWith( + *II, Builder.CreateBinaryIntrinsic( + IID, X, ConstantInt::get(Arg1->getType(), NewC))); + } + break; + } + case Intrinsic::umul_with_overflow: case Intrinsic::smul_with_overflow: if (Instruction *I = canonicalizeConstantArg0ToArg1(CI)) @@ -2026,16 +2043,29 @@ Instruction *InstCombiner::visitCallInst(CallInst &CI) { LLVM_FALLTHROUGH; case Intrinsic::usub_with_overflow: + if (Instruction *I = foldIntrinsicWithOverflowCommon(II)) + return I; + break; + case Intrinsic::ssub_with_overflow: { - OverflowCheckFlavor OCF = - IntrinsicIDToOverflowCheckFlavor(II->getIntrinsicID()); - assert(OCF != OCF_INVALID && "unexpected!"); + if (Instruction *I = foldIntrinsicWithOverflowCommon(II)) + return I; - Value *OperationResult = nullptr; - Constant *OverflowResult = nullptr; - if (OptimizeOverflowCheck(OCF, II->getArgOperand(0), II->getArgOperand(1), - *II, OperationResult, OverflowResult)) - return CreateOverflowTuple(II, OperationResult, OverflowResult); + Constant *C; + Value *Arg0 = II->getArgOperand(0); + Value *Arg1 = II->getArgOperand(1); + // Given a constant C that is not the minimum signed value + // for an integer of a given bit width: + // + // ssubo X, C -> saddo X, -C + if (match(Arg1, m_Constant(C)) && C->isNotMinSignedValue()) { + Value *NegVal = ConstantExpr::getNeg(C); + // Build a saddo call that is equivalent to the discovered + // ssubo call. + return replaceInstUsesWith( + *II, Builder.CreateBinaryIntrinsic(Intrinsic::sadd_with_overflow, + Arg0, NegVal)); + } break; } @@ -2047,39 +2077,32 @@ Instruction *InstCombiner::visitCallInst(CallInst &CI) { LLVM_FALLTHROUGH; case Intrinsic::usub_sat: case Intrinsic::ssub_sat: { - Value *Arg0 = II->getArgOperand(0); - Value *Arg1 = II->getArgOperand(1); - Intrinsic::ID IID = II->getIntrinsicID(); + SaturatingInst *SI = cast<SaturatingInst>(II); + Type *Ty = SI->getType(); + Value *Arg0 = SI->getLHS(); + Value *Arg1 = SI->getRHS(); // Make use of known overflow information. - OverflowResult OR; - switch (IID) { - default: - llvm_unreachable("Unexpected intrinsic!"); - case Intrinsic::uadd_sat: - OR = computeOverflowForUnsignedAdd(Arg0, Arg1, II); - if (OR == OverflowResult::NeverOverflows) - return BinaryOperator::CreateNUWAdd(Arg0, Arg1); - if (OR == OverflowResult::AlwaysOverflows) - return replaceInstUsesWith(*II, - ConstantInt::getAllOnesValue(II->getType())); - break; - case Intrinsic::usub_sat: - OR = computeOverflowForUnsignedSub(Arg0, Arg1, II); - if (OR == OverflowResult::NeverOverflows) - return BinaryOperator::CreateNUWSub(Arg0, Arg1); - if (OR == OverflowResult::AlwaysOverflows) - return replaceInstUsesWith(*II, - ConstantInt::getNullValue(II->getType())); - break; - case Intrinsic::sadd_sat: - if (willNotOverflowSignedAdd(Arg0, Arg1, *II)) - return BinaryOperator::CreateNSWAdd(Arg0, Arg1); - break; - case Intrinsic::ssub_sat: - if (willNotOverflowSignedSub(Arg0, Arg1, *II)) - return BinaryOperator::CreateNSWSub(Arg0, Arg1); - break; + OverflowResult OR = computeOverflow(SI->getBinaryOp(), SI->isSigned(), + Arg0, Arg1, SI); + switch (OR) { + case OverflowResult::MayOverflow: + break; + case OverflowResult::NeverOverflows: + if (SI->isSigned()) + return BinaryOperator::CreateNSW(SI->getBinaryOp(), Arg0, Arg1); + else + return BinaryOperator::CreateNUW(SI->getBinaryOp(), Arg0, Arg1); + case OverflowResult::AlwaysOverflowsLow: { + unsigned BitWidth = Ty->getScalarSizeInBits(); + APInt Min = APSInt::getMinValue(BitWidth, !SI->isSigned()); + return replaceInstUsesWith(*SI, ConstantInt::get(Ty, Min)); + } + case OverflowResult::AlwaysOverflowsHigh: { + unsigned BitWidth = Ty->getScalarSizeInBits(); + APInt Max = APSInt::getMaxValue(BitWidth, !SI->isSigned()); + return replaceInstUsesWith(*SI, ConstantInt::get(Ty, Max)); + } } // ssub.sat(X, C) -> sadd.sat(X, -C) if C != MIN @@ -2101,7 +2124,7 @@ Instruction *InstCombiner::visitCallInst(CallInst &CI) { APInt NewVal; bool IsUnsigned = IID == Intrinsic::uadd_sat || IID == Intrinsic::usub_sat; - if (Other->getIntrinsicID() == II->getIntrinsicID() && + if (Other->getIntrinsicID() == IID && match(Arg1, m_APInt(Val)) && match(Other->getArgOperand(0), m_Value(X)) && match(Other->getArgOperand(1), m_APInt(Val2))) { @@ -2136,7 +2159,6 @@ Instruction *InstCombiner::visitCallInst(CallInst &CI) { return I; Value *Arg0 = II->getArgOperand(0); Value *Arg1 = II->getArgOperand(1); - Intrinsic::ID IID = II->getIntrinsicID(); Value *X, *Y; if (match(Arg0, m_FNeg(m_Value(X))) && match(Arg1, m_FNeg(m_Value(Y))) && (Arg0->hasOneUse() || Arg1->hasOneUse())) { @@ -2266,8 +2288,7 @@ Instruction *InstCombiner::visitCallInst(CallInst &CI) { Value *ExtSrc; if (match(II->getArgOperand(0), m_OneUse(m_FPExt(m_Value(ExtSrc))))) { // Narrow the call: intrinsic (fpext x) -> fpext (intrinsic x) - Value *NarrowII = - Builder.CreateUnaryIntrinsic(II->getIntrinsicID(), ExtSrc, II); + Value *NarrowII = Builder.CreateUnaryIntrinsic(IID, ExtSrc, II); return new FPExtInst(NarrowII, II->getType()); } break; @@ -2302,7 +2323,7 @@ Instruction *InstCombiner::visitCallInst(CallInst &CI) { &DT) >= 16) { Value *Ptr = Builder.CreateBitCast(II->getArgOperand(0), PointerType::getUnqual(II->getType())); - return new LoadInst(Ptr); + return new LoadInst(II->getType(), Ptr); } break; case Intrinsic::ppc_vsx_lxvw4x: @@ -2310,7 +2331,7 @@ Instruction *InstCombiner::visitCallInst(CallInst &CI) { // Turn PPC VSX loads into normal loads. Value *Ptr = Builder.CreateBitCast(II->getArgOperand(0), PointerType::getUnqual(II->getType())); - return new LoadInst(Ptr, Twine(""), false, 1); + return new LoadInst(II->getType(), Ptr, Twine(""), false, 1); } case Intrinsic::ppc_altivec_stvx: case Intrinsic::ppc_altivec_stvxl: @@ -2338,7 +2359,7 @@ Instruction *InstCombiner::visitCallInst(CallInst &CI) { II->getType()->getVectorNumElements()); Value *Ptr = Builder.CreateBitCast(II->getArgOperand(0), PointerType::getUnqual(VTy)); - Value *Load = Builder.CreateLoad(Ptr); + Value *Load = Builder.CreateLoad(VTy, Ptr); return new FPExtInst(Load, II->getType()); } break; @@ -2348,7 +2369,7 @@ Instruction *InstCombiner::visitCallInst(CallInst &CI) { &DT) >= 32) { Value *Ptr = Builder.CreateBitCast(II->getArgOperand(0), PointerType::getUnqual(II->getType())); - return new LoadInst(Ptr); + return new LoadInst(II->getType(), Ptr); } break; case Intrinsic::ppc_qpx_qvstfs: @@ -2499,22 +2520,6 @@ Instruction *InstCombiner::visitCallInst(CallInst &CI) { break; } - case Intrinsic::x86_sse41_round_ps: - case Intrinsic::x86_sse41_round_pd: - case Intrinsic::x86_avx_round_ps_256: - case Intrinsic::x86_avx_round_pd_256: - case Intrinsic::x86_avx512_mask_rndscale_ps_128: - case Intrinsic::x86_avx512_mask_rndscale_ps_256: - case Intrinsic::x86_avx512_mask_rndscale_ps_512: - case Intrinsic::x86_avx512_mask_rndscale_pd_128: - case Intrinsic::x86_avx512_mask_rndscale_pd_256: - case Intrinsic::x86_avx512_mask_rndscale_pd_512: - case Intrinsic::x86_avx512_mask_rndscale_ss: - case Intrinsic::x86_avx512_mask_rndscale_sd: - if (Value *V = simplifyX86round(*II, Builder)) - return replaceInstUsesWith(*II, V); - break; - case Intrinsic::x86_mmx_pmovmskb: case Intrinsic::x86_sse_movmsk_ps: case Intrinsic::x86_sse2_movmsk_pd: @@ -2620,7 +2625,7 @@ Instruction *InstCombiner::visitCallInst(CallInst &CI) { Value *Arg1 = II->getArgOperand(1); Value *V; - switch (II->getIntrinsicID()) { + switch (IID) { default: llvm_unreachable("Case stmts out of sync!"); case Intrinsic::x86_avx512_add_ps_512: case Intrinsic::x86_avx512_add_pd_512: @@ -2664,7 +2669,7 @@ Instruction *InstCombiner::visitCallInst(CallInst &CI) { Value *RHS = Builder.CreateExtractElement(Arg1, (uint64_t)0); Value *V; - switch (II->getIntrinsicID()) { + switch (IID) { default: llvm_unreachable("Case stmts out of sync!"); case Intrinsic::x86_avx512_mask_add_ss_round: case Intrinsic::x86_avx512_mask_add_sd_round: @@ -2706,44 +2711,7 @@ Instruction *InstCombiner::visitCallInst(CallInst &CI) { return replaceInstUsesWith(*II, V); } } - LLVM_FALLTHROUGH; - - // X86 scalar intrinsics simplified with SimplifyDemandedVectorElts. - case Intrinsic::x86_avx512_mask_max_ss_round: - case Intrinsic::x86_avx512_mask_min_ss_round: - case Intrinsic::x86_avx512_mask_max_sd_round: - case Intrinsic::x86_avx512_mask_min_sd_round: - case Intrinsic::x86_sse_cmp_ss: - case Intrinsic::x86_sse_min_ss: - case Intrinsic::x86_sse_max_ss: - case Intrinsic::x86_sse2_cmp_sd: - case Intrinsic::x86_sse2_min_sd: - case Intrinsic::x86_sse2_max_sd: - case Intrinsic::x86_xop_vfrcz_ss: - case Intrinsic::x86_xop_vfrcz_sd: { - unsigned VWidth = II->getType()->getVectorNumElements(); - APInt UndefElts(VWidth, 0); - APInt AllOnesEltMask(APInt::getAllOnesValue(VWidth)); - if (Value *V = SimplifyDemandedVectorElts(II, AllOnesEltMask, UndefElts)) { - if (V != II) - return replaceInstUsesWith(*II, V); - return II; - } - break; - } - case Intrinsic::x86_sse41_round_ss: - case Intrinsic::x86_sse41_round_sd: { - unsigned VWidth = II->getType()->getVectorNumElements(); - APInt UndefElts(VWidth, 0); - APInt AllOnesEltMask(APInt::getAllOnesValue(VWidth)); - if (Value *V = SimplifyDemandedVectorElts(II, AllOnesEltMask, UndefElts)) { - if (V != II) - return replaceInstUsesWith(*II, V); - return II; - } else if (Value *V = simplifyX86round(*II, Builder)) - return replaceInstUsesWith(*II, V); break; - } // Constant fold ashr( <A x Bi>, Ci ). // Constant fold lshr( <A x Bi>, Ci ). @@ -2860,7 +2828,7 @@ Instruction *InstCombiner::visitCallInst(CallInst &CI) { case Intrinsic::x86_avx2_packsswb: case Intrinsic::x86_avx512_packssdw_512: case Intrinsic::x86_avx512_packsswb_512: - if (Value *V = simplifyX86pack(*II, true)) + if (Value *V = simplifyX86pack(*II, Builder, true)) return replaceInstUsesWith(*II, V); break; @@ -2870,7 +2838,7 @@ Instruction *InstCombiner::visitCallInst(CallInst &CI) { case Intrinsic::x86_avx2_packuswb: case Intrinsic::x86_avx512_packusdw_512: case Intrinsic::x86_avx512_packuswb_512: - if (Value *V = simplifyX86pack(*II, false)) + if (Value *V = simplifyX86pack(*II, Builder, false)) return replaceInstUsesWith(*II, V); break; @@ -3168,19 +3136,9 @@ Instruction *InstCombiner::visitCallInst(CallInst &CI) { return nullptr; break; - case Intrinsic::x86_xop_vpcomb: - case Intrinsic::x86_xop_vpcomd: - case Intrinsic::x86_xop_vpcomq: - case Intrinsic::x86_xop_vpcomw: - if (Value *V = simplifyX86vpcom(*II, Builder, true)) - return replaceInstUsesWith(*II, V); - break; - - case Intrinsic::x86_xop_vpcomub: - case Intrinsic::x86_xop_vpcomud: - case Intrinsic::x86_xop_vpcomuq: - case Intrinsic::x86_xop_vpcomuw: - if (Value *V = simplifyX86vpcom(*II, Builder, false)) + case Intrinsic::x86_addcarry_32: + case Intrinsic::x86_addcarry_64: + if (Value *V = simplifyX86addcarry(*II, Builder)) return replaceInstUsesWith(*II, V); break; @@ -3296,8 +3254,8 @@ Instruction *InstCombiner::visitCallInst(CallInst &CI) { } // Check for constant LHS & RHS - in this case we just simplify. - bool Zext = (II->getIntrinsicID() == Intrinsic::arm_neon_vmullu || - II->getIntrinsicID() == Intrinsic::aarch64_neon_umull); + bool Zext = (IID == Intrinsic::arm_neon_vmullu || + IID == Intrinsic::aarch64_neon_umull); VectorType *NewVT = cast<VectorType>(II->getType()); if (Constant *CV0 = dyn_cast<Constant>(Arg0)) { if (Constant *CV1 = dyn_cast<Constant>(Arg1)) { @@ -3374,7 +3332,7 @@ Instruction *InstCombiner::visitCallInst(CallInst &CI) { APFloat Significand = frexp(C->getValueAPF(), Exp, APFloat::rmNearestTiesToEven); - if (II->getIntrinsicID() == Intrinsic::amdgcn_frexp_mant) { + if (IID == Intrinsic::amdgcn_frexp_mant) { return replaceInstUsesWith(CI, ConstantFP::get(II->getContext(), Significand)); } @@ -3559,7 +3517,7 @@ Instruction *InstCombiner::visitCallInst(CallInst &CI) { } } - bool Signed = II->getIntrinsicID() == Intrinsic::amdgcn_sbfe; + bool Signed = IID == Intrinsic::amdgcn_sbfe; if (!CWidth || !COffset) break; @@ -3587,15 +3545,12 @@ Instruction *InstCombiner::visitCallInst(CallInst &CI) { } case Intrinsic::amdgcn_exp: case Intrinsic::amdgcn_exp_compr: { - ConstantInt *En = dyn_cast<ConstantInt>(II->getArgOperand(1)); - if (!En) // Illegal. - break; - + ConstantInt *En = cast<ConstantInt>(II->getArgOperand(1)); unsigned EnBits = En->getZExtValue(); if (EnBits == 0xf) break; // All inputs enabled. - bool IsCompr = II->getIntrinsicID() == Intrinsic::amdgcn_exp_compr; + bool IsCompr = IID == Intrinsic::amdgcn_exp_compr; bool Changed = false; for (int I = 0; I < (IsCompr ? 2 : 4); ++I) { if ((!IsCompr && (EnBits & (1 << I)) == 0) || @@ -3680,13 +3635,10 @@ Instruction *InstCombiner::visitCallInst(CallInst &CI) { } case Intrinsic::amdgcn_icmp: case Intrinsic::amdgcn_fcmp: { - const ConstantInt *CC = dyn_cast<ConstantInt>(II->getArgOperand(2)); - if (!CC) - break; - + const ConstantInt *CC = cast<ConstantInt>(II->getArgOperand(2)); // Guard against invalid arguments. int64_t CCVal = CC->getZExtValue(); - bool IsInteger = II->getIntrinsicID() == Intrinsic::amdgcn_icmp; + bool IsInteger = IID == Intrinsic::amdgcn_icmp; if ((IsInteger && (CCVal < CmpInst::FIRST_ICMP_PREDICATE || CCVal > CmpInst::LAST_ICMP_PREDICATE)) || (!IsInteger && (CCVal < CmpInst::FIRST_FCMP_PREDICATE || @@ -3709,7 +3661,7 @@ Instruction *InstCombiner::visitCallInst(CallInst &CI) { // register (which contains the bitmask of live threads). So a // comparison that always returns true is the same as a read of the // EXEC register. - Value *NewF = Intrinsic::getDeclaration( + Function *NewF = Intrinsic::getDeclaration( II->getModule(), Intrinsic::read_register, II->getType()); Metadata *MDArgs[] = {MDString::get(II->getContext(), "exec")}; MDNode *MD = MDNode::get(II->getContext(), MDArgs); @@ -3804,8 +3756,10 @@ Instruction *InstCombiner::visitCallInst(CallInst &CI) { } else if (!Ty->isFloatTy() && !Ty->isDoubleTy() && !Ty->isHalfTy()) break; - Value *NewF = Intrinsic::getDeclaration(II->getModule(), NewIID, - SrcLHS->getType()); + Function *NewF = + Intrinsic::getDeclaration(II->getModule(), NewIID, + { II->getType(), + SrcLHS->getType() }); Value *Args[] = { SrcLHS, SrcRHS, ConstantInt::get(CC->getType(), SrcPred) }; CallInst *NewCall = Builder.CreateCall(NewF, Args); @@ -3833,11 +3787,10 @@ Instruction *InstCombiner::visitCallInst(CallInst &CI) { case Intrinsic::amdgcn_update_dpp: { Value *Old = II->getArgOperand(0); - auto BC = dyn_cast<ConstantInt>(II->getArgOperand(5)); - auto RM = dyn_cast<ConstantInt>(II->getArgOperand(3)); - auto BM = dyn_cast<ConstantInt>(II->getArgOperand(4)); - if (!BC || !RM || !BM || - BC->isZeroValue() || + auto BC = cast<ConstantInt>(II->getArgOperand(5)); + auto RM = cast<ConstantInt>(II->getArgOperand(3)); + auto BM = cast<ConstantInt>(II->getArgOperand(4)); + if (BC->isZeroValue() || RM->getZExtValue() != 0xF || BM->getZExtValue() != 0xF || isa<UndefValue>(Old)) @@ -3847,6 +3800,37 @@ Instruction *InstCombiner::visitCallInst(CallInst &CI) { II->setOperand(0, UndefValue::get(Old->getType())); return II; } + case Intrinsic::amdgcn_readfirstlane: + case Intrinsic::amdgcn_readlane: { + // A constant value is trivially uniform. + if (Constant *C = dyn_cast<Constant>(II->getArgOperand(0))) + return replaceInstUsesWith(*II, C); + + // The rest of these may not be safe if the exec may not be the same between + // the def and use. + Value *Src = II->getArgOperand(0); + Instruction *SrcInst = dyn_cast<Instruction>(Src); + if (SrcInst && SrcInst->getParent() != II->getParent()) + break; + + // readfirstlane (readfirstlane x) -> readfirstlane x + // readlane (readfirstlane x), y -> readfirstlane x + if (match(Src, m_Intrinsic<Intrinsic::amdgcn_readfirstlane>())) + return replaceInstUsesWith(*II, Src); + + if (IID == Intrinsic::amdgcn_readfirstlane) { + // readfirstlane (readlane x, y) -> readlane x, y + if (match(Src, m_Intrinsic<Intrinsic::amdgcn_readlane>())) + return replaceInstUsesWith(*II, Src); + } else { + // readlane (readlane x, y), y -> readlane x, y + if (match(Src, m_Intrinsic<Intrinsic::amdgcn_readlane>( + m_Value(), m_Specific(II->getArgOperand(1))))) + return replaceInstUsesWith(*II, Src); + } + + break; + } case Intrinsic::stackrestore: { // If the save is right next to the restore, remove the restore. This can // happen when variable allocas are DCE'd. @@ -3870,14 +3854,14 @@ Instruction *InstCombiner::visitCallInst(CallInst &CI) { break; } if (CallInst *BCI = dyn_cast<CallInst>(BI)) { - if (IntrinsicInst *II = dyn_cast<IntrinsicInst>(BCI)) { + if (auto *II2 = dyn_cast<IntrinsicInst>(BCI)) { // If there is a stackrestore below this one, remove this one. - if (II->getIntrinsicID() == Intrinsic::stackrestore) + if (II2->getIntrinsicID() == Intrinsic::stackrestore) return eraseInstFromFunction(CI); // Bail if we cross over an intrinsic with side effects, such as // llvm.stacksave, llvm.read_register, or llvm.setjmp. - if (II->mayHaveSideEffects()) { + if (II2->mayHaveSideEffects()) { CannotRemove = true; break; } @@ -3920,16 +3904,20 @@ Instruction *InstCombiner::visitCallInst(CallInst &CI) { // Canonicalize assume(a && b) -> assume(a); assume(b); // Note: New assumption intrinsics created here are registered by // the InstCombineIRInserter object. - Value *AssumeIntrinsic = II->getCalledValue(), *A, *B; + FunctionType *AssumeIntrinsicTy = II->getFunctionType(); + Value *AssumeIntrinsic = II->getCalledValue(); + Value *A, *B; if (match(IIOperand, m_And(m_Value(A), m_Value(B)))) { - Builder.CreateCall(AssumeIntrinsic, A, II->getName()); - Builder.CreateCall(AssumeIntrinsic, B, II->getName()); + Builder.CreateCall(AssumeIntrinsicTy, AssumeIntrinsic, A, II->getName()); + Builder.CreateCall(AssumeIntrinsicTy, AssumeIntrinsic, B, II->getName()); return eraseInstFromFunction(*II); } // assume(!(a || b)) -> assume(!a); assume(!b); if (match(IIOperand, m_Not(m_Or(m_Value(A), m_Value(B))))) { - Builder.CreateCall(AssumeIntrinsic, Builder.CreateNot(A), II->getName()); - Builder.CreateCall(AssumeIntrinsic, Builder.CreateNot(B), II->getName()); + Builder.CreateCall(AssumeIntrinsicTy, AssumeIntrinsic, + Builder.CreateNot(A), II->getName()); + Builder.CreateCall(AssumeIntrinsicTy, AssumeIntrinsic, + Builder.CreateNot(B), II->getName()); return eraseInstFromFunction(*II); } @@ -4036,7 +4024,7 @@ Instruction *InstCombiner::visitCallInst(CallInst &CI) { break; } } - return visitCallSite(II); + return visitCallBase(*II); } // Fence instruction simplification @@ -4051,12 +4039,17 @@ Instruction *InstCombiner::visitFenceInst(FenceInst &FI) { // InvokeInst simplification Instruction *InstCombiner::visitInvokeInst(InvokeInst &II) { - return visitCallSite(&II); + return visitCallBase(II); +} + +// CallBrInst simplification +Instruction *InstCombiner::visitCallBrInst(CallBrInst &CBI) { + return visitCallBase(CBI); } /// If this cast does not affect the value passed through the varargs area, we /// can eliminate the use of the cast. -static bool isSafeToEliminateVarargsCast(const CallSite CS, +static bool isSafeToEliminateVarargsCast(const CallBase &Call, const DataLayout &DL, const CastInst *const CI, const int ix) { @@ -4068,18 +4061,20 @@ static bool isSafeToEliminateVarargsCast(const CallSite CS, // TODO: This is probably something which should be expanded to all // intrinsics since the entire point of intrinsics is that // they are understandable by the optimizer. - if (isStatepoint(CS) || isGCRelocate(CS) || isGCResult(CS)) + if (isStatepoint(&Call) || isGCRelocate(&Call) || isGCResult(&Call)) return false; // The size of ByVal or InAlloca arguments is derived from the type, so we // can't change to a type with a different size. If the size were // passed explicitly we could avoid this check. - if (!CS.isByValOrInAllocaArgument(ix)) + if (!Call.isByValOrInAllocaArgument(ix)) return true; Type* SrcTy = cast<PointerType>(CI->getOperand(0)->getType())->getElementType(); - Type* DstTy = cast<PointerType>(CI->getType())->getElementType(); + Type *DstTy = Call.isByValArgument(ix) + ? Call.getParamByValType(ix) + : cast<PointerType>(CI->getType())->getElementType(); if (!SrcTy->isSized() || !DstTy->isSized()) return false; if (DL.getTypeAllocSize(SrcTy) != DL.getTypeAllocSize(DstTy)) @@ -4096,7 +4091,7 @@ Instruction *InstCombiner::tryOptimizeCall(CallInst *CI) { auto InstCombineErase = [this](Instruction *I) { eraseInstFromFunction(*I); }; - LibCallSimplifier Simplifier(DL, &TLI, ORE, InstCombineRAUW, + LibCallSimplifier Simplifier(DL, &TLI, ORE, BFI, PSI, InstCombineRAUW, InstCombineErase); if (Value *With = Simplifier.optimizeCall(CI)) { ++NumSimplified; @@ -4182,10 +4177,10 @@ static IntrinsicInst *findInitTrampoline(Value *Callee) { return nullptr; } -/// Improvements for call and invoke instructions. -Instruction *InstCombiner::visitCallSite(CallSite CS) { - if (isAllocLikeFn(CS.getInstruction(), &TLI)) - return visitAllocSite(*CS.getInstruction()); +/// Improvements for call, callbr and invoke instructions. +Instruction *InstCombiner::visitCallBase(CallBase &Call) { + if (isAllocLikeFn(&Call, &TLI)) + return visitAllocSite(Call); bool Changed = false; @@ -4195,52 +4190,50 @@ Instruction *InstCombiner::visitCallSite(CallSite CS) { SmallVector<unsigned, 4> ArgNos; unsigned ArgNo = 0; - for (Value *V : CS.args()) { + for (Value *V : Call.args()) { if (V->getType()->isPointerTy() && - !CS.paramHasAttr(ArgNo, Attribute::NonNull) && - isKnownNonZero(V, DL, 0, &AC, CS.getInstruction(), &DT)) + !Call.paramHasAttr(ArgNo, Attribute::NonNull) && + isKnownNonZero(V, DL, 0, &AC, &Call, &DT)) ArgNos.push_back(ArgNo); ArgNo++; } - assert(ArgNo == CS.arg_size() && "sanity check"); + assert(ArgNo == Call.arg_size() && "sanity check"); if (!ArgNos.empty()) { - AttributeList AS = CS.getAttributes(); - LLVMContext &Ctx = CS.getInstruction()->getContext(); + AttributeList AS = Call.getAttributes(); + LLVMContext &Ctx = Call.getContext(); AS = AS.addParamAttribute(Ctx, ArgNos, Attribute::get(Ctx, Attribute::NonNull)); - CS.setAttributes(AS); + Call.setAttributes(AS); Changed = true; } // If the callee is a pointer to a function, attempt to move any casts to the - // arguments of the call/invoke. - Value *Callee = CS.getCalledValue(); - if (!isa<Function>(Callee) && transformConstExprCastCall(CS)) + // arguments of the call/callbr/invoke. + Value *Callee = Call.getCalledValue(); + if (!isa<Function>(Callee) && transformConstExprCastCall(Call)) return nullptr; if (Function *CalleeF = dyn_cast<Function>(Callee)) { // Remove the convergent attr on calls when the callee is not convergent. - if (CS.isConvergent() && !CalleeF->isConvergent() && + if (Call.isConvergent() && !CalleeF->isConvergent() && !CalleeF->isIntrinsic()) { - LLVM_DEBUG(dbgs() << "Removing convergent attr from instr " - << CS.getInstruction() << "\n"); - CS.setNotConvergent(); - return CS.getInstruction(); + LLVM_DEBUG(dbgs() << "Removing convergent attr from instr " << Call + << "\n"); + Call.setNotConvergent(); + return &Call; } // If the call and callee calling conventions don't match, this call must // be unreachable, as the call is undefined. - if (CalleeF->getCallingConv() != CS.getCallingConv() && + if (CalleeF->getCallingConv() != Call.getCallingConv() && // Only do this for calls to a function with a body. A prototype may // not actually end up matching the implementation's calling conv for a // variety of reasons (e.g. it may be written in assembly). !CalleeF->isDeclaration()) { - Instruction *OldCall = CS.getInstruction(); - new StoreInst(ConstantInt::getTrue(Callee->getContext()), - UndefValue::get(Type::getInt1PtrTy(Callee->getContext())), - OldCall); + Instruction *OldCall = &Call; + CreateNonTerminatorUnreachable(OldCall); // If OldCall does not return void then replaceAllUsesWith undef. // This allows ValueHandlers and custom metadata to adjust itself. if (!OldCall->getType()->isVoidTy()) @@ -4248,40 +4241,35 @@ Instruction *InstCombiner::visitCallSite(CallSite CS) { if (isa<CallInst>(OldCall)) return eraseInstFromFunction(*OldCall); - // We cannot remove an invoke, because it would change the CFG, just - // change the callee to a null pointer. - cast<InvokeInst>(OldCall)->setCalledFunction( - Constant::getNullValue(CalleeF->getType())); + // We cannot remove an invoke or a callbr, because it would change thexi + // CFG, just change the callee to a null pointer. + cast<CallBase>(OldCall)->setCalledFunction( + CalleeF->getFunctionType(), + Constant::getNullValue(CalleeF->getType())); return nullptr; } } if ((isa<ConstantPointerNull>(Callee) && - !NullPointerIsDefined(CS.getInstruction()->getFunction())) || + !NullPointerIsDefined(Call.getFunction())) || isa<UndefValue>(Callee)) { - // If CS does not return void then replaceAllUsesWith undef. + // If Call does not return void then replaceAllUsesWith undef. // This allows ValueHandlers and custom metadata to adjust itself. - if (!CS.getInstruction()->getType()->isVoidTy()) - replaceInstUsesWith(*CS.getInstruction(), - UndefValue::get(CS.getInstruction()->getType())); + if (!Call.getType()->isVoidTy()) + replaceInstUsesWith(Call, UndefValue::get(Call.getType())); - if (isa<InvokeInst>(CS.getInstruction())) { - // Can't remove an invoke because we cannot change the CFG. + if (Call.isTerminator()) { + // Can't remove an invoke or callbr because we cannot change the CFG. return nullptr; } - // This instruction is not reachable, just remove it. We insert a store to - // undef so that we know that this code is not reachable, despite the fact - // that we can't modify the CFG here. - new StoreInst(ConstantInt::getTrue(Callee->getContext()), - UndefValue::get(Type::getInt1PtrTy(Callee->getContext())), - CS.getInstruction()); - - return eraseInstFromFunction(*CS.getInstruction()); + // This instruction is not reachable, just remove it. + CreateNonTerminatorUnreachable(&Call); + return eraseInstFromFunction(Call); } if (IntrinsicInst *II = findInitTrampoline(Callee)) - return transformCallThroughTrampoline(CS, II); + return transformCallThroughTrampoline(Call, *II); PointerType *PTy = cast<PointerType>(Callee->getType()); FunctionType *FTy = cast<FunctionType>(PTy->getElementType()); @@ -4289,39 +4277,48 @@ Instruction *InstCombiner::visitCallSite(CallSite CS) { int ix = FTy->getNumParams(); // See if we can optimize any arguments passed through the varargs area of // the call. - for (CallSite::arg_iterator I = CS.arg_begin() + FTy->getNumParams(), - E = CS.arg_end(); I != E; ++I, ++ix) { + for (auto I = Call.arg_begin() + FTy->getNumParams(), E = Call.arg_end(); + I != E; ++I, ++ix) { CastInst *CI = dyn_cast<CastInst>(*I); - if (CI && isSafeToEliminateVarargsCast(CS, DL, CI, ix)) { + if (CI && isSafeToEliminateVarargsCast(Call, DL, CI, ix)) { *I = CI->getOperand(0); + + // Update the byval type to match the argument type. + if (Call.isByValArgument(ix)) { + Call.removeParamAttr(ix, Attribute::ByVal); + Call.addParamAttr( + ix, Attribute::getWithByValType( + Call.getContext(), + CI->getOperand(0)->getType()->getPointerElementType())); + } Changed = true; } } } - if (isa<InlineAsm>(Callee) && !CS.doesNotThrow()) { + if (isa<InlineAsm>(Callee) && !Call.doesNotThrow()) { // Inline asm calls cannot throw - mark them 'nounwind'. - CS.setDoesNotThrow(); + Call.setDoesNotThrow(); Changed = true; } // Try to optimize the call if possible, we require DataLayout for most of // this. None of these calls are seen as possibly dead so go ahead and // delete the instruction now. - if (CallInst *CI = dyn_cast<CallInst>(CS.getInstruction())) { + if (CallInst *CI = dyn_cast<CallInst>(&Call)) { Instruction *I = tryOptimizeCall(CI); // If we changed something return the result, etc. Otherwise let // the fallthrough check. if (I) return eraseInstFromFunction(*I); } - return Changed ? CS.getInstruction() : nullptr; + return Changed ? &Call : nullptr; } /// If the callee is a constexpr cast of a function, attempt to move the cast to -/// the arguments of the call/invoke. -bool InstCombiner::transformConstExprCastCall(CallSite CS) { - auto *Callee = dyn_cast<Function>(CS.getCalledValue()->stripPointerCasts()); +/// the arguments of the call/callbr/invoke. +bool InstCombiner::transformConstExprCastCall(CallBase &Call) { + auto *Callee = dyn_cast<Function>(Call.getCalledValue()->stripPointerCasts()); if (!Callee) return false; @@ -4335,11 +4332,11 @@ bool InstCombiner::transformConstExprCastCall(CallSite CS) { // prototype with the exception of pointee types. The code below doesn't // implement that, so we can't do this transform. // TODO: Do the transform if it only requires adding pointer casts. - if (CS.isMustTailCall()) + if (Call.isMustTailCall()) return false; - Instruction *Caller = CS.getInstruction(); - const AttributeList &CallerPAL = CS.getAttributes(); + Instruction *Caller = &Call; + const AttributeList &CallerPAL = Call.getAttributes(); // Okay, this is a cast from a function to a different type. Unless doing so // would cause a type conversion of one of our arguments, change this call to @@ -4370,20 +4367,24 @@ bool InstCombiner::transformConstExprCastCall(CallSite CS) { return false; // Attribute not compatible with transformed value. } - // If the callsite is an invoke instruction, and the return value is used by - // a PHI node in a successor, we cannot change the return type of the call - // because there is no place to put the cast instruction (without breaking - // the critical edge). Bail out in this case. - if (!Caller->use_empty()) + // If the callbase is an invoke/callbr instruction, and the return value is + // used by a PHI node in a successor, we cannot change the return type of + // the call because there is no place to put the cast instruction (without + // breaking the critical edge). Bail out in this case. + if (!Caller->use_empty()) { if (InvokeInst *II = dyn_cast<InvokeInst>(Caller)) for (User *U : II->users()) if (PHINode *PN = dyn_cast<PHINode>(U)) if (PN->getParent() == II->getNormalDest() || PN->getParent() == II->getUnwindDest()) return false; + // FIXME: Be conservative for callbr to avoid a quadratic search. + if (isa<CallBrInst>(Caller)) + return false; + } } - unsigned NumActualArgs = CS.arg_size(); + unsigned NumActualArgs = Call.arg_size(); unsigned NumCommonArgs = std::min(FT->getNumParams(), NumActualArgs); // Prevent us turning: @@ -4398,7 +4399,7 @@ bool InstCombiner::transformConstExprCastCall(CallSite CS) { Callee->getAttributes().hasAttrSomewhere(Attribute::ByVal)) return false; - CallSite::arg_iterator AI = CS.arg_begin(); + auto AI = Call.arg_begin(); for (unsigned i = 0, e = NumCommonArgs; i != e; ++i, ++AI) { Type *ParamTy = FT->getParamType(i); Type *ActTy = (*AI)->getType(); @@ -4410,7 +4411,7 @@ bool InstCombiner::transformConstExprCastCall(CallSite CS) { .overlaps(AttributeFuncs::typeIncompatible(ParamTy))) return false; // Attribute not compatible with transformed value. - if (CS.isInAllocaArgument(i)) + if (Call.isInAllocaArgument(i)) return false; // Cannot transform to and from inalloca. // If the parameter is passed as a byval argument, then we have to have a @@ -4420,7 +4421,7 @@ bool InstCombiner::transformConstExprCastCall(CallSite CS) { if (!ParamPTy || !ParamPTy->getElementType()->isSized()) return false; - Type *CurElTy = ActTy->getPointerElementType(); + Type *CurElTy = Call.getParamByValType(i); if (DL.getTypeAllocSize(CurElTy) != DL.getTypeAllocSize(ParamPTy->getElementType())) return false; @@ -4435,7 +4436,7 @@ bool InstCombiner::transformConstExprCastCall(CallSite CS) { // If the callee is just a declaration, don't change the varargsness of the // call. We don't want to introduce a varargs call where one doesn't // already exist. - PointerType *APTy = cast<PointerType>(CS.getCalledValue()->getType()); + PointerType *APTy = cast<PointerType>(Call.getCalledValue()->getType()); if (FT->isVarArg()!=cast<FunctionType>(APTy->getElementType())->isVarArg()) return false; @@ -4474,7 +4475,8 @@ bool InstCombiner::transformConstExprCastCall(CallSite CS) { // with the existing attributes. Wipe out any problematic attributes. RAttrs.remove(AttributeFuncs::typeIncompatible(NewRetTy)); - AI = CS.arg_begin(); + LLVMContext &Ctx = Call.getContext(); + AI = Call.arg_begin(); for (unsigned i = 0; i != NumCommonArgs; ++i, ++AI) { Type *ParamTy = FT->getParamType(i); @@ -4484,7 +4486,12 @@ bool InstCombiner::transformConstExprCastCall(CallSite CS) { Args.push_back(NewArg); // Add any parameter attributes. - ArgAttrs.push_back(CallerPAL.getParamAttributes(i)); + if (CallerPAL.hasParamAttribute(i, Attribute::ByVal)) { + AttrBuilder AB(CallerPAL.getParamAttributes(i)); + AB.addByValAttr(NewArg->getType()->getPointerElementType()); + ArgAttrs.push_back(AttributeSet::get(Ctx, AB)); + } else + ArgAttrs.push_back(CallerPAL.getParamAttributes(i)); } // If the function takes more arguments than the call was taking, add them @@ -4523,45 +4530,50 @@ bool InstCombiner::transformConstExprCastCall(CallSite CS) { assert((ArgAttrs.size() == FT->getNumParams() || FT->isVarArg()) && "missing argument attributes"); - LLVMContext &Ctx = Callee->getContext(); AttributeList NewCallerPAL = AttributeList::get( Ctx, FnAttrs, AttributeSet::get(Ctx, RAttrs), ArgAttrs); SmallVector<OperandBundleDef, 1> OpBundles; - CS.getOperandBundlesAsDefs(OpBundles); + Call.getOperandBundlesAsDefs(OpBundles); - CallSite NewCS; + CallBase *NewCall; if (InvokeInst *II = dyn_cast<InvokeInst>(Caller)) { - NewCS = Builder.CreateInvoke(Callee, II->getNormalDest(), - II->getUnwindDest(), Args, OpBundles); + NewCall = Builder.CreateInvoke(Callee, II->getNormalDest(), + II->getUnwindDest(), Args, OpBundles); + } else if (CallBrInst *CBI = dyn_cast<CallBrInst>(Caller)) { + NewCall = Builder.CreateCallBr(Callee, CBI->getDefaultDest(), + CBI->getIndirectDests(), Args, OpBundles); } else { - NewCS = Builder.CreateCall(Callee, Args, OpBundles); - cast<CallInst>(NewCS.getInstruction()) - ->setTailCallKind(cast<CallInst>(Caller)->getTailCallKind()); + NewCall = Builder.CreateCall(Callee, Args, OpBundles); + cast<CallInst>(NewCall)->setTailCallKind( + cast<CallInst>(Caller)->getTailCallKind()); } - NewCS->takeName(Caller); - NewCS.setCallingConv(CS.getCallingConv()); - NewCS.setAttributes(NewCallerPAL); + NewCall->takeName(Caller); + NewCall->setCallingConv(Call.getCallingConv()); + NewCall->setAttributes(NewCallerPAL); // Preserve the weight metadata for the new call instruction. The metadata // is used by SamplePGO to check callsite's hotness. uint64_t W; if (Caller->extractProfTotalWeight(W)) - NewCS->setProfWeight(W); + NewCall->setProfWeight(W); // Insert a cast of the return type as necessary. - Instruction *NC = NewCS.getInstruction(); + Instruction *NC = NewCall; Value *NV = NC; if (OldRetTy != NV->getType() && !Caller->use_empty()) { if (!NV->getType()->isVoidTy()) { NV = NC = CastInst::CreateBitOrPointerCast(NC, OldRetTy); NC->setDebugLoc(Caller->getDebugLoc()); - // If this is an invoke instruction, we should insert it after the first - // non-phi, instruction in the normal successor block. + // If this is an invoke/callbr instruction, we should insert it after the + // first non-phi instruction in the normal successor block. if (InvokeInst *II = dyn_cast<InvokeInst>(Caller)) { BasicBlock::iterator I = II->getNormalDest()->getFirstInsertionPt(); InsertNewInstBefore(NC, *I); + } else if (CallBrInst *CBI = dyn_cast<CallBrInst>(Caller)) { + BasicBlock::iterator I = CBI->getDefaultDest()->getFirstInsertionPt(); + InsertNewInstBefore(NC, *I); } else { // Otherwise, it's a call, just insert cast right after the call. InsertNewInstBefore(NC, *Caller); @@ -4590,23 +4602,20 @@ bool InstCombiner::transformConstExprCastCall(CallSite CS) { /// Turn a call to a function created by init_trampoline / adjust_trampoline /// intrinsic pair into a direct call to the underlying function. Instruction * -InstCombiner::transformCallThroughTrampoline(CallSite CS, - IntrinsicInst *Tramp) { - Value *Callee = CS.getCalledValue(); - PointerType *PTy = cast<PointerType>(Callee->getType()); - FunctionType *FTy = cast<FunctionType>(PTy->getElementType()); - AttributeList Attrs = CS.getAttributes(); +InstCombiner::transformCallThroughTrampoline(CallBase &Call, + IntrinsicInst &Tramp) { + Value *Callee = Call.getCalledValue(); + Type *CalleeTy = Callee->getType(); + FunctionType *FTy = Call.getFunctionType(); + AttributeList Attrs = Call.getAttributes(); // If the call already has the 'nest' attribute somewhere then give up - // otherwise 'nest' would occur twice after splicing in the chain. if (Attrs.hasAttrSomewhere(Attribute::Nest)) return nullptr; - assert(Tramp && - "transformCallThroughTrampoline called with incorrect CallSite."); - - Function *NestF =cast<Function>(Tramp->getArgOperand(1)->stripPointerCasts()); - FunctionType *NestFTy = cast<FunctionType>(NestF->getValueType()); + Function *NestF = cast<Function>(Tramp.getArgOperand(1)->stripPointerCasts()); + FunctionType *NestFTy = NestF->getFunctionType(); AttributeList NestAttrs = NestF->getAttributes(); if (!NestAttrs.isEmpty()) { @@ -4628,22 +4637,21 @@ InstCombiner::transformCallThroughTrampoline(CallSite CS, } if (NestTy) { - Instruction *Caller = CS.getInstruction(); std::vector<Value*> NewArgs; std::vector<AttributeSet> NewArgAttrs; - NewArgs.reserve(CS.arg_size() + 1); - NewArgAttrs.reserve(CS.arg_size()); + NewArgs.reserve(Call.arg_size() + 1); + NewArgAttrs.reserve(Call.arg_size()); // Insert the nest argument into the call argument list, which may // mean appending it. Likewise for attributes. { unsigned ArgNo = 0; - CallSite::arg_iterator I = CS.arg_begin(), E = CS.arg_end(); + auto I = Call.arg_begin(), E = Call.arg_end(); do { if (ArgNo == NestArgNo) { // Add the chain argument and attributes. - Value *NestVal = Tramp->getArgOperand(2); + Value *NestVal = Tramp.getArgOperand(2); if (NestVal->getType() != NestTy) NestVal = Builder.CreateBitCast(NestVal, NestTy, "nest"); NewArgs.push_back(NestVal); @@ -4705,24 +4713,30 @@ InstCombiner::transformCallThroughTrampoline(CallSite CS, Attrs.getRetAttributes(), NewArgAttrs); SmallVector<OperandBundleDef, 1> OpBundles; - CS.getOperandBundlesAsDefs(OpBundles); + Call.getOperandBundlesAsDefs(OpBundles); Instruction *NewCaller; - if (InvokeInst *II = dyn_cast<InvokeInst>(Caller)) { - NewCaller = InvokeInst::Create(NewCallee, + if (InvokeInst *II = dyn_cast<InvokeInst>(&Call)) { + NewCaller = InvokeInst::Create(NewFTy, NewCallee, II->getNormalDest(), II->getUnwindDest(), NewArgs, OpBundles); cast<InvokeInst>(NewCaller)->setCallingConv(II->getCallingConv()); cast<InvokeInst>(NewCaller)->setAttributes(NewPAL); + } else if (CallBrInst *CBI = dyn_cast<CallBrInst>(&Call)) { + NewCaller = + CallBrInst::Create(NewFTy, NewCallee, CBI->getDefaultDest(), + CBI->getIndirectDests(), NewArgs, OpBundles); + cast<CallBrInst>(NewCaller)->setCallingConv(CBI->getCallingConv()); + cast<CallBrInst>(NewCaller)->setAttributes(NewPAL); } else { - NewCaller = CallInst::Create(NewCallee, NewArgs, OpBundles); + NewCaller = CallInst::Create(NewFTy, NewCallee, NewArgs, OpBundles); cast<CallInst>(NewCaller)->setTailCallKind( - cast<CallInst>(Caller)->getTailCallKind()); + cast<CallInst>(Call).getTailCallKind()); cast<CallInst>(NewCaller)->setCallingConv( - cast<CallInst>(Caller)->getCallingConv()); + cast<CallInst>(Call).getCallingConv()); cast<CallInst>(NewCaller)->setAttributes(NewPAL); } - NewCaller->setDebugLoc(Caller->getDebugLoc()); + NewCaller->setDebugLoc(Call.getDebugLoc()); return NewCaller; } @@ -4731,9 +4745,7 @@ InstCombiner::transformCallThroughTrampoline(CallSite CS, // Replace the trampoline call with a direct call. Since there is no 'nest' // parameter, there is no need to adjust the argument list. Let the generic // code sort out any function type mismatches. - Constant *NewCallee = - NestF->getType() == PTy ? NestF : - ConstantExpr::getBitCast(NestF, PTy); - CS.setCalledFunction(NewCallee); - return CS.getInstruction(); + Constant *NewCallee = ConstantExpr::getBitCast(NestF, CalleeTy); + Call.setCalledFunction(FTy, NewCallee); + return &Call; } |
