diff options
Diffstat (limited to 'llvm/lib/CodeGen/GlobalISel/CombinerHelper.cpp')
| -rw-r--r-- | llvm/lib/CodeGen/GlobalISel/CombinerHelper.cpp | 1019 |
1 files changed, 858 insertions, 161 deletions
diff --git a/llvm/lib/CodeGen/GlobalISel/CombinerHelper.cpp b/llvm/lib/CodeGen/GlobalISel/CombinerHelper.cpp index df0219fcfa64..06d827de2e96 100644 --- a/llvm/lib/CodeGen/GlobalISel/CombinerHelper.cpp +++ b/llvm/lib/CodeGen/GlobalISel/CombinerHelper.cpp @@ -6,13 +6,18 @@ // //===----------------------------------------------------------------------===// #include "llvm/CodeGen/GlobalISel/CombinerHelper.h" +#include "llvm/ADT/SetVector.h" +#include "llvm/ADT/SmallBitVector.h" #include "llvm/CodeGen/GlobalISel/Combiner.h" #include "llvm/CodeGen/GlobalISel/GISelChangeObserver.h" #include "llvm/CodeGen/GlobalISel/GISelKnownBits.h" +#include "llvm/CodeGen/GlobalISel/GenericMachineInstrs.h" #include "llvm/CodeGen/GlobalISel/LegalizerInfo.h" #include "llvm/CodeGen/GlobalISel/MIPatternMatch.h" #include "llvm/CodeGen/GlobalISel/MachineIRBuilder.h" #include "llvm/CodeGen/GlobalISel/Utils.h" +#include "llvm/CodeGen/LowLevelType.h" +#include "llvm/CodeGen/MachineBasicBlock.h" #include "llvm/CodeGen/MachineDominators.h" #include "llvm/CodeGen/MachineFrameInfo.h" #include "llvm/CodeGen/MachineInstr.h" @@ -20,8 +25,10 @@ #include "llvm/CodeGen/MachineRegisterInfo.h" #include "llvm/CodeGen/TargetInstrInfo.h" #include "llvm/CodeGen/TargetLowering.h" +#include "llvm/CodeGen/TargetOpcodes.h" #include "llvm/Support/MathExtras.h" #include "llvm/Target/TargetMachine.h" +#include <tuple> #define DEBUG_TYPE "gi-combiner" @@ -436,16 +443,13 @@ bool CombinerHelper::matchCombineExtendingLoads(MachineInstr &MI, // to find a safe place to sink it) whereas the extend is freely movable. // It also prevents us from duplicating the load for the volatile case or just // for performance. - - if (MI.getOpcode() != TargetOpcode::G_LOAD && - MI.getOpcode() != TargetOpcode::G_SEXTLOAD && - MI.getOpcode() != TargetOpcode::G_ZEXTLOAD) + GAnyLoad *LoadMI = dyn_cast<GAnyLoad>(&MI); + if (!LoadMI) return false; - auto &LoadValue = MI.getOperand(0); - assert(LoadValue.isReg() && "Result wasn't a register?"); + Register LoadReg = LoadMI->getDstReg(); - LLT LoadValueTy = MRI.getType(LoadValue.getReg()); + LLT LoadValueTy = MRI.getType(LoadReg); if (!LoadValueTy.isScalar()) return false; @@ -467,27 +471,29 @@ bool CombinerHelper::matchCombineExtendingLoads(MachineInstr &MI, // and emit a variant of (extend (trunc X)) for the others according to the // relative type sizes. At the same time, pick an extend to use based on the // extend involved in the chosen type. - unsigned PreferredOpcode = MI.getOpcode() == TargetOpcode::G_LOAD - ? TargetOpcode::G_ANYEXT - : MI.getOpcode() == TargetOpcode::G_SEXTLOAD - ? TargetOpcode::G_SEXT - : TargetOpcode::G_ZEXT; + unsigned PreferredOpcode = + isa<GLoad>(&MI) + ? TargetOpcode::G_ANYEXT + : isa<GSExtLoad>(&MI) ? TargetOpcode::G_SEXT : TargetOpcode::G_ZEXT; Preferred = {LLT(), PreferredOpcode, nullptr}; - for (auto &UseMI : MRI.use_nodbg_instructions(LoadValue.getReg())) { + for (auto &UseMI : MRI.use_nodbg_instructions(LoadReg)) { if (UseMI.getOpcode() == TargetOpcode::G_SEXT || UseMI.getOpcode() == TargetOpcode::G_ZEXT || (UseMI.getOpcode() == TargetOpcode::G_ANYEXT)) { + const auto &MMO = LoadMI->getMMO(); + // For atomics, only form anyextending loads. + if (MMO.isAtomic() && UseMI.getOpcode() != TargetOpcode::G_ANYEXT) + continue; // Check for legality. if (LI) { LegalityQuery::MemDesc MMDesc; - const auto &MMO = **MI.memoperands_begin(); - MMDesc.SizeInBits = MMO.getSizeInBits(); + MMDesc.MemoryTy = MMO.getMemoryType(); MMDesc.AlignInBits = MMO.getAlign().value() * 8; - MMDesc.Ordering = MMO.getOrdering(); + MMDesc.Ordering = MMO.getSuccessOrdering(); LLT UseTy = MRI.getType(UseMI.getOperand(0).getReg()); - LLT SrcTy = MRI.getType(MI.getOperand(1).getReg()); - if (LI->getAction({MI.getOpcode(), {UseTy, SrcTy}, {MMDesc}}).Action != - LegalizeActions::Legal) + LLT SrcTy = MRI.getType(LoadMI->getPointerReg()); + if (LI->getAction({LoadMI->getOpcode(), {UseTy, SrcTy}, {MMDesc}}) + .Action != LegalizeActions::Legal) continue; } Preferred = ChoosePreferredUse(Preferred, @@ -660,23 +666,22 @@ bool CombinerHelper::matchSextTruncSextLoad(MachineInstr &MI) { uint64_t SizeInBits = MI.getOperand(2).getImm(); // If the source is a G_SEXTLOAD from the same bit width, then we don't // need any extend at all, just a truncate. - if (auto *LoadMI = getOpcodeDef(TargetOpcode::G_SEXTLOAD, LoadUser, MRI)) { - const auto &MMO = **LoadMI->memoperands_begin(); + if (auto *LoadMI = getOpcodeDef<GSExtLoad>(LoadUser, MRI)) { // If truncating more than the original extended value, abort. - if (TruncSrc && MRI.getType(TruncSrc).getSizeInBits() < MMO.getSizeInBits()) + auto LoadSizeBits = LoadMI->getMemSizeInBits(); + if (TruncSrc && MRI.getType(TruncSrc).getSizeInBits() < LoadSizeBits) return false; - if (MMO.getSizeInBits() == SizeInBits) + if (LoadSizeBits == SizeInBits) return true; } return false; } -bool CombinerHelper::applySextTruncSextLoad(MachineInstr &MI) { +void CombinerHelper::applySextTruncSextLoad(MachineInstr &MI) { assert(MI.getOpcode() == TargetOpcode::G_SEXT_INREG); Builder.setInstrAndDebugLoc(MI); Builder.buildCopy(MI.getOperand(0).getReg(), MI.getOperand(1).getReg()); MI.eraseFromParent(); - return true; } bool CombinerHelper::matchSextInRegOfLoad( @@ -688,20 +693,16 @@ bool CombinerHelper::matchSextInRegOfLoad( return false; Register SrcReg = MI.getOperand(1).getReg(); - MachineInstr *LoadDef = getOpcodeDef(TargetOpcode::G_LOAD, SrcReg, MRI); - if (!LoadDef || !MRI.hasOneNonDBGUse(LoadDef->getOperand(0).getReg())) + auto *LoadDef = getOpcodeDef<GLoad>(SrcReg, MRI); + if (!LoadDef || !MRI.hasOneNonDBGUse(LoadDef->getOperand(0).getReg()) || + !LoadDef->isSimple()) return false; // If the sign extend extends from a narrower width than the load's width, // then we can narrow the load width when we combine to a G_SEXTLOAD. - auto &MMO = **LoadDef->memoperands_begin(); - // Don't do this for non-simple loads. - if (MMO.isAtomic() || MMO.isVolatile()) - return false; - // Avoid widening the load at all. - unsigned NewSizeBits = - std::min((uint64_t)MI.getOperand(2).getImm(), MMO.getSizeInBits()); + unsigned NewSizeBits = std::min((uint64_t)MI.getOperand(2).getImm(), + LoadDef->getMemSizeInBits()); // Don't generate G_SEXTLOADs with a < 1 byte width. if (NewSizeBits < 8) @@ -710,18 +711,17 @@ bool CombinerHelper::matchSextInRegOfLoad( // anyway for most targets. if (!isPowerOf2_32(NewSizeBits)) return false; - MatchInfo = std::make_tuple(LoadDef->getOperand(0).getReg(), NewSizeBits); + MatchInfo = std::make_tuple(LoadDef->getDstReg(), NewSizeBits); return true; } -bool CombinerHelper::applySextInRegOfLoad( +void CombinerHelper::applySextInRegOfLoad( MachineInstr &MI, std::tuple<Register, unsigned> &MatchInfo) { assert(MI.getOpcode() == TargetOpcode::G_SEXT_INREG); Register LoadReg; unsigned ScalarSizeBits; std::tie(LoadReg, ScalarSizeBits) = MatchInfo; - auto *LoadDef = MRI.getVRegDef(LoadReg); - assert(LoadDef && "Expected a load reg"); + GLoad *LoadDef = cast<GLoad>(MRI.getVRegDef(LoadReg)); // If we have the following: // %ld = G_LOAD %ptr, (load 2) @@ -729,15 +729,14 @@ bool CombinerHelper::applySextInRegOfLoad( // ==> // %ld = G_SEXTLOAD %ptr (load 1) - auto &MMO = **LoadDef->memoperands_begin(); - Builder.setInstrAndDebugLoc(MI); + auto &MMO = LoadDef->getMMO(); + Builder.setInstrAndDebugLoc(*LoadDef); auto &MF = Builder.getMF(); auto PtrInfo = MMO.getPointerInfo(); auto *NewMMO = MF.getMachineMemOperand(&MMO, PtrInfo, ScalarSizeBits / 8); Builder.buildLoadInstr(TargetOpcode::G_SEXTLOAD, MI.getOperand(0).getReg(), - LoadDef->getOperand(1).getReg(), *NewMMO); + LoadDef->getPointerReg(), *NewMMO); MI.eraseFromParent(); - return true; } bool CombinerHelper::findPostIndexCandidate(MachineInstr &MI, Register &Addr, @@ -941,10 +940,104 @@ void CombinerHelper::applyCombineIndexedLoadStore( LLVM_DEBUG(dbgs() << " Combinined to indexed operation"); } -bool CombinerHelper::matchOptBrCondByInvertingCond(MachineInstr &MI) { - if (MI.getOpcode() != TargetOpcode::G_BR) +bool CombinerHelper::matchCombineDivRem(MachineInstr &MI, + MachineInstr *&OtherMI) { + unsigned Opcode = MI.getOpcode(); + bool IsDiv, IsSigned; + + switch (Opcode) { + default: + llvm_unreachable("Unexpected opcode!"); + case TargetOpcode::G_SDIV: + case TargetOpcode::G_UDIV: { + IsDiv = true; + IsSigned = Opcode == TargetOpcode::G_SDIV; + break; + } + case TargetOpcode::G_SREM: + case TargetOpcode::G_UREM: { + IsDiv = false; + IsSigned = Opcode == TargetOpcode::G_SREM; + break; + } + } + + Register Src1 = MI.getOperand(1).getReg(); + unsigned DivOpcode, RemOpcode, DivremOpcode; + if (IsSigned) { + DivOpcode = TargetOpcode::G_SDIV; + RemOpcode = TargetOpcode::G_SREM; + DivremOpcode = TargetOpcode::G_SDIVREM; + } else { + DivOpcode = TargetOpcode::G_UDIV; + RemOpcode = TargetOpcode::G_UREM; + DivremOpcode = TargetOpcode::G_UDIVREM; + } + + if (!isLegalOrBeforeLegalizer({DivremOpcode, {MRI.getType(Src1)}})) return false; + // Combine: + // %div:_ = G_[SU]DIV %src1:_, %src2:_ + // %rem:_ = G_[SU]REM %src1:_, %src2:_ + // into: + // %div:_, %rem:_ = G_[SU]DIVREM %src1:_, %src2:_ + + // Combine: + // %rem:_ = G_[SU]REM %src1:_, %src2:_ + // %div:_ = G_[SU]DIV %src1:_, %src2:_ + // into: + // %div:_, %rem:_ = G_[SU]DIVREM %src1:_, %src2:_ + + for (auto &UseMI : MRI.use_nodbg_instructions(Src1)) { + if (MI.getParent() == UseMI.getParent() && + ((IsDiv && UseMI.getOpcode() == RemOpcode) || + (!IsDiv && UseMI.getOpcode() == DivOpcode)) && + matchEqualDefs(MI.getOperand(2), UseMI.getOperand(2))) { + OtherMI = &UseMI; + return true; + } + } + + return false; +} + +void CombinerHelper::applyCombineDivRem(MachineInstr &MI, + MachineInstr *&OtherMI) { + unsigned Opcode = MI.getOpcode(); + assert(OtherMI && "OtherMI shouldn't be empty."); + + Register DestDivReg, DestRemReg; + if (Opcode == TargetOpcode::G_SDIV || Opcode == TargetOpcode::G_UDIV) { + DestDivReg = MI.getOperand(0).getReg(); + DestRemReg = OtherMI->getOperand(0).getReg(); + } else { + DestDivReg = OtherMI->getOperand(0).getReg(); + DestRemReg = MI.getOperand(0).getReg(); + } + + bool IsSigned = + Opcode == TargetOpcode::G_SDIV || Opcode == TargetOpcode::G_SREM; + + // Check which instruction is first in the block so we don't break def-use + // deps by "moving" the instruction incorrectly. + if (dominates(MI, *OtherMI)) + Builder.setInstrAndDebugLoc(MI); + else + Builder.setInstrAndDebugLoc(*OtherMI); + + Builder.buildInstr(IsSigned ? TargetOpcode::G_SDIVREM + : TargetOpcode::G_UDIVREM, + {DestDivReg, DestRemReg}, + {MI.getOperand(1).getReg(), MI.getOperand(2).getReg()}); + MI.eraseFromParent(); + OtherMI->eraseFromParent(); +} + +bool CombinerHelper::matchOptBrCondByInvertingCond(MachineInstr &MI, + MachineInstr *&BrCond) { + assert(MI.getOpcode() == TargetOpcode::G_BR); + // Try to match the following: // bb1: // G_BRCOND %c1, %bb2 @@ -964,21 +1057,20 @@ bool CombinerHelper::matchOptBrCondByInvertingCond(MachineInstr &MI) { return false; assert(std::next(BrIt) == MBB->end() && "expected G_BR to be a terminator"); - MachineInstr *BrCond = &*std::prev(BrIt); + BrCond = &*std::prev(BrIt); if (BrCond->getOpcode() != TargetOpcode::G_BRCOND) return false; - // Check that the next block is the conditional branch target. - if (!MBB->isLayoutSuccessor(BrCond->getOperand(1).getMBB())) - return false; - return true; + // Check that the next block is the conditional branch target. Also make sure + // that it isn't the same as the G_BR's target (otherwise, this will loop.) + MachineBasicBlock *BrCondTarget = BrCond->getOperand(1).getMBB(); + return BrCondTarget != MI.getOperand(0).getMBB() && + MBB->isLayoutSuccessor(BrCondTarget); } -void CombinerHelper::applyOptBrCondByInvertingCond(MachineInstr &MI) { +void CombinerHelper::applyOptBrCondByInvertingCond(MachineInstr &MI, + MachineInstr *&BrCond) { MachineBasicBlock *BrTarget = MI.getOperand(0).getMBB(); - MachineBasicBlock::iterator BrIt(MI); - MachineInstr *BrCond = &*std::prev(BrIt); - Builder.setInstrAndDebugLoc(*BrCond); LLT Ty = MRI.getType(BrCond->getOperand(0).getReg()); // FIXME: Does int/fp matter for this? If so, we might need to restrict @@ -1056,7 +1148,7 @@ static bool findGISelOptimalMemOpLowering(std::vector<LLT> &MemOps, MVT VT = getMVTForLLT(Ty); if (NumMemOps && Op.allowOverlap() && NewTySize < Size && TLI.allowsMisalignedMemoryAccesses( - VT, DstAS, Op.isFixedDstAlign() ? Op.getDstAlign().value() : 0, + VT, DstAS, Op.isFixedDstAlign() ? Op.getDstAlign() : Align(1), MachineMemOperand::MONone, &Fast) && Fast) TySize = Size; @@ -1117,7 +1209,7 @@ static Register getMemsetValue(Register Val, LLT Ty, MachineIRBuilder &MIB) { } bool CombinerHelper::optimizeMemset(MachineInstr &MI, Register Dst, - Register Val, unsigned KnownLen, + Register Val, uint64_t KnownLen, Align Alignment, bool IsVolatile) { auto &MF = *MI.getParent()->getParent(); const auto &TLI = *MF.getSubtarget().getTargetLowering(); @@ -1211,7 +1303,7 @@ bool CombinerHelper::optimizeMemset(MachineInstr &MI, Register Dst, } auto *StoreMMO = - MF.getMachineMemOperand(&DstMMO, DstOff, Ty.getSizeInBytes()); + MF.getMachineMemOperand(&DstMMO, DstOff, Ty); Register Ptr = Dst; if (DstOff != 0) { @@ -1229,10 +1321,51 @@ bool CombinerHelper::optimizeMemset(MachineInstr &MI, Register Dst, return true; } +bool CombinerHelper::tryEmitMemcpyInline(MachineInstr &MI) { + assert(MI.getOpcode() == TargetOpcode::G_MEMCPY_INLINE); + + Register Dst = MI.getOperand(0).getReg(); + Register Src = MI.getOperand(1).getReg(); + Register Len = MI.getOperand(2).getReg(); + + const auto *MMOIt = MI.memoperands_begin(); + const MachineMemOperand *MemOp = *MMOIt; + bool IsVolatile = MemOp->isVolatile(); + + // See if this is a constant length copy + auto LenVRegAndVal = getConstantVRegValWithLookThrough(Len, MRI); + // FIXME: support dynamically sized G_MEMCPY_INLINE + assert(LenVRegAndVal.hasValue() && + "inline memcpy with dynamic size is not yet supported"); + uint64_t KnownLen = LenVRegAndVal->Value.getZExtValue(); + if (KnownLen == 0) { + MI.eraseFromParent(); + return true; + } + + const auto &DstMMO = **MI.memoperands_begin(); + const auto &SrcMMO = **std::next(MI.memoperands_begin()); + Align DstAlign = DstMMO.getBaseAlign(); + Align SrcAlign = SrcMMO.getBaseAlign(); + + return tryEmitMemcpyInline(MI, Dst, Src, KnownLen, DstAlign, SrcAlign, + IsVolatile); +} + +bool CombinerHelper::tryEmitMemcpyInline(MachineInstr &MI, Register Dst, + Register Src, uint64_t KnownLen, + Align DstAlign, Align SrcAlign, + bool IsVolatile) { + assert(MI.getOpcode() == TargetOpcode::G_MEMCPY_INLINE); + return optimizeMemcpy(MI, Dst, Src, KnownLen, + std::numeric_limits<uint64_t>::max(), DstAlign, + SrcAlign, IsVolatile); +} + bool CombinerHelper::optimizeMemcpy(MachineInstr &MI, Register Dst, - Register Src, unsigned KnownLen, - Align DstAlign, Align SrcAlign, - bool IsVolatile) { + Register Src, uint64_t KnownLen, + uint64_t Limit, Align DstAlign, + Align SrcAlign, bool IsVolatile) { auto &MF = *MI.getParent()->getParent(); const auto &TLI = *MF.getSubtarget().getTargetLowering(); auto &DL = MF.getDataLayout(); @@ -1242,7 +1375,6 @@ bool CombinerHelper::optimizeMemcpy(MachineInstr &MI, Register Dst, bool DstAlignCanChange = false; MachineFrameInfo &MFI = MF.getFrameInfo(); - bool OptSize = shouldLowerMemFuncForSize(MF); Align Alignment = commonAlignment(DstAlign, SrcAlign); MachineInstr *FIDef = getOpcodeDef(TargetOpcode::G_FRAME_INDEX, Dst, MRI); @@ -1253,7 +1385,6 @@ bool CombinerHelper::optimizeMemcpy(MachineInstr &MI, Register Dst, // FIXME: also use the equivalent of isMemSrcFromConstant and alwaysinlining // if the memcpy is in a tail call position. - unsigned Limit = TLI.getMaxStoresPerMemcpy(OptSize); std::vector<LLT> MemOps; const auto &DstMMO = **MI.memoperands_begin(); @@ -1277,7 +1408,7 @@ bool CombinerHelper::optimizeMemcpy(MachineInstr &MI, Register Dst, // Don't promote to an alignment that would require dynamic stack // realignment. const TargetRegisterInfo *TRI = MF.getSubtarget().getRegisterInfo(); - if (!TRI->needsStackRealignment(MF)) + if (!TRI->hasStackRealignment(MF)) while (NewAlign > Alignment && DL.exceedsNaturalStackAlignment(NewAlign)) NewAlign = NewAlign / 2; @@ -1336,7 +1467,7 @@ bool CombinerHelper::optimizeMemcpy(MachineInstr &MI, Register Dst, } bool CombinerHelper::optimizeMemmove(MachineInstr &MI, Register Dst, - Register Src, unsigned KnownLen, + Register Src, uint64_t KnownLen, Align DstAlign, Align SrcAlign, bool IsVolatile) { auto &MF = *MI.getParent()->getParent(); @@ -1382,7 +1513,7 @@ bool CombinerHelper::optimizeMemmove(MachineInstr &MI, Register Dst, // Don't promote to an alignment that would require dynamic stack // realignment. const TargetRegisterInfo *TRI = MF.getSubtarget().getRegisterInfo(); - if (!TRI->needsStackRealignment(MF)) + if (!TRI->hasStackRealignment(MF)) while (NewAlign > Alignment && DL.exceedsNaturalStackAlignment(NewAlign)) NewAlign = NewAlign / 2; @@ -1449,10 +1580,6 @@ bool CombinerHelper::tryCombineMemCpyFamily(MachineInstr &MI, unsigned MaxLen) { auto MMOIt = MI.memoperands_begin(); const MachineMemOperand *MemOp = *MMOIt; - bool IsVolatile = MemOp->isVolatile(); - // Don't try to optimize volatile. - if (IsVolatile) - return false; Align DstAlign = MemOp->getBaseAlign(); Align SrcAlign; @@ -1470,18 +1597,33 @@ bool CombinerHelper::tryCombineMemCpyFamily(MachineInstr &MI, unsigned MaxLen) { auto LenVRegAndVal = getConstantVRegValWithLookThrough(Len, MRI); if (!LenVRegAndVal) return false; // Leave it to the legalizer to lower it to a libcall. - unsigned KnownLen = LenVRegAndVal->Value.getZExtValue(); + uint64_t KnownLen = LenVRegAndVal->Value.getZExtValue(); if (KnownLen == 0) { MI.eraseFromParent(); return true; } + bool IsVolatile = MemOp->isVolatile(); + if (Opc == TargetOpcode::G_MEMCPY_INLINE) + return tryEmitMemcpyInline(MI, Dst, Src, KnownLen, DstAlign, SrcAlign, + IsVolatile); + + // Don't try to optimize volatile. + if (IsVolatile) + return false; + if (MaxLen && KnownLen > MaxLen) return false; - if (Opc == TargetOpcode::G_MEMCPY) - return optimizeMemcpy(MI, Dst, Src, KnownLen, DstAlign, SrcAlign, IsVolatile); + if (Opc == TargetOpcode::G_MEMCPY) { + auto &MF = *MI.getParent()->getParent(); + const auto &TLI = *MF.getSubtarget().getTargetLowering(); + bool OptSize = shouldLowerMemFuncForSize(MF); + uint64_t Limit = TLI.getMaxStoresPerMemcpy(OptSize); + return optimizeMemcpy(MI, Dst, Src, KnownLen, Limit, DstAlign, SrcAlign, + IsVolatile); + } if (Opc == TargetOpcode::G_MEMMOVE) return optimizeMemmove(MI, Dst, Src, KnownLen, DstAlign, SrcAlign, IsVolatile); if (Opc == TargetOpcode::G_MEMSET) @@ -1540,7 +1682,7 @@ bool CombinerHelper::matchCombineConstantFoldFpUnary(MachineInstr &MI, return Cst.hasValue(); } -bool CombinerHelper::applyCombineConstantFoldFpUnary(MachineInstr &MI, +void CombinerHelper::applyCombineConstantFoldFpUnary(MachineInstr &MI, Optional<APFloat> &Cst) { assert(Cst.hasValue() && "Optional is unexpectedly empty!"); Builder.setInstrAndDebugLoc(MI); @@ -1549,7 +1691,6 @@ bool CombinerHelper::applyCombineConstantFoldFpUnary(MachineInstr &MI, Register DstReg = MI.getOperand(0).getReg(); Builder.buildFConstant(DstReg, *FPVal); MI.eraseFromParent(); - return true; } bool CombinerHelper::matchPtrAddImmedChain(MachineInstr &MI, @@ -1569,6 +1710,13 @@ bool CombinerHelper::matchPtrAddImmedChain(MachineInstr &MI, if (!MaybeImmVal) return false; + // Don't do this combine if there multiple uses of the first PTR_ADD, + // since we may be able to compute the second PTR_ADD as an immediate + // offset anyway. Folding the first offset into the second may cause us + // to go beyond the bounds of our legal addressing modes. + if (!MRI.hasOneNonDBGUse(Add2)) + return false; + MachineInstr *Add2Def = MRI.getUniqueVRegDef(Add2); if (!Add2Def || Add2Def->getOpcode() != TargetOpcode::G_PTR_ADD) return false; @@ -1585,7 +1733,7 @@ bool CombinerHelper::matchPtrAddImmedChain(MachineInstr &MI, return true; } -bool CombinerHelper::applyPtrAddImmedChain(MachineInstr &MI, +void CombinerHelper::applyPtrAddImmedChain(MachineInstr &MI, PtrAddChain &MatchInfo) { assert(MI.getOpcode() == TargetOpcode::G_PTR_ADD && "Expected G_PTR_ADD"); MachineIRBuilder MIB(MI); @@ -1595,7 +1743,6 @@ bool CombinerHelper::applyPtrAddImmedChain(MachineInstr &MI, MI.getOperand(1).setReg(MatchInfo.Base); MI.getOperand(2).setReg(NewOffset.getReg(0)); Observer.changedInstr(MI); - return true; } bool CombinerHelper::matchShiftImmedChain(MachineInstr &MI, @@ -1643,7 +1790,7 @@ bool CombinerHelper::matchShiftImmedChain(MachineInstr &MI, return true; } -bool CombinerHelper::applyShiftImmedChain(MachineInstr &MI, +void CombinerHelper::applyShiftImmedChain(MachineInstr &MI, RegisterImmPair &MatchInfo) { unsigned Opcode = MI.getOpcode(); assert((Opcode == TargetOpcode::G_SHL || Opcode == TargetOpcode::G_ASHR || @@ -1661,7 +1808,7 @@ bool CombinerHelper::applyShiftImmedChain(MachineInstr &MI, if (Opcode == TargetOpcode::G_SHL || Opcode == TargetOpcode::G_LSHR) { Builder.buildConstant(MI.getOperand(0), 0); MI.eraseFromParent(); - return true; + return; } // Arithmetic shift and saturating signed left shift have no effect beyond // scalar size. @@ -1674,7 +1821,6 @@ bool CombinerHelper::applyShiftImmedChain(MachineInstr &MI, MI.getOperand(1).setReg(MatchInfo.Reg); MI.getOperand(2).setReg(NewImm); Observer.changedInstr(MI); - return true; } bool CombinerHelper::matchShiftOfShiftedLogic(MachineInstr &MI, @@ -1758,7 +1904,7 @@ bool CombinerHelper::matchShiftOfShiftedLogic(MachineInstr &MI, return true; } -bool CombinerHelper::applyShiftOfShiftedLogic(MachineInstr &MI, +void CombinerHelper::applyShiftOfShiftedLogic(MachineInstr &MI, ShiftOfShiftedLogic &MatchInfo) { unsigned Opcode = MI.getOpcode(); assert((Opcode == TargetOpcode::G_SHL || Opcode == TargetOpcode::G_ASHR || @@ -1790,7 +1936,6 @@ bool CombinerHelper::applyShiftOfShiftedLogic(MachineInstr &MI, MatchInfo.Logic->eraseFromParent(); MI.eraseFromParent(); - return true; } bool CombinerHelper::matchCombineMulToShl(MachineInstr &MI, @@ -1805,7 +1950,7 @@ bool CombinerHelper::matchCombineMulToShl(MachineInstr &MI, return (static_cast<int32_t>(ShiftVal) != -1); } -bool CombinerHelper::applyCombineMulToShl(MachineInstr &MI, +void CombinerHelper::applyCombineMulToShl(MachineInstr &MI, unsigned &ShiftVal) { assert(MI.getOpcode() == TargetOpcode::G_MUL && "Expected a G_MUL"); MachineIRBuilder MIB(MI); @@ -1815,7 +1960,6 @@ bool CombinerHelper::applyCombineMulToShl(MachineInstr &MI, MI.setDesc(MIB.getTII().get(TargetOpcode::G_SHL)); MI.getOperand(2).setReg(ShiftCst.getReg(0)); Observer.changedInstr(MI); - return true; } // shl ([sza]ext x), y => zext (shl x, y), if shift does not overflow source @@ -1856,7 +2000,7 @@ bool CombinerHelper::matchCombineShlOfExtend(MachineInstr &MI, return MinLeadingZeros >= ShiftAmt; } -bool CombinerHelper::applyCombineShlOfExtend(MachineInstr &MI, +void CombinerHelper::applyCombineShlOfExtend(MachineInstr &MI, const RegisterImmPair &MatchData) { Register ExtSrcReg = MatchData.Reg; int64_t ShiftAmtVal = MatchData.Imm; @@ -1868,6 +2012,24 @@ bool CombinerHelper::applyCombineShlOfExtend(MachineInstr &MI, Builder.buildShl(ExtSrcTy, ExtSrcReg, ShiftAmt, MI.getFlags()); Builder.buildZExt(MI.getOperand(0), NarrowShift); MI.eraseFromParent(); +} + +bool CombinerHelper::matchCombineMergeUnmerge(MachineInstr &MI, + Register &MatchInfo) { + GMerge &Merge = cast<GMerge>(MI); + SmallVector<Register, 16> MergedValues; + for (unsigned I = 0; I < Merge.getNumSources(); ++I) + MergedValues.emplace_back(Merge.getSourceReg(I)); + + auto *Unmerge = getOpcodeDef<GUnmerge>(MergedValues[0], MRI); + if (!Unmerge || Unmerge->getNumDefs() != Merge.getNumSources()) + return false; + + for (unsigned I = 0; I < MergedValues.size(); ++I) + if (MergedValues[I] != Unmerge->getReg(I)) + return false; + + MatchInfo = Unmerge->getSourceReg(); return true; } @@ -1906,7 +2068,7 @@ bool CombinerHelper::matchCombineUnmergeMergeToPlainValues( return true; } -bool CombinerHelper::applyCombineUnmergeMergeToPlainValues( +void CombinerHelper::applyCombineUnmergeMergeToPlainValues( MachineInstr &MI, SmallVectorImpl<Register> &Operands) { assert(MI.getOpcode() == TargetOpcode::G_UNMERGE_VALUES && "Expected an unmerge"); @@ -1927,7 +2089,6 @@ bool CombinerHelper::applyCombineUnmergeMergeToPlainValues( Builder.buildCast(DstReg, SrcReg); } MI.eraseFromParent(); - return true; } bool CombinerHelper::matchCombineUnmergeConstant(MachineInstr &MI, @@ -1955,7 +2116,7 @@ bool CombinerHelper::matchCombineUnmergeConstant(MachineInstr &MI, return true; } -bool CombinerHelper::applyCombineUnmergeConstant(MachineInstr &MI, +void CombinerHelper::applyCombineUnmergeConstant(MachineInstr &MI, SmallVectorImpl<APInt> &Csts) { assert(MI.getOpcode() == TargetOpcode::G_UNMERGE_VALUES && "Expected an unmerge"); @@ -1969,7 +2130,6 @@ bool CombinerHelper::applyCombineUnmergeConstant(MachineInstr &MI, } MI.eraseFromParent(); - return true; } bool CombinerHelper::matchCombineUnmergeWithDeadLanesToTrunc(MachineInstr &MI) { @@ -1983,7 +2143,7 @@ bool CombinerHelper::matchCombineUnmergeWithDeadLanesToTrunc(MachineInstr &MI) { return true; } -bool CombinerHelper::applyCombineUnmergeWithDeadLanesToTrunc(MachineInstr &MI) { +void CombinerHelper::applyCombineUnmergeWithDeadLanesToTrunc(MachineInstr &MI) { Builder.setInstrAndDebugLoc(MI); Register SrcReg = MI.getOperand(MI.getNumDefs()).getReg(); // Truncating a vector is going to truncate every single lane, @@ -2002,7 +2162,6 @@ bool CombinerHelper::applyCombineUnmergeWithDeadLanesToTrunc(MachineInstr &MI) { } else Builder.buildTrunc(Dst0Reg, SrcReg); MI.eraseFromParent(); - return true; } bool CombinerHelper::matchCombineUnmergeZExtToZExt(MachineInstr &MI) { @@ -2031,7 +2190,7 @@ bool CombinerHelper::matchCombineUnmergeZExtToZExt(MachineInstr &MI) { return ZExtSrcTy.getSizeInBits() <= Dst0Ty.getSizeInBits(); } -bool CombinerHelper::applyCombineUnmergeZExtToZExt(MachineInstr &MI) { +void CombinerHelper::applyCombineUnmergeZExtToZExt(MachineInstr &MI) { assert(MI.getOpcode() == TargetOpcode::G_UNMERGE_VALUES && "Expected an unmerge"); @@ -2063,7 +2222,6 @@ bool CombinerHelper::applyCombineUnmergeZExtToZExt(MachineInstr &MI) { replaceRegWith(MRI, MI.getOperand(Idx).getReg(), ZeroReg); } MI.eraseFromParent(); - return true; } bool CombinerHelper::matchCombineShiftToUnmerge(MachineInstr &MI, @@ -2091,7 +2249,7 @@ bool CombinerHelper::matchCombineShiftToUnmerge(MachineInstr &MI, return ShiftVal >= Size / 2 && ShiftVal < Size; } -bool CombinerHelper::applyCombineShiftToUnmerge(MachineInstr &MI, +void CombinerHelper::applyCombineShiftToUnmerge(MachineInstr &MI, const unsigned &ShiftVal) { Register DstReg = MI.getOperand(0).getReg(); Register SrcReg = MI.getOperand(1).getReg(); @@ -2162,7 +2320,6 @@ bool CombinerHelper::applyCombineShiftToUnmerge(MachineInstr &MI, } MI.eraseFromParent(); - return true; } bool CombinerHelper::tryCombineShiftToUnmerge(MachineInstr &MI, @@ -2185,13 +2342,12 @@ bool CombinerHelper::matchCombineI2PToP2I(MachineInstr &MI, Register &Reg) { m_GPtrToInt(m_all_of(m_SpecificType(DstTy), m_Reg(Reg)))); } -bool CombinerHelper::applyCombineI2PToP2I(MachineInstr &MI, Register &Reg) { +void CombinerHelper::applyCombineI2PToP2I(MachineInstr &MI, Register &Reg) { assert(MI.getOpcode() == TargetOpcode::G_INTTOPTR && "Expected a G_INTTOPTR"); Register DstReg = MI.getOperand(0).getReg(); Builder.setInstr(MI); Builder.buildCopy(DstReg, Reg); MI.eraseFromParent(); - return true; } bool CombinerHelper::matchCombineP2IToI2P(MachineInstr &MI, Register &Reg) { @@ -2200,13 +2356,12 @@ bool CombinerHelper::matchCombineP2IToI2P(MachineInstr &MI, Register &Reg) { return mi_match(SrcReg, MRI, m_GIntToPtr(m_Reg(Reg))); } -bool CombinerHelper::applyCombineP2IToI2P(MachineInstr &MI, Register &Reg) { +void CombinerHelper::applyCombineP2IToI2P(MachineInstr &MI, Register &Reg) { assert(MI.getOpcode() == TargetOpcode::G_PTRTOINT && "Expected a G_PTRTOINT"); Register DstReg = MI.getOperand(0).getReg(); Builder.setInstr(MI); Builder.buildZExtOrTrunc(DstReg, Reg); MI.eraseFromParent(); - return true; } bool CombinerHelper::matchCombineAddP2IToPtrAdd( @@ -2234,7 +2389,7 @@ bool CombinerHelper::matchCombineAddP2IToPtrAdd( return false; } -bool CombinerHelper::applyCombineAddP2IToPtrAdd( +void CombinerHelper::applyCombineAddP2IToPtrAdd( MachineInstr &MI, std::pair<Register, bool> &PtrReg) { Register Dst = MI.getOperand(0).getReg(); Register LHS = MI.getOperand(1).getReg(); @@ -2251,7 +2406,6 @@ bool CombinerHelper::applyCombineAddP2IToPtrAdd( auto PtrAdd = Builder.buildPtrAdd(PtrTy, LHS, RHS); Builder.buildPtrToInt(Dst, PtrAdd); MI.eraseFromParent(); - return true; } bool CombinerHelper::matchCombineConstPtrAddToI2P(MachineInstr &MI, @@ -2272,7 +2426,7 @@ bool CombinerHelper::matchCombineConstPtrAddToI2P(MachineInstr &MI, return false; } -bool CombinerHelper::applyCombineConstPtrAddToI2P(MachineInstr &MI, +void CombinerHelper::applyCombineConstPtrAddToI2P(MachineInstr &MI, int64_t &NewCst) { assert(MI.getOpcode() == TargetOpcode::G_PTR_ADD && "Expected a G_PTR_ADD"); Register Dst = MI.getOperand(0).getReg(); @@ -2280,7 +2434,6 @@ bool CombinerHelper::applyCombineConstPtrAddToI2P(MachineInstr &MI, Builder.setInstrAndDebugLoc(MI); Builder.buildConstant(Dst, NewCst); MI.eraseFromParent(); - return true; } bool CombinerHelper::matchCombineAnyExtTrunc(MachineInstr &MI, Register &Reg) { @@ -2292,12 +2445,18 @@ bool CombinerHelper::matchCombineAnyExtTrunc(MachineInstr &MI, Register &Reg) { m_GTrunc(m_all_of(m_Reg(Reg), m_SpecificType(DstTy)))); } -bool CombinerHelper::applyCombineAnyExtTrunc(MachineInstr &MI, Register &Reg) { - assert(MI.getOpcode() == TargetOpcode::G_ANYEXT && "Expected a G_ANYEXT"); +bool CombinerHelper::matchCombineZextTrunc(MachineInstr &MI, Register &Reg) { + assert(MI.getOpcode() == TargetOpcode::G_ZEXT && "Expected a G_ZEXT"); Register DstReg = MI.getOperand(0).getReg(); - MI.eraseFromParent(); - replaceRegWith(MRI, DstReg, Reg); - return true; + Register SrcReg = MI.getOperand(1).getReg(); + LLT DstTy = MRI.getType(DstReg); + if (mi_match(SrcReg, MRI, + m_GTrunc(m_all_of(m_Reg(Reg), m_SpecificType(DstTy))))) { + unsigned DstSize = DstTy.getScalarSizeInBits(); + unsigned SrcSize = MRI.getType(SrcReg).getScalarSizeInBits(); + return KB->getKnownBits(Reg).countMinLeadingZeros() >= DstSize - SrcSize; + } + return false; } bool CombinerHelper::matchCombineExtOfExt( @@ -2321,7 +2480,7 @@ bool CombinerHelper::matchCombineExtOfExt( return false; } -bool CombinerHelper::applyCombineExtOfExt( +void CombinerHelper::applyCombineExtOfExt( MachineInstr &MI, std::tuple<Register, unsigned> &MatchInfo) { assert((MI.getOpcode() == TargetOpcode::G_ANYEXT || MI.getOpcode() == TargetOpcode::G_SEXT || @@ -2336,7 +2495,7 @@ bool CombinerHelper::applyCombineExtOfExt( Observer.changingInstr(MI); MI.getOperand(1).setReg(Reg); Observer.changedInstr(MI); - return true; + return; } // Combine: @@ -2349,13 +2508,10 @@ bool CombinerHelper::applyCombineExtOfExt( Builder.setInstrAndDebugLoc(MI); Builder.buildInstr(SrcExtOp, {DstReg}, {Reg}); MI.eraseFromParent(); - return true; } - - return false; } -bool CombinerHelper::applyCombineMulByNegativeOne(MachineInstr &MI) { +void CombinerHelper::applyCombineMulByNegativeOne(MachineInstr &MI) { assert(MI.getOpcode() == TargetOpcode::G_MUL && "Expected a G_MUL"); Register DstReg = MI.getOperand(0).getReg(); Register SrcReg = MI.getOperand(1).getReg(); @@ -2365,7 +2521,6 @@ bool CombinerHelper::applyCombineMulByNegativeOne(MachineInstr &MI) { Builder.buildSub(DstReg, Builder.buildConstant(DstTy, 0), SrcReg, MI.getFlags()); MI.eraseFromParent(); - return true; } bool CombinerHelper::matchCombineFNegOfFNeg(MachineInstr &MI, Register &Reg) { @@ -2381,14 +2536,6 @@ bool CombinerHelper::matchCombineFAbsOfFAbs(MachineInstr &MI, Register &Src) { return mi_match(Src, MRI, m_GFabs(m_Reg(AbsSrc))); } -bool CombinerHelper::applyCombineFAbsOfFAbs(MachineInstr &MI, Register &Src) { - assert(MI.getOpcode() == TargetOpcode::G_FABS && "Expected a G_FABS"); - Register Dst = MI.getOperand(0).getReg(); - MI.eraseFromParent(); - replaceRegWith(MRI, Dst, Src); - return true; -} - bool CombinerHelper::matchCombineTruncOfExt( MachineInstr &MI, std::pair<Register, unsigned> &MatchInfo) { assert(MI.getOpcode() == TargetOpcode::G_TRUNC && "Expected a G_TRUNC"); @@ -2403,7 +2550,7 @@ bool CombinerHelper::matchCombineTruncOfExt( return false; } -bool CombinerHelper::applyCombineTruncOfExt( +void CombinerHelper::applyCombineTruncOfExt( MachineInstr &MI, std::pair<Register, unsigned> &MatchInfo) { assert(MI.getOpcode() == TargetOpcode::G_TRUNC && "Expected a G_TRUNC"); Register SrcReg = MatchInfo.first; @@ -2414,7 +2561,7 @@ bool CombinerHelper::applyCombineTruncOfExt( if (SrcTy == DstTy) { MI.eraseFromParent(); replaceRegWith(MRI, DstReg, SrcReg); - return true; + return; } Builder.setInstrAndDebugLoc(MI); if (SrcTy.getSizeInBits() < DstTy.getSizeInBits()) @@ -2422,7 +2569,6 @@ bool CombinerHelper::applyCombineTruncOfExt( else Builder.buildTrunc(DstReg, SrcReg); MI.eraseFromParent(); - return true; } bool CombinerHelper::matchCombineTruncOfShl( @@ -2449,7 +2595,7 @@ bool CombinerHelper::matchCombineTruncOfShl( return false; } -bool CombinerHelper::applyCombineTruncOfShl( +void CombinerHelper::applyCombineTruncOfShl( MachineInstr &MI, std::pair<Register, Register> &MatchInfo) { assert(MI.getOpcode() == TargetOpcode::G_TRUNC && "Expected a G_TRUNC"); Register DstReg = MI.getOperand(0).getReg(); @@ -2463,7 +2609,6 @@ bool CombinerHelper::applyCombineTruncOfShl( auto TruncShiftSrc = Builder.buildTrunc(DstTy, ShiftSrc); Builder.buildShl(DstReg, TruncShiftSrc, ShiftAmt, SrcMI->getFlags()); MI.eraseFromParent(); - return true; } bool CombinerHelper::matchAnyExplicitUseIsUndef(MachineInstr &MI) { @@ -2662,6 +2807,14 @@ bool CombinerHelper::replaceInstWithConstant(MachineInstr &MI, int64_t C) { return true; } +bool CombinerHelper::replaceInstWithConstant(MachineInstr &MI, APInt C) { + assert(MI.getNumDefs() == 1 && "Expected only one def?"); + Builder.setInstr(MI); + Builder.buildConstant(MI.getOperand(0), C); + MI.eraseFromParent(); + return true; +} + bool CombinerHelper::replaceInstWithUndef(MachineInstr &MI) { assert(MI.getNumDefs() == 1 && "Expected only one def?"); Builder.setInstr(MI); @@ -2731,7 +2884,7 @@ bool CombinerHelper::matchCombineInsertVecElts( return TmpInst->getOpcode() == TargetOpcode::G_IMPLICIT_DEF; } -bool CombinerHelper::applyCombineInsertVecElts( +void CombinerHelper::applyCombineInsertVecElts( MachineInstr &MI, SmallVectorImpl<Register> &MatchInfo) { Builder.setInstr(MI); Register UndefReg; @@ -2748,17 +2901,15 @@ bool CombinerHelper::applyCombineInsertVecElts( } Builder.buildBuildVector(MI.getOperand(0).getReg(), MatchInfo); MI.eraseFromParent(); - return true; } -bool CombinerHelper::applySimplifyAddToSub( +void CombinerHelper::applySimplifyAddToSub( MachineInstr &MI, std::tuple<Register, Register> &MatchInfo) { Builder.setInstr(MI); Register SubLHS, SubRHS; std::tie(SubLHS, SubRHS) = MatchInfo; Builder.buildSub(MI.getOperand(0).getReg(), SubLHS, SubRHS); MI.eraseFromParent(); - return true; } bool CombinerHelper::matchHoistLogicOpWithSameOpcodeHands( @@ -2852,7 +3003,7 @@ bool CombinerHelper::matchHoistLogicOpWithSameOpcodeHands( return true; } -bool CombinerHelper::applyBuildInstructionSteps( +void CombinerHelper::applyBuildInstructionSteps( MachineInstr &MI, InstructionStepsMatchInfo &MatchInfo) { assert(MatchInfo.InstrsToBuild.size() && "Expected at least one instr to build?"); @@ -2865,7 +3016,6 @@ bool CombinerHelper::applyBuildInstructionSteps( OperandFn(Instr); } MI.eraseFromParent(); - return true; } bool CombinerHelper::matchAshrShlToSextInreg( @@ -2885,7 +3035,8 @@ bool CombinerHelper::matchAshrShlToSextInreg( MatchInfo = std::make_tuple(Src, ShlCst); return true; } -bool CombinerHelper::applyAshShlToSextInreg( + +void CombinerHelper::applyAshShlToSextInreg( MachineInstr &MI, std::tuple<Register, int64_t> &MatchInfo) { assert(MI.getOpcode() == TargetOpcode::G_ASHR); Register Src; @@ -2895,6 +3046,32 @@ bool CombinerHelper::applyAshShlToSextInreg( Builder.setInstrAndDebugLoc(MI); Builder.buildSExtInReg(MI.getOperand(0).getReg(), Src, Size - ShiftAmt); MI.eraseFromParent(); +} + +/// and(and(x, C1), C2) -> C1&C2 ? and(x, C1&C2) : 0 +bool CombinerHelper::matchOverlappingAnd( + MachineInstr &MI, std::function<void(MachineIRBuilder &)> &MatchInfo) { + assert(MI.getOpcode() == TargetOpcode::G_AND); + + Register Dst = MI.getOperand(0).getReg(); + LLT Ty = MRI.getType(Dst); + + Register R; + int64_t C1; + int64_t C2; + if (!mi_match( + Dst, MRI, + m_GAnd(m_GAnd(m_Reg(R), m_ICst(C1)), m_ICst(C2)))) + return false; + + MatchInfo = [=](MachineIRBuilder &B) { + if (C1 & C2) { + B.buildAnd(Dst, R, B.buildConstant(Ty, C1 & C2)); + return; + } + auto Zero = B.buildConstant(Ty, 0); + replaceRegWith(MRI, Dst, Zero->getOperand(0).getReg()); + }; return true; } @@ -3091,7 +3268,7 @@ bool CombinerHelper::matchNotCmp(MachineInstr &MI, return true; } -bool CombinerHelper::applyNotCmp(MachineInstr &MI, +void CombinerHelper::applyNotCmp(MachineInstr &MI, SmallVectorImpl<Register> &RegsToNegate) { for (Register Reg : RegsToNegate) { MachineInstr *Def = MRI.getVRegDef(Reg); @@ -3121,7 +3298,6 @@ bool CombinerHelper::applyNotCmp(MachineInstr &MI, replaceRegWith(MRI, MI.getOperand(0).getReg(), MI.getOperand(1).getReg()); MI.eraseFromParent(); - return true; } bool CombinerHelper::matchXorOfAndWithSameReg( @@ -3155,7 +3331,7 @@ bool CombinerHelper::matchXorOfAndWithSameReg( return Y == SharedReg; } -bool CombinerHelper::applyXorOfAndWithSameReg( +void CombinerHelper::applyXorOfAndWithSameReg( MachineInstr &MI, std::pair<Register, Register> &MatchInfo) { // Fold (xor (and x, y), y) -> (and (not x), y) Builder.setInstrAndDebugLoc(MI); @@ -3167,7 +3343,6 @@ bool CombinerHelper::applyXorOfAndWithSameReg( MI.getOperand(1).setReg(Not->getOperand(0).getReg()); MI.getOperand(2).setReg(Y); Observer.changedInstr(MI); - return true; } bool CombinerHelper::matchPtrAddZero(MachineInstr &MI) { @@ -3188,16 +3363,15 @@ bool CombinerHelper::matchPtrAddZero(MachineInstr &MI) { return isBuildVectorAllZeros(*VecMI, MRI); } -bool CombinerHelper::applyPtrAddZero(MachineInstr &MI) { +void CombinerHelper::applyPtrAddZero(MachineInstr &MI) { assert(MI.getOpcode() == TargetOpcode::G_PTR_ADD); Builder.setInstrAndDebugLoc(MI); Builder.buildIntToPtr(MI.getOperand(0), MI.getOperand(2)); MI.eraseFromParent(); - return true; } /// The second source operand is known to be a power of 2. -bool CombinerHelper::applySimplifyURemByPow2(MachineInstr &MI) { +void CombinerHelper::applySimplifyURemByPow2(MachineInstr &MI) { Register DstReg = MI.getOperand(0).getReg(); Register Src0 = MI.getOperand(1).getReg(); Register Pow2Src1 = MI.getOperand(2).getReg(); @@ -3209,7 +3383,6 @@ bool CombinerHelper::applySimplifyURemByPow2(MachineInstr &MI) { auto Add = Builder.buildAdd(Ty, Pow2Src1, NegOne); Builder.buildAnd(DstReg, Src0, Add); MI.eraseFromParent(); - return true; } Optional<SmallVector<Register, 8>> @@ -3283,7 +3456,7 @@ CombinerHelper::findCandidatesForLoadOrCombine(const MachineInstr *Root) const { /// e.g. x[i] << 24 /// /// \returns The load instruction and the byte offset it is moved into. -static Optional<std::pair<MachineInstr *, int64_t>> +static Optional<std::pair<GZExtLoad *, int64_t>> matchLoadAndBytePosition(Register Reg, unsigned MemSizeInBits, const MachineRegisterInfo &MRI) { assert(MRI.hasOneNonDBGUse(Reg) && @@ -3300,18 +3473,17 @@ matchLoadAndBytePosition(Register Reg, unsigned MemSizeInBits, return None; // TODO: Handle other types of loads. - auto *Load = getOpcodeDef(TargetOpcode::G_ZEXTLOAD, MaybeLoad, MRI); + auto *Load = getOpcodeDef<GZExtLoad>(MaybeLoad, MRI); if (!Load) return None; - const auto &MMO = **Load->memoperands_begin(); - if (!MMO.isUnordered() || MMO.getSizeInBits() != MemSizeInBits) + if (!Load->isUnordered() || Load->getMemSizeInBits() != MemSizeInBits) return None; return std::make_pair(Load, Shift / MemSizeInBits); } -Optional<std::pair<MachineInstr *, int64_t>> +Optional<std::tuple<GZExtLoad *, int64_t, GZExtLoad *>> CombinerHelper::findLoadOffsetsForLoadOrCombine( SmallDenseMap<int64_t, int64_t, 8> &MemOffset2Idx, const SmallVector<Register, 8> &RegsToVisit, const unsigned MemSizeInBits) { @@ -3323,7 +3495,7 @@ CombinerHelper::findLoadOffsetsForLoadOrCombine( int64_t LowestIdx = INT64_MAX; // The load which uses the lowest index. - MachineInstr *LowestIdxLoad = nullptr; + GZExtLoad *LowestIdxLoad = nullptr; // Keeps track of the load indices we see. We shouldn't see any indices twice. SmallSet<int64_t, 8> SeenIdx; @@ -3334,10 +3506,10 @@ CombinerHelper::findLoadOffsetsForLoadOrCombine( const MachineMemOperand *MMO = nullptr; // Earliest instruction-order load in the pattern. - MachineInstr *EarliestLoad = nullptr; + GZExtLoad *EarliestLoad = nullptr; // Latest instruction-order load in the pattern. - MachineInstr *LatestLoad = nullptr; + GZExtLoad *LatestLoad = nullptr; // Base pointer which every load should share. Register BasePtr; @@ -3352,7 +3524,7 @@ CombinerHelper::findLoadOffsetsForLoadOrCombine( auto LoadAndPos = matchLoadAndBytePosition(Reg, MemSizeInBits, MRI); if (!LoadAndPos) return None; - MachineInstr *Load; + GZExtLoad *Load; int64_t DstPos; std::tie(Load, DstPos) = *LoadAndPos; @@ -3365,10 +3537,10 @@ CombinerHelper::findLoadOffsetsForLoadOrCombine( return None; // Make sure that the MachineMemOperands of every seen load are compatible. - const MachineMemOperand *LoadMMO = *Load->memoperands_begin(); + auto &LoadMMO = Load->getMMO(); if (!MMO) - MMO = LoadMMO; - if (MMO->getAddrSpace() != LoadMMO->getAddrSpace()) + MMO = &LoadMMO; + if (MMO->getAddrSpace() != LoadMMO.getAddrSpace()) return None; // Find out what the base pointer and index for the load is. @@ -3442,7 +3614,7 @@ CombinerHelper::findLoadOffsetsForLoadOrCombine( return None; } - return std::make_pair(LowestIdxLoad, LowestIdx); + return std::make_tuple(LowestIdxLoad, LowestIdx, LatestLoad); } bool CombinerHelper::matchLoadOrCombine( @@ -3490,13 +3662,13 @@ bool CombinerHelper::matchLoadOrCombine( // Also verify that each of these ends up putting a[i] into the same memory // offset as a load into a wide type would. SmallDenseMap<int64_t, int64_t, 8> MemOffset2Idx; - MachineInstr *LowestIdxLoad; + GZExtLoad *LowestIdxLoad, *LatestLoad; int64_t LowestIdx; auto MaybeLoadInfo = findLoadOffsetsForLoadOrCombine( MemOffset2Idx, *RegsToVisit, NarrowMemSizeInBits); if (!MaybeLoadInfo) return false; - std::tie(LowestIdxLoad, LowestIdx) = *MaybeLoadInfo; + std::tie(LowestIdxLoad, LowestIdx, LatestLoad) = *MaybeLoadInfo; // We have a bunch of loads being OR'd together. Using the addresses + offsets // we found before, check if this corresponds to a big or little endian byte @@ -3530,12 +3702,12 @@ bool CombinerHelper::matchLoadOrCombine( // We wil reuse the pointer from the load which ends up at byte offset 0. It // may not use index 0. - Register Ptr = LowestIdxLoad->getOperand(1).getReg(); - const MachineMemOperand &MMO = **LowestIdxLoad->memoperands_begin(); + Register Ptr = LowestIdxLoad->getPointerReg(); + const MachineMemOperand &MMO = LowestIdxLoad->getMMO(); LegalityQuery::MemDesc MMDesc; - MMDesc.SizeInBits = WideMemSizeInBits; + MMDesc.MemoryTy = Ty; MMDesc.AlignInBits = MMO.getAlign().value() * 8; - MMDesc.Ordering = MMO.getOrdering(); + MMDesc.Ordering = MMO.getSuccessOrdering(); if (!isLegalOrBeforeLegalizer( {TargetOpcode::G_LOAD, {Ty, MRI.getType(Ptr)}, {MMDesc}})) return false; @@ -3551,6 +3723,7 @@ bool CombinerHelper::matchLoadOrCombine( return false; MatchInfo = [=](MachineIRBuilder &MIB) { + MIB.setInstrAndDebugLoc(*LatestLoad); Register LoadDst = NeedsBSwap ? MRI.cloneVirtualRegister(Dst) : Dst; MIB.buildLoad(LoadDst, Ptr, *NewMMO); if (NeedsBSwap) @@ -3559,11 +3732,535 @@ bool CombinerHelper::matchLoadOrCombine( return true; } -bool CombinerHelper::applyLoadOrCombine( +bool CombinerHelper::matchExtendThroughPhis(MachineInstr &MI, + MachineInstr *&ExtMI) { + assert(MI.getOpcode() == TargetOpcode::G_PHI); + + Register DstReg = MI.getOperand(0).getReg(); + + // TODO: Extending a vector may be expensive, don't do this until heuristics + // are better. + if (MRI.getType(DstReg).isVector()) + return false; + + // Try to match a phi, whose only use is an extend. + if (!MRI.hasOneNonDBGUse(DstReg)) + return false; + ExtMI = &*MRI.use_instr_nodbg_begin(DstReg); + switch (ExtMI->getOpcode()) { + case TargetOpcode::G_ANYEXT: + return true; // G_ANYEXT is usually free. + case TargetOpcode::G_ZEXT: + case TargetOpcode::G_SEXT: + break; + default: + return false; + } + + // If the target is likely to fold this extend away, don't propagate. + if (Builder.getTII().isExtendLikelyToBeFolded(*ExtMI, MRI)) + return false; + + // We don't want to propagate the extends unless there's a good chance that + // they'll be optimized in some way. + // Collect the unique incoming values. + SmallPtrSet<MachineInstr *, 4> InSrcs; + for (unsigned Idx = 1; Idx < MI.getNumOperands(); Idx += 2) { + auto *DefMI = getDefIgnoringCopies(MI.getOperand(Idx).getReg(), MRI); + switch (DefMI->getOpcode()) { + case TargetOpcode::G_LOAD: + case TargetOpcode::G_TRUNC: + case TargetOpcode::G_SEXT: + case TargetOpcode::G_ZEXT: + case TargetOpcode::G_ANYEXT: + case TargetOpcode::G_CONSTANT: + InSrcs.insert(getDefIgnoringCopies(MI.getOperand(Idx).getReg(), MRI)); + // Don't try to propagate if there are too many places to create new + // extends, chances are it'll increase code size. + if (InSrcs.size() > 2) + return false; + break; + default: + return false; + } + } + return true; +} + +void CombinerHelper::applyExtendThroughPhis(MachineInstr &MI, + MachineInstr *&ExtMI) { + assert(MI.getOpcode() == TargetOpcode::G_PHI); + Register DstReg = ExtMI->getOperand(0).getReg(); + LLT ExtTy = MRI.getType(DstReg); + + // Propagate the extension into the block of each incoming reg's block. + // Use a SetVector here because PHIs can have duplicate edges, and we want + // deterministic iteration order. + SmallSetVector<MachineInstr *, 8> SrcMIs; + SmallDenseMap<MachineInstr *, MachineInstr *, 8> OldToNewSrcMap; + for (unsigned SrcIdx = 1; SrcIdx < MI.getNumOperands(); SrcIdx += 2) { + auto *SrcMI = MRI.getVRegDef(MI.getOperand(SrcIdx).getReg()); + if (!SrcMIs.insert(SrcMI)) + continue; + + // Build an extend after each src inst. + auto *MBB = SrcMI->getParent(); + MachineBasicBlock::iterator InsertPt = ++SrcMI->getIterator(); + if (InsertPt != MBB->end() && InsertPt->isPHI()) + InsertPt = MBB->getFirstNonPHI(); + + Builder.setInsertPt(*SrcMI->getParent(), InsertPt); + Builder.setDebugLoc(MI.getDebugLoc()); + auto NewExt = Builder.buildExtOrTrunc(ExtMI->getOpcode(), ExtTy, + SrcMI->getOperand(0).getReg()); + OldToNewSrcMap[SrcMI] = NewExt; + } + + // Create a new phi with the extended inputs. + Builder.setInstrAndDebugLoc(MI); + auto NewPhi = Builder.buildInstrNoInsert(TargetOpcode::G_PHI); + NewPhi.addDef(DstReg); + for (unsigned SrcIdx = 1; SrcIdx < MI.getNumOperands(); ++SrcIdx) { + auto &MO = MI.getOperand(SrcIdx); + if (!MO.isReg()) { + NewPhi.addMBB(MO.getMBB()); + continue; + } + auto *NewSrc = OldToNewSrcMap[MRI.getVRegDef(MO.getReg())]; + NewPhi.addUse(NewSrc->getOperand(0).getReg()); + } + Builder.insertInstr(NewPhi); + ExtMI->eraseFromParent(); +} + +bool CombinerHelper::matchExtractVecEltBuildVec(MachineInstr &MI, + Register &Reg) { + assert(MI.getOpcode() == TargetOpcode::G_EXTRACT_VECTOR_ELT); + // If we have a constant index, look for a G_BUILD_VECTOR source + // and find the source register that the index maps to. + Register SrcVec = MI.getOperand(1).getReg(); + LLT SrcTy = MRI.getType(SrcVec); + if (!isLegalOrBeforeLegalizer( + {TargetOpcode::G_BUILD_VECTOR, {SrcTy, SrcTy.getElementType()}})) + return false; + + auto Cst = getConstantVRegValWithLookThrough(MI.getOperand(2).getReg(), MRI); + if (!Cst || Cst->Value.getZExtValue() >= SrcTy.getNumElements()) + return false; + + unsigned VecIdx = Cst->Value.getZExtValue(); + MachineInstr *BuildVecMI = + getOpcodeDef(TargetOpcode::G_BUILD_VECTOR, SrcVec, MRI); + if (!BuildVecMI) { + BuildVecMI = getOpcodeDef(TargetOpcode::G_BUILD_VECTOR_TRUNC, SrcVec, MRI); + if (!BuildVecMI) + return false; + LLT ScalarTy = MRI.getType(BuildVecMI->getOperand(1).getReg()); + if (!isLegalOrBeforeLegalizer( + {TargetOpcode::G_BUILD_VECTOR_TRUNC, {SrcTy, ScalarTy}})) + return false; + } + + EVT Ty(getMVTForLLT(SrcTy)); + if (!MRI.hasOneNonDBGUse(SrcVec) && + !getTargetLowering().aggressivelyPreferBuildVectorSources(Ty)) + return false; + + Reg = BuildVecMI->getOperand(VecIdx + 1).getReg(); + return true; +} + +void CombinerHelper::applyExtractVecEltBuildVec(MachineInstr &MI, + Register &Reg) { + // Check the type of the register, since it may have come from a + // G_BUILD_VECTOR_TRUNC. + LLT ScalarTy = MRI.getType(Reg); + Register DstReg = MI.getOperand(0).getReg(); + LLT DstTy = MRI.getType(DstReg); + + Builder.setInstrAndDebugLoc(MI); + if (ScalarTy != DstTy) { + assert(ScalarTy.getSizeInBits() > DstTy.getSizeInBits()); + Builder.buildTrunc(DstReg, Reg); + MI.eraseFromParent(); + return; + } + replaceSingleDefInstWithReg(MI, Reg); +} + +bool CombinerHelper::matchExtractAllEltsFromBuildVector( + MachineInstr &MI, + SmallVectorImpl<std::pair<Register, MachineInstr *>> &SrcDstPairs) { + assert(MI.getOpcode() == TargetOpcode::G_BUILD_VECTOR); + // This combine tries to find build_vector's which have every source element + // extracted using G_EXTRACT_VECTOR_ELT. This can happen when transforms like + // the masked load scalarization is run late in the pipeline. There's already + // a combine for a similar pattern starting from the extract, but that + // doesn't attempt to do it if there are multiple uses of the build_vector, + // which in this case is true. Starting the combine from the build_vector + // feels more natural than trying to find sibling nodes of extracts. + // E.g. + // %vec(<4 x s32>) = G_BUILD_VECTOR %s1(s32), %s2, %s3, %s4 + // %ext1 = G_EXTRACT_VECTOR_ELT %vec, 0 + // %ext2 = G_EXTRACT_VECTOR_ELT %vec, 1 + // %ext3 = G_EXTRACT_VECTOR_ELT %vec, 2 + // %ext4 = G_EXTRACT_VECTOR_ELT %vec, 3 + // ==> + // replace ext{1,2,3,4} with %s{1,2,3,4} + + Register DstReg = MI.getOperand(0).getReg(); + LLT DstTy = MRI.getType(DstReg); + unsigned NumElts = DstTy.getNumElements(); + + SmallBitVector ExtractedElts(NumElts); + for (auto &II : make_range(MRI.use_instr_nodbg_begin(DstReg), + MRI.use_instr_nodbg_end())) { + if (II.getOpcode() != TargetOpcode::G_EXTRACT_VECTOR_ELT) + return false; + auto Cst = getConstantVRegVal(II.getOperand(2).getReg(), MRI); + if (!Cst) + return false; + unsigned Idx = Cst.getValue().getZExtValue(); + if (Idx >= NumElts) + return false; // Out of range. + ExtractedElts.set(Idx); + SrcDstPairs.emplace_back( + std::make_pair(MI.getOperand(Idx + 1).getReg(), &II)); + } + // Match if every element was extracted. + return ExtractedElts.all(); +} + +void CombinerHelper::applyExtractAllEltsFromBuildVector( + MachineInstr &MI, + SmallVectorImpl<std::pair<Register, MachineInstr *>> &SrcDstPairs) { + assert(MI.getOpcode() == TargetOpcode::G_BUILD_VECTOR); + for (auto &Pair : SrcDstPairs) { + auto *ExtMI = Pair.second; + replaceRegWith(MRI, ExtMI->getOperand(0).getReg(), Pair.first); + ExtMI->eraseFromParent(); + } + MI.eraseFromParent(); +} + +void CombinerHelper::applyBuildFn( MachineInstr &MI, std::function<void(MachineIRBuilder &)> &MatchInfo) { Builder.setInstrAndDebugLoc(MI); MatchInfo(Builder); MI.eraseFromParent(); +} + +void CombinerHelper::applyBuildFnNoErase( + MachineInstr &MI, std::function<void(MachineIRBuilder &)> &MatchInfo) { + Builder.setInstrAndDebugLoc(MI); + MatchInfo(Builder); +} + +/// Match an FSHL or FSHR that can be combined to a ROTR or ROTL rotate. +bool CombinerHelper::matchFunnelShiftToRotate(MachineInstr &MI) { + unsigned Opc = MI.getOpcode(); + assert(Opc == TargetOpcode::G_FSHL || Opc == TargetOpcode::G_FSHR); + Register X = MI.getOperand(1).getReg(); + Register Y = MI.getOperand(2).getReg(); + if (X != Y) + return false; + unsigned RotateOpc = + Opc == TargetOpcode::G_FSHL ? TargetOpcode::G_ROTL : TargetOpcode::G_ROTR; + return isLegalOrBeforeLegalizer({RotateOpc, {MRI.getType(X), MRI.getType(Y)}}); +} + +void CombinerHelper::applyFunnelShiftToRotate(MachineInstr &MI) { + unsigned Opc = MI.getOpcode(); + assert(Opc == TargetOpcode::G_FSHL || Opc == TargetOpcode::G_FSHR); + bool IsFSHL = Opc == TargetOpcode::G_FSHL; + Observer.changingInstr(MI); + MI.setDesc(Builder.getTII().get(IsFSHL ? TargetOpcode::G_ROTL + : TargetOpcode::G_ROTR)); + MI.RemoveOperand(2); + Observer.changedInstr(MI); +} + +// Fold (rot x, c) -> (rot x, c % BitSize) +bool CombinerHelper::matchRotateOutOfRange(MachineInstr &MI) { + assert(MI.getOpcode() == TargetOpcode::G_ROTL || + MI.getOpcode() == TargetOpcode::G_ROTR); + unsigned Bitsize = + MRI.getType(MI.getOperand(0).getReg()).getScalarSizeInBits(); + Register AmtReg = MI.getOperand(2).getReg(); + bool OutOfRange = false; + auto MatchOutOfRange = [Bitsize, &OutOfRange](const Constant *C) { + if (auto *CI = dyn_cast<ConstantInt>(C)) + OutOfRange |= CI->getValue().uge(Bitsize); + return true; + }; + return matchUnaryPredicate(MRI, AmtReg, MatchOutOfRange) && OutOfRange; +} + +void CombinerHelper::applyRotateOutOfRange(MachineInstr &MI) { + assert(MI.getOpcode() == TargetOpcode::G_ROTL || + MI.getOpcode() == TargetOpcode::G_ROTR); + unsigned Bitsize = + MRI.getType(MI.getOperand(0).getReg()).getScalarSizeInBits(); + Builder.setInstrAndDebugLoc(MI); + Register Amt = MI.getOperand(2).getReg(); + LLT AmtTy = MRI.getType(Amt); + auto Bits = Builder.buildConstant(AmtTy, Bitsize); + Amt = Builder.buildURem(AmtTy, MI.getOperand(2).getReg(), Bits).getReg(0); + Observer.changingInstr(MI); + MI.getOperand(2).setReg(Amt); + Observer.changedInstr(MI); +} + +bool CombinerHelper::matchICmpToTrueFalseKnownBits(MachineInstr &MI, + int64_t &MatchInfo) { + assert(MI.getOpcode() == TargetOpcode::G_ICMP); + auto Pred = static_cast<CmpInst::Predicate>(MI.getOperand(1).getPredicate()); + auto KnownLHS = KB->getKnownBits(MI.getOperand(2).getReg()); + auto KnownRHS = KB->getKnownBits(MI.getOperand(3).getReg()); + Optional<bool> KnownVal; + switch (Pred) { + default: + llvm_unreachable("Unexpected G_ICMP predicate?"); + case CmpInst::ICMP_EQ: + KnownVal = KnownBits::eq(KnownLHS, KnownRHS); + break; + case CmpInst::ICMP_NE: + KnownVal = KnownBits::ne(KnownLHS, KnownRHS); + break; + case CmpInst::ICMP_SGE: + KnownVal = KnownBits::sge(KnownLHS, KnownRHS); + break; + case CmpInst::ICMP_SGT: + KnownVal = KnownBits::sgt(KnownLHS, KnownRHS); + break; + case CmpInst::ICMP_SLE: + KnownVal = KnownBits::sle(KnownLHS, KnownRHS); + break; + case CmpInst::ICMP_SLT: + KnownVal = KnownBits::slt(KnownLHS, KnownRHS); + break; + case CmpInst::ICMP_UGE: + KnownVal = KnownBits::uge(KnownLHS, KnownRHS); + break; + case CmpInst::ICMP_UGT: + KnownVal = KnownBits::ugt(KnownLHS, KnownRHS); + break; + case CmpInst::ICMP_ULE: + KnownVal = KnownBits::ule(KnownLHS, KnownRHS); + break; + case CmpInst::ICMP_ULT: + KnownVal = KnownBits::ult(KnownLHS, KnownRHS); + break; + } + if (!KnownVal) + return false; + MatchInfo = + *KnownVal + ? getICmpTrueVal(getTargetLowering(), + /*IsVector = */ + MRI.getType(MI.getOperand(0).getReg()).isVector(), + /* IsFP = */ false) + : 0; + return true; +} + +/// Form a G_SBFX from a G_SEXT_INREG fed by a right shift. +bool CombinerHelper::matchBitfieldExtractFromSExtInReg( + MachineInstr &MI, std::function<void(MachineIRBuilder &)> &MatchInfo) { + assert(MI.getOpcode() == TargetOpcode::G_SEXT_INREG); + Register Dst = MI.getOperand(0).getReg(); + Register Src = MI.getOperand(1).getReg(); + LLT Ty = MRI.getType(Src); + LLT ExtractTy = getTargetLowering().getPreferredShiftAmountTy(Ty); + if (!LI || !LI->isLegalOrCustom({TargetOpcode::G_SBFX, {Ty, ExtractTy}})) + return false; + int64_t Width = MI.getOperand(2).getImm(); + Register ShiftSrc; + int64_t ShiftImm; + if (!mi_match( + Src, MRI, + m_OneNonDBGUse(m_any_of(m_GAShr(m_Reg(ShiftSrc), m_ICst(ShiftImm)), + m_GLShr(m_Reg(ShiftSrc), m_ICst(ShiftImm)))))) + return false; + if (ShiftImm < 0 || ShiftImm + Width > Ty.getScalarSizeInBits()) + return false; + + MatchInfo = [=](MachineIRBuilder &B) { + auto Cst1 = B.buildConstant(ExtractTy, ShiftImm); + auto Cst2 = B.buildConstant(ExtractTy, Width); + B.buildSbfx(Dst, ShiftSrc, Cst1, Cst2); + }; + return true; +} + +/// Form a G_UBFX from "(a srl b) & mask", where b and mask are constants. +bool CombinerHelper::matchBitfieldExtractFromAnd( + MachineInstr &MI, std::function<void(MachineIRBuilder &)> &MatchInfo) { + assert(MI.getOpcode() == TargetOpcode::G_AND); + Register Dst = MI.getOperand(0).getReg(); + LLT Ty = MRI.getType(Dst); + if (!getTargetLowering().isConstantUnsignedBitfieldExtactLegal( + TargetOpcode::G_UBFX, Ty, Ty)) + return false; + + int64_t AndImm, LSBImm; + Register ShiftSrc; + const unsigned Size = Ty.getScalarSizeInBits(); + if (!mi_match(MI.getOperand(0).getReg(), MRI, + m_GAnd(m_OneNonDBGUse(m_GLShr(m_Reg(ShiftSrc), m_ICst(LSBImm))), + m_ICst(AndImm)))) + return false; + + // The mask is a mask of the low bits iff imm & (imm+1) == 0. + auto MaybeMask = static_cast<uint64_t>(AndImm); + if (MaybeMask & (MaybeMask + 1)) + return false; + + // LSB must fit within the register. + if (static_cast<uint64_t>(LSBImm) >= Size) + return false; + + LLT ExtractTy = getTargetLowering().getPreferredShiftAmountTy(Ty); + uint64_t Width = APInt(Size, AndImm).countTrailingOnes(); + MatchInfo = [=](MachineIRBuilder &B) { + auto WidthCst = B.buildConstant(ExtractTy, Width); + auto LSBCst = B.buildConstant(ExtractTy, LSBImm); + B.buildInstr(TargetOpcode::G_UBFX, {Dst}, {ShiftSrc, LSBCst, WidthCst}); + }; + return true; +} + +bool CombinerHelper::reassociationCanBreakAddressingModePattern( + MachineInstr &PtrAdd) { + assert(PtrAdd.getOpcode() == TargetOpcode::G_PTR_ADD); + + Register Src1Reg = PtrAdd.getOperand(1).getReg(); + MachineInstr *Src1Def = getOpcodeDef(TargetOpcode::G_PTR_ADD, Src1Reg, MRI); + if (!Src1Def) + return false; + + Register Src2Reg = PtrAdd.getOperand(2).getReg(); + + if (MRI.hasOneNonDBGUse(Src1Reg)) + return false; + + auto C1 = getConstantVRegVal(Src1Def->getOperand(2).getReg(), MRI); + if (!C1) + return false; + auto C2 = getConstantVRegVal(Src2Reg, MRI); + if (!C2) + return false; + + const APInt &C1APIntVal = *C1; + const APInt &C2APIntVal = *C2; + const int64_t CombinedValue = (C1APIntVal + C2APIntVal).getSExtValue(); + + for (auto &UseMI : MRI.use_nodbg_instructions(Src1Reg)) { + // This combine may end up running before ptrtoint/inttoptr combines + // manage to eliminate redundant conversions, so try to look through them. + MachineInstr *ConvUseMI = &UseMI; + unsigned ConvUseOpc = ConvUseMI->getOpcode(); + while (ConvUseOpc == TargetOpcode::G_INTTOPTR || + ConvUseOpc == TargetOpcode::G_PTRTOINT) { + Register DefReg = ConvUseMI->getOperand(0).getReg(); + if (!MRI.hasOneNonDBGUse(DefReg)) + break; + ConvUseMI = &*MRI.use_instr_nodbg_begin(DefReg); + ConvUseOpc = ConvUseMI->getOpcode(); + } + auto LoadStore = ConvUseOpc == TargetOpcode::G_LOAD || + ConvUseOpc == TargetOpcode::G_STORE; + if (!LoadStore) + continue; + // Is x[offset2] already not a legal addressing mode? If so then + // reassociating the constants breaks nothing (we test offset2 because + // that's the one we hope to fold into the load or store). + TargetLoweringBase::AddrMode AM; + AM.HasBaseReg = true; + AM.BaseOffs = C2APIntVal.getSExtValue(); + unsigned AS = + MRI.getType(ConvUseMI->getOperand(1).getReg()).getAddressSpace(); + Type *AccessTy = + getTypeForLLT(MRI.getType(ConvUseMI->getOperand(0).getReg()), + PtrAdd.getMF()->getFunction().getContext()); + const auto &TLI = *PtrAdd.getMF()->getSubtarget().getTargetLowering(); + if (!TLI.isLegalAddressingMode(PtrAdd.getMF()->getDataLayout(), AM, + AccessTy, AS)) + continue; + + // Would x[offset1+offset2] still be a legal addressing mode? + AM.BaseOffs = CombinedValue; + if (!TLI.isLegalAddressingMode(PtrAdd.getMF()->getDataLayout(), AM, + AccessTy, AS)) + return true; + } + + return false; +} + +bool CombinerHelper::matchReassocPtrAdd( + MachineInstr &MI, std::function<void(MachineIRBuilder &)> &MatchInfo) { + assert(MI.getOpcode() == TargetOpcode::G_PTR_ADD); + // We're trying to match a few pointer computation patterns here for + // re-association opportunities. + // 1) Isolating a constant operand to be on the RHS, e.g.: + // G_PTR_ADD(BASE, G_ADD(X, C)) -> G_PTR_ADD(G_PTR_ADD(BASE, X), C) + // + // 2) Folding two constants in each sub-tree as long as such folding + // doesn't break a legal addressing mode. + // G_PTR_ADD(G_PTR_ADD(BASE, C1), C2) -> G_PTR_ADD(BASE, C1+C2) + Register Src1Reg = MI.getOperand(1).getReg(); + Register Src2Reg = MI.getOperand(2).getReg(); + MachineInstr *LHS = MRI.getVRegDef(Src1Reg); + MachineInstr *RHS = MRI.getVRegDef(Src2Reg); + + if (LHS->getOpcode() != TargetOpcode::G_PTR_ADD) { + // Try to match example 1). + if (RHS->getOpcode() != TargetOpcode::G_ADD) + return false; + auto C2 = getConstantVRegVal(RHS->getOperand(2).getReg(), MRI); + if (!C2) + return false; + + MatchInfo = [=,&MI](MachineIRBuilder &B) { + LLT PtrTy = MRI.getType(MI.getOperand(0).getReg()); + + auto NewBase = + Builder.buildPtrAdd(PtrTy, Src1Reg, RHS->getOperand(1).getReg()); + Observer.changingInstr(MI); + MI.getOperand(1).setReg(NewBase.getReg(0)); + MI.getOperand(2).setReg(RHS->getOperand(2).getReg()); + Observer.changedInstr(MI); + }; + } else { + // Try to match example 2. + Register LHSSrc1 = LHS->getOperand(1).getReg(); + Register LHSSrc2 = LHS->getOperand(2).getReg(); + auto C1 = getConstantVRegVal(LHSSrc2, MRI); + if (!C1) + return false; + auto C2 = getConstantVRegVal(Src2Reg, MRI); + if (!C2) + return false; + + MatchInfo = [=, &MI](MachineIRBuilder &B) { + auto NewCst = B.buildConstant(MRI.getType(Src2Reg), *C1 + *C2); + Observer.changingInstr(MI); + MI.getOperand(1).setReg(LHSSrc1); + MI.getOperand(2).setReg(NewCst.getReg(0)); + Observer.changedInstr(MI); + }; + } + return !reassociationCanBreakAddressingModePattern(MI); +} + +bool CombinerHelper::matchConstantFold(MachineInstr &MI, APInt &MatchInfo) { + Register Op1 = MI.getOperand(1).getReg(); + Register Op2 = MI.getOperand(2).getReg(); + auto MaybeCst = ConstantFoldBinOp(MI.getOpcode(), Op1, Op2, MRI); + if (!MaybeCst) + return false; + MatchInfo = *MaybeCst; return true; } |
