| 1 | //===- GCNRegPressure.cpp -------------------------------------------------===// |
| 2 | // |
| 3 | // Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions. |
| 4 | // See https://llvm.org/LICENSE.txt for license information. |
| 5 | // SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception |
| 6 | // |
| 7 | //===----------------------------------------------------------------------===// |
| 8 | /// |
| 9 | /// \file |
| 10 | /// This file implements the GCNRegPressure class. |
| 11 | /// |
| 12 | //===----------------------------------------------------------------------===// |
| 13 | |
| 14 | #include "GCNRegPressure.h" |
| 15 | #include "AMDGPU.h" |
| 16 | #include "SIMachineFunctionInfo.h" |
| 17 | #include "llvm/ADT/SetVector.h" |
| 18 | #include "llvm/CodeGen/LiveIntervalUnion.h" |
| 19 | #include "llvm/CodeGen/MachineBasicBlock.h" |
| 20 | #include "llvm/CodeGen/MachineLoopInfo.h" |
| 21 | #include "llvm/CodeGen/RegisterPressure.h" |
| 22 | |
| 23 | using namespace llvm; |
| 24 | |
| 25 | #define DEBUG_TYPE "machine-scheduler" |
| 26 | |
| 27 | bool llvm::isEqual(const GCNRPTracker::LiveRegSet &S1, |
| 28 | const GCNRPTracker::LiveRegSet &S2) { |
| 29 | if (S1.size() != S2.size()) |
| 30 | return false; |
| 31 | |
| 32 | for (const auto &P : S1) { |
| 33 | auto I = S2.find(Val: P.first); |
| 34 | if (I == S2.end() || I->second != P.second) |
| 35 | return false; |
| 36 | } |
| 37 | return true; |
| 38 | } |
| 39 | |
| 40 | /////////////////////////////////////////////////////////////////////////////// |
| 41 | // GCNRegPressure |
| 42 | |
| 43 | unsigned GCNRegPressure::getRegKind(const TargetRegisterClass *RC, |
| 44 | const SIRegisterInfo *STI) { |
| 45 | return STI->isSGPRClass(RC) |
| 46 | ? SGPR |
| 47 | : (STI->isAGPRClass(RC) |
| 48 | ? AGPR |
| 49 | : (STI->isVectorSuperClass(RC) ? AVGPR : VGPR)); |
| 50 | } |
| 51 | |
| 52 | void GCNRegPressure::inc(unsigned Reg, |
| 53 | LaneBitmask PrevMask, |
| 54 | LaneBitmask NewMask, |
| 55 | const MachineRegisterInfo &MRI) { |
| 56 | unsigned NewNumCoveredRegs = SIRegisterInfo::getNumCoveredRegs(LM: NewMask); |
| 57 | unsigned PrevNumCoveredRegs = SIRegisterInfo::getNumCoveredRegs(LM: PrevMask); |
| 58 | if (NewNumCoveredRegs == PrevNumCoveredRegs) |
| 59 | return; |
| 60 | |
| 61 | int Sign = 1; |
| 62 | if (NewMask < PrevMask) { |
| 63 | std::swap(a&: NewMask, b&: PrevMask); |
| 64 | std::swap(a&: NewNumCoveredRegs, b&: PrevNumCoveredRegs); |
| 65 | Sign = -1; |
| 66 | } |
| 67 | assert(PrevMask < NewMask && PrevNumCoveredRegs < NewNumCoveredRegs && |
| 68 | "prev mask should always be lesser than new" ); |
| 69 | |
| 70 | const TargetRegisterClass *RC = MRI.getRegClass(Reg); |
| 71 | const TargetRegisterInfo *TRI = MRI.getTargetRegisterInfo(); |
| 72 | const SIRegisterInfo *STI = static_cast<const SIRegisterInfo *>(TRI); |
| 73 | unsigned RegKind = getRegKind(RC, STI); |
| 74 | if (TRI->getRegSizeInBits(RC: *RC) != 32) { |
| 75 | // Reg is from a tuple register class. |
| 76 | if (PrevMask.none()) { |
| 77 | unsigned TupleIdx = TOTAL_KINDS + RegKind; |
| 78 | Value[TupleIdx] += Sign * TRI->getRegClassWeight(RC).RegWeight; |
| 79 | } |
| 80 | // Pressure scales with number of new registers covered by the new mask. |
| 81 | // Note when true16 is enabled, we can no longer safely use the following |
| 82 | // approach to calculate the difference in the number of 32-bit registers |
| 83 | // between two masks: |
| 84 | // |
| 85 | // Sign *= SIRegisterInfo::getNumCoveredRegs(~PrevMask & NewMask); |
| 86 | // |
| 87 | // The issue is that the mask calculation `~PrevMask & NewMask` doesn't |
| 88 | // properly account for partial usage of a 32-bit register when dealing with |
| 89 | // 16-bit registers. |
| 90 | // |
| 91 | // Consider this example: |
| 92 | // Assume PrevMask = 0b0010 and NewMask = 0b1111. Here, the correct register |
| 93 | // usage difference should be 1, because even though PrevMask uses only half |
| 94 | // of a 32-bit register, it should still be counted as a full register use. |
| 95 | // However, the mask calculation yields `~PrevMask & NewMask = 0b1101`, and |
| 96 | // calling `getNumCoveredRegs` returns 2 instead of 1. This incorrect |
| 97 | // calculation can lead to integer overflow when Sign = -1. |
| 98 | Sign *= NewNumCoveredRegs - PrevNumCoveredRegs; |
| 99 | } |
| 100 | Value[RegKind] += Sign; |
| 101 | } |
| 102 | |
| 103 | namespace { |
| 104 | struct RegExcess { |
| 105 | unsigned SGPR = 0; |
| 106 | unsigned VGPR = 0; |
| 107 | unsigned ArchVGPR = 0; |
| 108 | unsigned AGPR = 0; |
| 109 | |
| 110 | bool anyExcess() const { return SGPR || VGPR || ArchVGPR || AGPR; } |
| 111 | bool hasVectorRegisterExcess() const { return VGPR || ArchVGPR || AGPR; } |
| 112 | |
| 113 | RegExcess(const MachineFunction &MF, const GCNRegPressure &RP) |
| 114 | : RegExcess(MF, RP, GCNRPTarget(MF, RP)) {} |
| 115 | RegExcess(const MachineFunction &MF, const GCNRegPressure &RP, |
| 116 | const GCNRPTarget &Target) { |
| 117 | unsigned MaxSGPRs = Target.getMaxSGPRs(); |
| 118 | unsigned MaxVGPRs = Target.getMaxVGPRs(); |
| 119 | |
| 120 | const GCNSubtarget &ST = MF.getSubtarget<GCNSubtarget>(); |
| 121 | SGPR = std::max(a: static_cast<int>(RP.getSGPRNum() - MaxSGPRs), b: 0); |
| 122 | |
| 123 | // The number of virtual VGPRs required to handle excess SGPR |
| 124 | unsigned WaveSize = ST.getWavefrontSize(); |
| 125 | unsigned VGPRForSGPRSpills = divideCeil(Numerator: SGPR, Denominator: WaveSize); |
| 126 | |
| 127 | unsigned MaxArchVGPRs = ST.getAddressableNumArchVGPRs(); |
| 128 | |
| 129 | // Unified excess pressure conditions, accounting for VGPRs used for SGPR |
| 130 | // spills |
| 131 | VGPR = std::max(a: static_cast<int>(RP.getVGPRNum(UnifiedVGPRFile: ST.hasGFX90AInsts()) + |
| 132 | VGPRForSGPRSpills - MaxVGPRs), |
| 133 | b: 0); |
| 134 | |
| 135 | unsigned ArchVGPRLimit = ST.hasGFX90AInsts() ? MaxArchVGPRs : MaxVGPRs; |
| 136 | // Arch VGPR excess pressure conditions, accounting for VGPRs used for SGPR |
| 137 | // spills |
| 138 | ArchVGPR = std::max(a: static_cast<int>(RP.getArchVGPRNum() + |
| 139 | VGPRForSGPRSpills - ArchVGPRLimit), |
| 140 | b: 0); |
| 141 | |
| 142 | // AGPR excess pressure conditions |
| 143 | AGPR = std::max(a: static_cast<int>(RP.getAGPRNum() - ArchVGPRLimit), b: 0); |
| 144 | } |
| 145 | }; |
| 146 | } // namespace |
| 147 | |
| 148 | bool GCNRegPressure::less(const MachineFunction &MF, const GCNRegPressure &O, |
| 149 | unsigned MaxOccupancy) const { |
| 150 | const GCNSubtarget &ST = MF.getSubtarget<GCNSubtarget>(); |
| 151 | unsigned DynamicVGPRBlockSize = |
| 152 | MF.getInfo<SIMachineFunctionInfo>()->getDynamicVGPRBlockSize(); |
| 153 | |
| 154 | const auto SGPROcc = std::min(a: MaxOccupancy, |
| 155 | b: ST.getOccupancyWithNumSGPRs(SGPRs: getSGPRNum())); |
| 156 | const auto VGPROcc = std::min( |
| 157 | a: MaxOccupancy, b: ST.getOccupancyWithNumVGPRs(VGPRs: getVGPRNum(UnifiedVGPRFile: ST.hasGFX90AInsts()), |
| 158 | DynamicVGPRBlockSize)); |
| 159 | const auto OtherSGPROcc = std::min(a: MaxOccupancy, |
| 160 | b: ST.getOccupancyWithNumSGPRs(SGPRs: O.getSGPRNum())); |
| 161 | const auto OtherVGPROcc = |
| 162 | std::min(a: MaxOccupancy, |
| 163 | b: ST.getOccupancyWithNumVGPRs(VGPRs: O.getVGPRNum(UnifiedVGPRFile: ST.hasGFX90AInsts()), |
| 164 | DynamicVGPRBlockSize)); |
| 165 | |
| 166 | const auto Occ = std::min(a: SGPROcc, b: VGPROcc); |
| 167 | const auto OtherOcc = std::min(a: OtherSGPROcc, b: OtherVGPROcc); |
| 168 | |
| 169 | // Give first precedence to the better occupancy. |
| 170 | if (Occ != OtherOcc) |
| 171 | return Occ > OtherOcc; |
| 172 | |
| 173 | unsigned MaxVGPRs = ST.getMaxNumVGPRs(MF); |
| 174 | |
| 175 | RegExcess Excess(MF, *this); |
| 176 | RegExcess OtherExcess(MF, O); |
| 177 | |
| 178 | unsigned MaxArchVGPRs = ST.getAddressableNumArchVGPRs(); |
| 179 | |
| 180 | bool ExcessRP = Excess.anyExcess(); |
| 181 | bool OtherExcessRP = OtherExcess.anyExcess(); |
| 182 | |
| 183 | // Give second precedence to the reduced number of spills to hold the register |
| 184 | // pressure. |
| 185 | if (ExcessRP || OtherExcessRP) { |
| 186 | // The difference in excess VGPR pressure, after including VGPRs used for |
| 187 | // SGPR spills |
| 188 | int VGPRDiff = |
| 189 | ((OtherExcess.VGPR + OtherExcess.ArchVGPR + OtherExcess.AGPR) - |
| 190 | (Excess.VGPR + Excess.ArchVGPR + Excess.AGPR)); |
| 191 | |
| 192 | int SGPRDiff = OtherExcess.SGPR - Excess.SGPR; |
| 193 | |
| 194 | if (VGPRDiff != 0) |
| 195 | return VGPRDiff > 0; |
| 196 | if (SGPRDiff != 0) { |
| 197 | unsigned PureExcessVGPR = |
| 198 | std::max(a: static_cast<int>(getVGPRNum(UnifiedVGPRFile: ST.hasGFX90AInsts()) - MaxVGPRs), |
| 199 | b: 0) + |
| 200 | std::max(a: static_cast<int>(getVGPRNum(UnifiedVGPRFile: false) - MaxArchVGPRs), b: 0); |
| 201 | unsigned OtherPureExcessVGPR = |
| 202 | std::max( |
| 203 | a: static_cast<int>(O.getVGPRNum(UnifiedVGPRFile: ST.hasGFX90AInsts()) - MaxVGPRs), |
| 204 | b: 0) + |
| 205 | std::max(a: static_cast<int>(O.getVGPRNum(UnifiedVGPRFile: false) - MaxArchVGPRs), b: 0); |
| 206 | |
| 207 | // If we have a special case where there is a tie in excess VGPR, but one |
| 208 | // of the pressures has VGPR usage from SGPR spills, prefer the pressure |
| 209 | // with SGPR spills. |
| 210 | if (PureExcessVGPR != OtherPureExcessVGPR) |
| 211 | return SGPRDiff < 0; |
| 212 | // If both pressures have the same excess pressure before and after |
| 213 | // accounting for SGPR spills, prefer fewer SGPR spills. |
| 214 | return SGPRDiff > 0; |
| 215 | } |
| 216 | } |
| 217 | |
| 218 | bool SGPRImportant = SGPROcc < VGPROcc; |
| 219 | const bool OtherSGPRImportant = OtherSGPROcc < OtherVGPROcc; |
| 220 | |
| 221 | // If both pressures disagree on what is more important compare vgprs. |
| 222 | if (SGPRImportant != OtherSGPRImportant) { |
| 223 | SGPRImportant = false; |
| 224 | } |
| 225 | |
| 226 | // Give third precedence to lower register tuple pressure. |
| 227 | bool SGPRFirst = SGPRImportant; |
| 228 | for (int I = 2; I > 0; --I, SGPRFirst = !SGPRFirst) { |
| 229 | if (SGPRFirst) { |
| 230 | auto SW = getSGPRTuplesWeight(); |
| 231 | auto OtherSW = O.getSGPRTuplesWeight(); |
| 232 | if (SW != OtherSW) |
| 233 | return SW < OtherSW; |
| 234 | } else { |
| 235 | auto VW = getVGPRTuplesWeight(); |
| 236 | auto OtherVW = O.getVGPRTuplesWeight(); |
| 237 | if (VW != OtherVW) |
| 238 | return VW < OtherVW; |
| 239 | } |
| 240 | } |
| 241 | |
| 242 | // Give final precedence to lower general RP. |
| 243 | return SGPRImportant ? (getSGPRNum() < O.getSGPRNum()): |
| 244 | (getVGPRNum(UnifiedVGPRFile: ST.hasGFX90AInsts()) < |
| 245 | O.getVGPRNum(UnifiedVGPRFile: ST.hasGFX90AInsts())); |
| 246 | } |
| 247 | |
| 248 | Printable llvm::print(const GCNRegPressure &RP, const GCNSubtarget *ST, |
| 249 | unsigned DynamicVGPRBlockSize) { |
| 250 | return Printable([&RP, ST, DynamicVGPRBlockSize](raw_ostream &OS) { |
| 251 | OS << "VGPRs: " << RP.getArchVGPRNum() << ' ' |
| 252 | << "AGPRs: " << RP.getAGPRNum(); |
| 253 | if (ST) |
| 254 | OS << "(O" |
| 255 | << ST->getOccupancyWithNumVGPRs(VGPRs: RP.getVGPRNum(UnifiedVGPRFile: ST->hasGFX90AInsts()), |
| 256 | DynamicVGPRBlockSize) |
| 257 | << ')'; |
| 258 | OS << ", SGPRs: " << RP.getSGPRNum(); |
| 259 | if (ST) |
| 260 | OS << "(O" << ST->getOccupancyWithNumSGPRs(SGPRs: RP.getSGPRNum()) << ')'; |
| 261 | OS << ", LVGPR WT: " << RP.getVGPRTuplesWeight() |
| 262 | << ", LSGPR WT: " << RP.getSGPRTuplesWeight(); |
| 263 | if (ST) |
| 264 | OS << " -> Occ: " << RP.getOccupancy(ST: *ST, DynamicVGPRBlockSize); |
| 265 | OS << '\n'; |
| 266 | }); |
| 267 | } |
| 268 | |
| 269 | static LaneBitmask getDefRegMask(const MachineOperand &MO, |
| 270 | const MachineRegisterInfo &MRI) { |
| 271 | assert(MO.isDef() && MO.isReg() && MO.getReg().isVirtual()); |
| 272 | |
| 273 | // We don't rely on read-undef flag because in case of tentative schedule |
| 274 | // tracking it isn't set correctly yet. This works correctly however since |
| 275 | // use mask has been tracked before using LIS. |
| 276 | return MO.getSubReg() == 0 ? |
| 277 | MRI.getMaxLaneMaskForVReg(Reg: MO.getReg()) : |
| 278 | MRI.getTargetRegisterInfo()->getSubRegIndexLaneMask(SubIdx: MO.getSubReg()); |
| 279 | } |
| 280 | |
| 281 | static void |
| 282 | collectVirtualRegUses(SmallVectorImpl<VRegMaskOrUnit> &VRegMaskOrUnits, |
| 283 | const MachineInstr &MI, const LiveIntervals &LIS, |
| 284 | const MachineRegisterInfo &MRI) { |
| 285 | |
| 286 | auto &TRI = *MRI.getTargetRegisterInfo(); |
| 287 | for (const auto &MO : MI.operands()) { |
| 288 | if (!MO.isReg() || !MO.getReg().isVirtual()) |
| 289 | continue; |
| 290 | if (!MO.isUse() || !MO.readsReg()) |
| 291 | continue; |
| 292 | |
| 293 | Register Reg = MO.getReg(); |
| 294 | auto I = llvm::find_if(Range&: VRegMaskOrUnits, P: [Reg](const VRegMaskOrUnit &RM) { |
| 295 | return RM.VRegOrUnit.asVirtualReg() == Reg; |
| 296 | }); |
| 297 | |
| 298 | auto &P = I == VRegMaskOrUnits.end() |
| 299 | ? VRegMaskOrUnits.emplace_back(Args: VirtRegOrUnit(Reg), |
| 300 | Args: LaneBitmask::getNone()) |
| 301 | : *I; |
| 302 | |
| 303 | P.LaneMask |= MO.getSubReg() ? TRI.getSubRegIndexLaneMask(SubIdx: MO.getSubReg()) |
| 304 | : MRI.getMaxLaneMaskForVReg(Reg); |
| 305 | } |
| 306 | |
| 307 | SlotIndex InstrSI; |
| 308 | for (auto &P : VRegMaskOrUnits) { |
| 309 | auto &LI = LIS.getInterval(Reg: P.VRegOrUnit.asVirtualReg()); |
| 310 | if (!LI.hasSubRanges()) |
| 311 | continue; |
| 312 | |
| 313 | // For a tentative schedule LIS isn't updated yet but livemask should |
| 314 | // remain the same on any schedule. Subreg defs can be reordered but they |
| 315 | // all must dominate uses anyway. |
| 316 | if (!InstrSI) |
| 317 | InstrSI = LIS.getInstructionIndex(Instr: MI).getBaseIndex(); |
| 318 | |
| 319 | P.LaneMask = getLiveLaneMask(LI, SI: InstrSI, MRI, LaneMaskFilter: P.LaneMask); |
| 320 | } |
| 321 | } |
| 322 | |
| 323 | /// Mostly copy/paste from CodeGen/RegisterPressure.cpp |
| 324 | static LaneBitmask getLanesWithProperty( |
| 325 | const LiveIntervals &LIS, const MachineRegisterInfo &MRI, |
| 326 | bool TrackLaneMasks, Register Reg, SlotIndex Pos, |
| 327 | function_ref<bool(const LiveRange &LR, SlotIndex Pos)> Property) { |
| 328 | assert(Reg.isVirtual()); |
| 329 | const LiveInterval &LI = LIS.getInterval(Reg); |
| 330 | LaneBitmask Result; |
| 331 | if (TrackLaneMasks && LI.hasSubRanges()) { |
| 332 | for (const LiveInterval::SubRange &SR : LI.subranges()) { |
| 333 | if (Property(SR, Pos)) |
| 334 | Result |= SR.LaneMask; |
| 335 | } |
| 336 | } else if (Property(LI, Pos)) { |
| 337 | Result = |
| 338 | TrackLaneMasks ? MRI.getMaxLaneMaskForVReg(Reg) : LaneBitmask::getAll(); |
| 339 | } |
| 340 | |
| 341 | return Result; |
| 342 | } |
| 343 | |
| 344 | /// Mostly copy/paste from CodeGen/RegisterPressure.cpp |
| 345 | /// Helper to find a vreg use between two indices {PriorUseIdx, NextUseIdx}. |
| 346 | /// The query starts with a lane bitmask which gets lanes/bits removed for every |
| 347 | /// use we find. |
| 348 | static LaneBitmask findUseBetween(unsigned Reg, LaneBitmask LastUseMask, |
| 349 | SlotIndex PriorUseIdx, SlotIndex NextUseIdx, |
| 350 | const MachineRegisterInfo &MRI, |
| 351 | const SIRegisterInfo *TRI, |
| 352 | const LiveIntervals *LIS, |
| 353 | bool Upward = false) { |
| 354 | for (const MachineOperand &MO : MRI.use_nodbg_operands(Reg)) { |
| 355 | if (MO.isUndef()) |
| 356 | continue; |
| 357 | const MachineInstr *MI = MO.getParent(); |
| 358 | SlotIndex InstSlot = LIS->getInstructionIndex(Instr: *MI).getRegSlot(); |
| 359 | bool InRange = Upward ? (InstSlot > PriorUseIdx && InstSlot <= NextUseIdx) |
| 360 | : (InstSlot >= PriorUseIdx && InstSlot < NextUseIdx); |
| 361 | if (!InRange) |
| 362 | continue; |
| 363 | |
| 364 | unsigned SubRegIdx = MO.getSubReg(); |
| 365 | LaneBitmask UseMask = TRI->getSubRegIndexLaneMask(SubIdx: SubRegIdx); |
| 366 | LastUseMask &= ~UseMask; |
| 367 | if (LastUseMask.none()) |
| 368 | return LaneBitmask::getNone(); |
| 369 | } |
| 370 | return LastUseMask; |
| 371 | } |
| 372 | |
| 373 | //////////////////////////////////////////////////////////////////////////////// |
| 374 | // GCNRPTarget |
| 375 | |
| 376 | GCNRPTarget::GCNRPTarget(const MachineFunction &MF, const GCNRegPressure &RP) |
| 377 | : GCNRPTarget(RP, MF) { |
| 378 | const Function &F = MF.getFunction(); |
| 379 | const GCNSubtarget &ST = MF.getSubtarget<GCNSubtarget>(); |
| 380 | setTarget(NumSGPRs: ST.getMaxNumSGPRs(F), NumVGPRs: ST.getMaxNumVGPRs(F)); |
| 381 | } |
| 382 | |
| 383 | GCNRPTarget::GCNRPTarget(unsigned NumSGPRs, unsigned NumVGPRs, |
| 384 | const MachineFunction &MF, const GCNRegPressure &RP) |
| 385 | : GCNRPTarget(RP, MF) { |
| 386 | setTarget(NumSGPRs, NumVGPRs); |
| 387 | } |
| 388 | |
| 389 | GCNRPTarget::GCNRPTarget(unsigned Occupancy, const MachineFunction &MF, |
| 390 | const GCNRegPressure &RP) |
| 391 | : GCNRPTarget(RP, MF) { |
| 392 | const GCNSubtarget &ST = MF.getSubtarget<GCNSubtarget>(); |
| 393 | unsigned DynamicVGPRBlockSize = |
| 394 | MF.getInfo<SIMachineFunctionInfo>()->getDynamicVGPRBlockSize(); |
| 395 | setTarget(NumSGPRs: ST.getMaxNumSGPRs(WavesPerEU: Occupancy, /*Addressable=*/false), |
| 396 | NumVGPRs: ST.getMaxNumVGPRs(WavesPerEU: Occupancy, DynamicVGPRBlockSize)); |
| 397 | } |
| 398 | |
| 399 | void GCNRPTarget::setTarget(unsigned NumSGPRs, unsigned NumVGPRs) { |
| 400 | const GCNSubtarget &ST = MF.getSubtarget<GCNSubtarget>(); |
| 401 | MaxSGPRs = std::min(a: ST.getAddressableNumSGPRs(), b: NumSGPRs); |
| 402 | MaxVGPRs = std::min(a: ST.getAddressableNumArchVGPRs(), b: NumVGPRs); |
| 403 | if (UnifiedRF) { |
| 404 | unsigned DynamicVGPRBlockSize = |
| 405 | MF.getInfo<SIMachineFunctionInfo>()->getDynamicVGPRBlockSize(); |
| 406 | MaxUnifiedVGPRs = |
| 407 | std::min(a: ST.getAddressableNumVGPRs(DynamicVGPRBlockSize), b: NumVGPRs); |
| 408 | } else { |
| 409 | MaxUnifiedVGPRs = 0; |
| 410 | } |
| 411 | } |
| 412 | |
| 413 | bool GCNRPTarget::isSaveBeneficial(Register Reg) const { |
| 414 | const MachineRegisterInfo &MRI = MF.getRegInfo(); |
| 415 | const TargetRegisterClass *RC = MRI.getRegClass(Reg); |
| 416 | const TargetRegisterInfo *TRI = MRI.getTargetRegisterInfo(); |
| 417 | const SIRegisterInfo *SRI = static_cast<const SIRegisterInfo *>(TRI); |
| 418 | |
| 419 | RegExcess Excess(MF, RP, *this); |
| 420 | |
| 421 | if (SRI->isSGPRClass(RC)) |
| 422 | return Excess.SGPR; |
| 423 | |
| 424 | if (SRI->isAGPRClass(RC)) |
| 425 | return (UnifiedRF && Excess.VGPR) || Excess.AGPR; |
| 426 | |
| 427 | return (UnifiedRF && Excess.VGPR) || Excess.ArchVGPR; |
| 428 | } |
| 429 | |
| 430 | bool GCNRPTarget::isSaveBeneficial(const GCNRegPressure &SaveRP) const { |
| 431 | RegExcess Excess(MF, RP, *this); |
| 432 | if (SaveRP.getSGPRNum() != 0 && Excess.SGPR != 0) |
| 433 | return true; |
| 434 | if (SaveRP.getArchVGPRNum() != 0 && Excess.ArchVGPR != 0) |
| 435 | return true; |
| 436 | if (SaveRP.getAGPRNum() != 0 && Excess.AGPR != 0) |
| 437 | return true; |
| 438 | if (UnifiedRF && Excess.VGPR != 0) |
| 439 | return SaveRP.getArchVGPRNum() != 0 || SaveRP.getAGPRNum() != 0; |
| 440 | return false; |
| 441 | } |
| 442 | |
| 443 | unsigned GCNRPTarget::getNumRegsBenefit(const GCNRegPressure &SaveRP) const { |
| 444 | RegExcess Excess(MF, RP, *this); |
| 445 | const unsigned NumVGPRAboveAddrLimit = |
| 446 | std::min(a: Excess.ArchVGPR, b: SaveRP.getArchVGPRNum()) + |
| 447 | std::min(a: Excess.AGPR, b: SaveRP.getAGPRNum()); |
| 448 | unsigned NumRegsSaved = |
| 449 | std::min(a: Excess.SGPR, b: SaveRP.getSGPRNum()) + NumVGPRAboveAddrLimit; |
| 450 | |
| 451 | if (UnifiedRF && Excess.VGPR) { |
| 452 | // We have already accounted for excess pressure above addressive limits for |
| 453 | // the individual VGPR classes. However for targets with unified RFs there |
| 454 | // is also a unified VGPR pressure (ArchVGPR + AGPR combination) limit to |
| 455 | // honor that may be more restrictive that the per-VGPR-class limits. We |
| 456 | // must also be careful not to double-count VGPR saves that may contribute |
| 457 | // to lowering pressure both above the addressable limit in their respective |
| 458 | // class as well as in the unified VGPR limit. |
| 459 | const unsigned VGPRSave = SaveRP.getArchVGPRNum() + SaveRP.getAGPRNum(); |
| 460 | if (NumVGPRAboveAddrLimit < VGPRSave) |
| 461 | NumRegsSaved += std::min(a: Excess.VGPR, b: VGPRSave - NumVGPRAboveAddrLimit); |
| 462 | } |
| 463 | |
| 464 | return NumRegsSaved; |
| 465 | } |
| 466 | |
| 467 | bool GCNRPTarget::satisfied(const GCNRegPressure &TestRP) const { |
| 468 | if (TestRP.getSGPRNum() > MaxSGPRs || TestRP.getVGPRNum(UnifiedVGPRFile: false) > MaxVGPRs) |
| 469 | return false; |
| 470 | if (UnifiedRF && TestRP.getVGPRNum(UnifiedVGPRFile: true) > MaxUnifiedVGPRs) |
| 471 | return false; |
| 472 | return true; |
| 473 | } |
| 474 | |
| 475 | bool GCNRPTarget::hasVectorRegisterExcess() const { |
| 476 | RegExcess Excess(MF, RP, *this); |
| 477 | return Excess.hasVectorRegisterExcess(); |
| 478 | } |
| 479 | |
| 480 | /////////////////////////////////////////////////////////////////////////////// |
| 481 | // GCNRPTracker |
| 482 | |
| 483 | LaneBitmask llvm::getLiveLaneMask(unsigned Reg, SlotIndex SI, |
| 484 | const LiveIntervals &LIS, |
| 485 | const MachineRegisterInfo &MRI, |
| 486 | LaneBitmask LaneMaskFilter) { |
| 487 | return getLiveLaneMask(LI: LIS.getInterval(Reg), SI, MRI, LaneMaskFilter); |
| 488 | } |
| 489 | |
| 490 | LaneBitmask llvm::getLiveLaneMask(const LiveInterval &LI, SlotIndex SI, |
| 491 | const MachineRegisterInfo &MRI, |
| 492 | LaneBitmask LaneMaskFilter) { |
| 493 | LaneBitmask LiveMask; |
| 494 | if (LI.hasSubRanges()) { |
| 495 | for (const auto &S : LI.subranges()) |
| 496 | if ((S.LaneMask & LaneMaskFilter).any() && S.liveAt(index: SI)) { |
| 497 | LiveMask |= S.LaneMask; |
| 498 | assert(LiveMask == (LiveMask & MRI.getMaxLaneMaskForVReg(LI.reg()))); |
| 499 | } |
| 500 | } else if (LI.liveAt(index: SI)) { |
| 501 | LiveMask = MRI.getMaxLaneMaskForVReg(Reg: LI.reg()); |
| 502 | } |
| 503 | LiveMask &= LaneMaskFilter; |
| 504 | return LiveMask; |
| 505 | } |
| 506 | |
| 507 | GCNRPTracker::LiveRegSet llvm::getLiveRegs(SlotIndex SI, |
| 508 | const LiveIntervals &LIS, |
| 509 | const MachineRegisterInfo &MRI, |
| 510 | GCNRegPressure::RegKind RegKind) { |
| 511 | GCNRPTracker::LiveRegSet LiveRegs; |
| 512 | for (unsigned I = 0, E = MRI.getNumVirtRegs(); I != E; ++I) { |
| 513 | auto Reg = Register::index2VirtReg(Index: I); |
| 514 | if (RegKind != GCNRegPressure::TOTAL_KINDS && |
| 515 | GCNRegPressure::getRegKind(Reg, MRI) != RegKind) |
| 516 | continue; |
| 517 | if (!LIS.hasInterval(Reg)) |
| 518 | continue; |
| 519 | auto LiveMask = getLiveLaneMask(Reg, SI, LIS, MRI); |
| 520 | if (LiveMask.any()) |
| 521 | LiveRegs[Reg] = LiveMask; |
| 522 | } |
| 523 | return LiveRegs; |
| 524 | } |
| 525 | |
| 526 | void GCNRPTracker::reset(const MachineInstr &MI, bool After) { |
| 527 | const MachineRegisterInfo &MRI = MI.getMF()->getRegInfo(); |
| 528 | if (!MI.isDebugInstr()) { |
| 529 | SlotIndex SI = LIS.getInstructionIndex(Instr: MI); |
| 530 | if (After) |
| 531 | SI = SI.getDeadSlot(); |
| 532 | reset(MRI, SI); |
| 533 | return; |
| 534 | } |
| 535 | |
| 536 | // Look for the first valid index after the provided debug MI. |
| 537 | MachineBasicBlock::const_iterator It = MI.getIterator(), |
| 538 | MBBEnd = MI.getParent()->end(); |
| 539 | MachineBasicBlock::const_iterator NonDbgMI = |
| 540 | skipDebugInstructionsForward(It, End: MBBEnd); |
| 541 | if (NonDbgMI == MBBEnd) { |
| 542 | // There are no non-debug instruction between MI and the end of the |
| 543 | // block, so we reset the tracker at the end of the block. |
| 544 | reset(MBB: *MI.getParent(), /*End=*/true); |
| 545 | return; |
| 546 | } |
| 547 | // MI is a debug instruction so register pressure before or after it is |
| 548 | // identical. Since we moved forward to finding a non-debug instruction |
| 549 | // in the block, we reset the tracker before that instruction i.e., at its |
| 550 | // base index. |
| 551 | reset(MRI, SI: LIS.getInstructionIndex(Instr: *NonDbgMI)); |
| 552 | } |
| 553 | |
| 554 | void GCNRPTracker::reset(const MachineBasicBlock &MBB, bool End) { |
| 555 | SlotIndex SI = End ? LIS.getSlotIndexes()->getMBBLastIdx(MBB: &MBB) |
| 556 | : LIS.getMBBStartIdx(mbb: &MBB); |
| 557 | reset(MRI: MBB.getParent()->getRegInfo(), SI); |
| 558 | } |
| 559 | |
| 560 | void GCNRPTracker::reset(const MachineRegisterInfo &MRI, SlotIndex SI) { |
| 561 | this->MRI = &MRI; |
| 562 | LastTrackedMI = nullptr; |
| 563 | LiveRegs = llvm::getLiveRegs(SI, LIS, MRI); |
| 564 | MaxPressure = CurPressure = getRegPressure(MRI, LiveRegs); |
| 565 | } |
| 566 | |
| 567 | void GCNRPTracker::reset(const MachineRegisterInfo &MRI, |
| 568 | const LiveRegSet &LiveRegs) { |
| 569 | this->MRI = &MRI; |
| 570 | LastTrackedMI = nullptr; |
| 571 | if (&this->LiveRegs != &LiveRegs) |
| 572 | this->LiveRegs = LiveRegs; |
| 573 | MaxPressure = CurPressure = getRegPressure(MRI, LiveRegs); |
| 574 | } |
| 575 | |
| 576 | /// Mostly copy/paste from CodeGen/RegisterPressure.cpp |
| 577 | LaneBitmask GCNRPTracker::getLastUsedLanes(Register Reg, SlotIndex Pos) const { |
| 578 | return getLanesWithProperty( |
| 579 | LIS, MRI: *MRI, TrackLaneMasks: true, Reg, Pos: Pos.getBaseIndex(), |
| 580 | Property: [](const LiveRange &LR, SlotIndex Pos) { |
| 581 | const LiveRange::Segment *S = LR.getSegmentContaining(Idx: Pos); |
| 582 | return S != nullptr && S->end == Pos.getRegSlot(); |
| 583 | }); |
| 584 | } |
| 585 | |
| 586 | //////////////////////////////////////////////////////////////////////////////// |
| 587 | // GCNUpwardRPTracker |
| 588 | |
| 589 | void GCNUpwardRPTracker::recede(const MachineInstr &MI) { |
| 590 | assert(MRI && "call reset first" ); |
| 591 | |
| 592 | LastTrackedMI = &MI; |
| 593 | |
| 594 | if (MI.isDebugInstr()) |
| 595 | return; |
| 596 | |
| 597 | // Kill all defs. |
| 598 | GCNRegPressure DefPressure, ECDefPressure; |
| 599 | bool HasECDefs = false; |
| 600 | for (const MachineOperand &MO : MI.all_defs()) { |
| 601 | if (!MO.getReg().isVirtual()) |
| 602 | continue; |
| 603 | |
| 604 | Register Reg = MO.getReg(); |
| 605 | LaneBitmask DefMask = getDefRegMask(MO, MRI: *MRI); |
| 606 | |
| 607 | // Treat a def as fully live at the moment of definition: keep a record. |
| 608 | if (MO.isEarlyClobber()) { |
| 609 | ECDefPressure.inc(Reg, PrevMask: LaneBitmask::getNone(), NewMask: DefMask, MRI: *MRI); |
| 610 | HasECDefs = true; |
| 611 | } else |
| 612 | DefPressure.inc(Reg, PrevMask: LaneBitmask::getNone(), NewMask: DefMask, MRI: *MRI); |
| 613 | |
| 614 | auto I = LiveRegs.find(Val: Reg); |
| 615 | if (I == LiveRegs.end()) |
| 616 | continue; |
| 617 | |
| 618 | LaneBitmask &LiveMask = I->second; |
| 619 | LaneBitmask PrevMask = LiveMask; |
| 620 | LiveMask &= ~DefMask; |
| 621 | CurPressure.inc(Reg, PrevMask, NewMask: LiveMask, MRI: *MRI); |
| 622 | if (LiveMask.none()) |
| 623 | LiveRegs.erase(I); |
| 624 | } |
| 625 | |
| 626 | // Update MaxPressure with defs pressure. |
| 627 | DefPressure += CurPressure; |
| 628 | if (HasECDefs) |
| 629 | DefPressure += ECDefPressure; |
| 630 | MaxPressure = max(P1: DefPressure, P2: MaxPressure); |
| 631 | |
| 632 | // Make uses alive. |
| 633 | SmallVector<VRegMaskOrUnit, 8> RegUses; |
| 634 | collectVirtualRegUses(VRegMaskOrUnits&: RegUses, MI, LIS, MRI: *MRI); |
| 635 | for (const VRegMaskOrUnit &U : RegUses) { |
| 636 | LaneBitmask &LiveMask = LiveRegs[U.VRegOrUnit.asVirtualReg()]; |
| 637 | LaneBitmask PrevMask = LiveMask; |
| 638 | LiveMask |= U.LaneMask; |
| 639 | CurPressure.inc(Reg: U.VRegOrUnit.asVirtualReg(), PrevMask, NewMask: LiveMask, MRI: *MRI); |
| 640 | } |
| 641 | |
| 642 | // Update MaxPressure with uses plus early-clobber defs pressure. |
| 643 | MaxPressure = HasECDefs ? max(P1: CurPressure + ECDefPressure, P2: MaxPressure) |
| 644 | : max(P1: CurPressure, P2: MaxPressure); |
| 645 | |
| 646 | assert(CurPressure == getRegPressure(*MRI, LiveRegs)); |
| 647 | } |
| 648 | |
| 649 | //////////////////////////////////////////////////////////////////////////////// |
| 650 | // GCNDownwardRPTracker |
| 651 | |
| 652 | bool GCNDownwardRPTracker::reset(const MachineInstr &MI, |
| 653 | MachineBasicBlock::const_iterator End, |
| 654 | const LiveRegSet *LiveRegsCopy) { |
| 655 | MBBEnd = MI.getParent()->end(); |
| 656 | assert((End == MBBEnd || End->getParent()->end() == MBBEnd) && |
| 657 | "end unrelated to MI block" ); |
| 658 | NextMI = &MI; |
| 659 | NextMI = skipDebugInstructionsForward(It: NextMI, End); |
| 660 | |
| 661 | // Do not use the MI to compute live registers when a set is provided. |
| 662 | // Otherwise the first non-debug instruction after the provided one (or the |
| 663 | // end of the block, if no such instruction exists) serves as the basis to |
| 664 | // compute a live register set. |
| 665 | if (LiveRegsCopy) |
| 666 | GCNRPTracker::reset(MRI: MI.getMF()->getRegInfo(), LiveRegs: *LiveRegsCopy); |
| 667 | else if (NextMI != MBBEnd) |
| 668 | GCNRPTracker::reset(MI: *NextMI, /*After=*/false); |
| 669 | else |
| 670 | GCNRPTracker::reset(MBB: *MI.getParent(), /*End=*/true); |
| 671 | return NextMI != End; |
| 672 | } |
| 673 | |
| 674 | void GCNDownwardRPTracker::retireVirtReg(Register Reg, SlotIndex SI) { |
| 675 | const LiveInterval &LI = LIS.getInterval(Reg); |
| 676 | if (LI.hasSubRanges()) { |
| 677 | auto It = LiveRegs.end(); |
| 678 | for (const auto &S : LI.subranges()) { |
| 679 | if (!S.liveAt(index: SI)) { |
| 680 | if (It == LiveRegs.end()) { |
| 681 | It = LiveRegs.find(Val: Reg); |
| 682 | if (It == LiveRegs.end()) |
| 683 | llvm_unreachable("register isn't live" ); |
| 684 | } |
| 685 | auto PrevMask = It->second; |
| 686 | It->second &= ~S.LaneMask; |
| 687 | CurPressure.inc(Reg, PrevMask, NewMask: It->second, MRI: *MRI); |
| 688 | } |
| 689 | } |
| 690 | if (It != LiveRegs.end() && It->second.none()) |
| 691 | LiveRegs.erase(I: It); |
| 692 | } else if (!LI.liveAt(index: SI)) { |
| 693 | auto It = LiveRegs.find(Val: Reg); |
| 694 | if (It == LiveRegs.end()) |
| 695 | llvm_unreachable("register isn't live" ); |
| 696 | CurPressure.inc(Reg, PrevMask: It->second, NewMask: LaneBitmask::getNone(), MRI: *MRI); |
| 697 | LiveRegs.erase(I: It); |
| 698 | } |
| 699 | } |
| 700 | |
| 701 | bool GCNDownwardRPTracker::advanceBeforeNext(MachineInstr *MI, |
| 702 | bool UseInternalIterator) { |
| 703 | assert(MRI && "call reset first" ); |
| 704 | SlotIndex SI; |
| 705 | const MachineInstr *CurrMI; |
| 706 | if (UseInternalIterator) { |
| 707 | if (!LastTrackedMI) |
| 708 | return NextMI == MBBEnd; |
| 709 | |
| 710 | assert(NextMI == MBBEnd || !NextMI->isDebugInstr()); |
| 711 | CurrMI = LastTrackedMI; |
| 712 | |
| 713 | SI = NextMI == MBBEnd |
| 714 | ? LIS.getInstructionIndex(Instr: *LastTrackedMI).getDeadSlot() |
| 715 | : LIS.getInstructionIndex(Instr: *NextMI).getBaseIndex(); |
| 716 | } else { //! UseInternalIterator |
| 717 | SI = LIS.getInstructionIndex(Instr: *MI).getBaseIndex(); |
| 718 | CurrMI = MI; |
| 719 | } |
| 720 | |
| 721 | assert(SI.isValid()); |
| 722 | |
| 723 | // Remove dead registers or mask bits. |
| 724 | SmallSet<Register, 8> SeenRegs; |
| 725 | for (auto &MO : CurrMI->operands()) { |
| 726 | if (!MO.isReg() || !MO.getReg().isVirtual()) |
| 727 | continue; |
| 728 | if (MO.isUse() && CurrMI->getOpcode() == AMDGPU::PHI) |
| 729 | break; |
| 730 | if (MO.isUse() && !MO.readsReg()) |
| 731 | continue; |
| 732 | if (!UseInternalIterator && MO.isDef()) |
| 733 | continue; |
| 734 | if (!SeenRegs.insert(V: MO.getReg()).second) |
| 735 | continue; |
| 736 | retireVirtReg(Reg: MO.getReg(), SI); |
| 737 | } |
| 738 | |
| 739 | MaxPressure = max(P1: MaxPressure, P2: CurPressure); |
| 740 | |
| 741 | LastTrackedMI = nullptr; |
| 742 | |
| 743 | return UseInternalIterator && (NextMI == MBBEnd); |
| 744 | } |
| 745 | |
| 746 | void GCNDownwardRPTracker::advanceToNext(MachineInstr *MI, |
| 747 | bool UseInternalIterator) { |
| 748 | if (UseInternalIterator) { |
| 749 | LastTrackedMI = &*NextMI++; |
| 750 | NextMI = skipDebugInstructionsForward(It: NextMI, End: MBBEnd); |
| 751 | } else { |
| 752 | LastTrackedMI = MI; |
| 753 | } |
| 754 | |
| 755 | const MachineInstr *CurrMI = LastTrackedMI; |
| 756 | |
| 757 | // Add new registers or mask bits. |
| 758 | for (const auto &MO : CurrMI->all_defs()) { |
| 759 | Register Reg = MO.getReg(); |
| 760 | if (!Reg.isVirtual()) |
| 761 | continue; |
| 762 | auto &LiveMask = LiveRegs[Reg]; |
| 763 | auto PrevMask = LiveMask; |
| 764 | LiveMask |= getDefRegMask(MO, MRI: *MRI); |
| 765 | CurPressure.inc(Reg, PrevMask, NewMask: LiveMask, MRI: *MRI); |
| 766 | } |
| 767 | |
| 768 | MaxPressure = max(P1: MaxPressure, P2: CurPressure); |
| 769 | } |
| 770 | |
| 771 | bool GCNDownwardRPTracker::advance(MachineInstr *MI, bool UseInternalIterator) { |
| 772 | if (UseInternalIterator && NextMI == MBBEnd) |
| 773 | return false; |
| 774 | |
| 775 | advanceBeforeNext(MI, UseInternalIterator); |
| 776 | advanceToNext(MI, UseInternalIterator); |
| 777 | if (!UseInternalIterator) { |
| 778 | const MachineInstr *SavedLastTrackedMI = LastTrackedMI; |
| 779 | // We must remove any dead def lanes from the current RP |
| 780 | advanceBeforeNext(MI, UseInternalIterator: true); |
| 781 | // Restore LastTrackedMI set by advanceToNext, otherwise |
| 782 | // speculative queries (bumpDownwardPressure) don't |
| 783 | // know the last scheduled instruction and fail to |
| 784 | // correctly estimate pressure change. |
| 785 | LastTrackedMI = SavedLastTrackedMI; |
| 786 | } |
| 787 | return true; |
| 788 | } |
| 789 | |
| 790 | bool GCNDownwardRPTracker::advance(MachineBasicBlock::const_iterator End) { |
| 791 | bool AnyAdvance = false; |
| 792 | while (NextMI != End && advance()) |
| 793 | AnyAdvance = true; |
| 794 | return AnyAdvance; |
| 795 | } |
| 796 | |
| 797 | bool GCNDownwardRPTracker::advance(MachineBasicBlock::const_iterator Begin, |
| 798 | MachineBasicBlock::const_iterator End, |
| 799 | const LiveRegSet *LiveRegsCopy) { |
| 800 | if (!reset(MI: *Begin, End, LiveRegsCopy)) |
| 801 | return false; |
| 802 | return advance(End); |
| 803 | } |
| 804 | |
| 805 | Printable llvm::reportMismatch(const GCNRPTracker::LiveRegSet &LISLR, |
| 806 | const GCNRPTracker::LiveRegSet &TrackedLR, |
| 807 | const TargetRegisterInfo *TRI, StringRef Pfx) { |
| 808 | return Printable([&LISLR, &TrackedLR, TRI, Pfx](raw_ostream &OS) { |
| 809 | for (auto const &P : TrackedLR) { |
| 810 | auto I = LISLR.find(Val: P.first); |
| 811 | if (I == LISLR.end()) { |
| 812 | OS << Pfx << printReg(Reg: P.first, TRI) << ":L" << PrintLaneMask(LaneMask: P.second) |
| 813 | << " isn't found in LIS reported set\n" ; |
| 814 | } else if (I->second != P.second) { |
| 815 | OS << Pfx << printReg(Reg: P.first, TRI) |
| 816 | << " masks doesn't match: LIS reported " << PrintLaneMask(LaneMask: I->second) |
| 817 | << ", tracked " << PrintLaneMask(LaneMask: P.second) << '\n'; |
| 818 | } |
| 819 | } |
| 820 | for (auto const &P : LISLR) { |
| 821 | auto I = TrackedLR.find(Val: P.first); |
| 822 | if (I == TrackedLR.end()) { |
| 823 | OS << Pfx << printReg(Reg: P.first, TRI) << ":L" << PrintLaneMask(LaneMask: P.second) |
| 824 | << " isn't found in tracked set\n" ; |
| 825 | } |
| 826 | } |
| 827 | }); |
| 828 | } |
| 829 | |
| 830 | GCNRegPressure |
| 831 | GCNDownwardRPTracker::bumpDownwardPressure(const MachineInstr *MI, |
| 832 | const SIRegisterInfo *TRI) const { |
| 833 | assert(!MI->isDebugOrPseudoInstr() && "Expect a nondebug instruction." ); |
| 834 | |
| 835 | SlotIndex SlotIdx; |
| 836 | SlotIdx = LIS.getInstructionIndex(Instr: *MI).getRegSlot(); |
| 837 | |
| 838 | SlotIndex CurrIdx; |
| 839 | const MachineBasicBlock *MBB = MI->getParent(); |
| 840 | MachineBasicBlock::const_iterator StartPos = |
| 841 | LastTrackedMI ? std::next(x: LastTrackedMI->getIterator()) : MBB->begin(); |
| 842 | MachineBasicBlock::const_iterator IdxPos = |
| 843 | skipDebugInstructionsForward(It: StartPos, End: MBB->end()); |
| 844 | if (IdxPos == MBB->end()) { |
| 845 | CurrIdx = LIS.getMBBEndIdx(mbb: MBB); |
| 846 | } else { |
| 847 | CurrIdx = LIS.getInstructionIndex(Instr: *IdxPos).getRegSlot(); |
| 848 | } |
| 849 | |
| 850 | // Account for register pressure similar to RegPressureTracker::recede(). |
| 851 | RegisterOperands RegOpers; |
| 852 | RegOpers.collect(MI: *MI, TRI: *TRI, MRI: *MRI, TrackLaneMasks: true, /*IgnoreDead=*/false); |
| 853 | RegOpers.adjustLaneLiveness(LIS, MRI: *MRI, Pos: SlotIdx); |
| 854 | GCNRegPressure TempPressure = CurPressure; |
| 855 | // Tracks the live mask reported by the use loop for redefined registers. |
| 856 | SmallDenseMap<Register, LaneBitmask, 8> PostUseMask; |
| 857 | |
| 858 | for (const VRegMaskOrUnit &Use : RegOpers.Uses) { |
| 859 | if (!Use.VRegOrUnit.isVirtualReg()) |
| 860 | continue; |
| 861 | Register Reg = Use.VRegOrUnit.asVirtualReg(); |
| 862 | LaneBitmask LastUseMask = getLastUsedLanes(Reg, Pos: SlotIdx); |
| 863 | if (LastUseMask.none()) |
| 864 | continue; |
| 865 | // The LastUseMask is queried from the liveness information of instruction |
| 866 | // which may be further down the schedule. Some lanes may actually not be |
| 867 | // last uses for the current position. |
| 868 | // FIXME: allow the caller to pass in the list of vreg uses that remain |
| 869 | // to be bottom-scheduled to avoid searching uses at each query. |
| 870 | LastUseMask = |
| 871 | findUseBetween(Reg, LastUseMask, PriorUseIdx: CurrIdx, NextUseIdx: SlotIdx, MRI: *MRI, TRI, LIS: &LIS); |
| 872 | if (LastUseMask.none()) |
| 873 | continue; |
| 874 | |
| 875 | auto It = LiveRegs.find(Val: Reg); |
| 876 | LaneBitmask LiveMask = It != LiveRegs.end() ? It->second : LaneBitmask(0); |
| 877 | LaneBitmask NewMask = LiveMask & ~LastUseMask; |
| 878 | PostUseMask[Reg] = NewMask; |
| 879 | TempPressure.inc(Reg, PrevMask: LiveMask, NewMask, MRI: *MRI); |
| 880 | } |
| 881 | |
| 882 | // Generate liveness for defs. |
| 883 | for (const VRegMaskOrUnit &Def : RegOpers.Defs) { |
| 884 | if (!Def.VRegOrUnit.isVirtualReg()) |
| 885 | continue; |
| 886 | Register Reg = Def.VRegOrUnit.asVirtualReg(); |
| 887 | auto PostIt = PostUseMask.find(Val: Reg); |
| 888 | LaneBitmask LiveMask; |
| 889 | if (PostIt != PostUseMask.end()) { |
| 890 | LiveMask = PostIt->second; |
| 891 | } else { |
| 892 | auto It = LiveRegs.find(Val: Reg); |
| 893 | LiveMask = It != LiveRegs.end() ? It->second : LaneBitmask(0); |
| 894 | } |
| 895 | |
| 896 | LaneBitmask NewMask = LiveMask | Def.LaneMask; |
| 897 | TempPressure.inc(Reg, PrevMask: LiveMask, NewMask, MRI: *MRI); |
| 898 | } |
| 899 | |
| 900 | return TempPressure; |
| 901 | } |
| 902 | |
| 903 | bool GCNUpwardRPTracker::isValid() const { |
| 904 | const auto &SI = LIS.getInstructionIndex(Instr: *LastTrackedMI).getBaseIndex(); |
| 905 | const auto LISLR = llvm::getLiveRegs(SI, LIS, MRI: *MRI); |
| 906 | const auto &TrackedLR = LiveRegs; |
| 907 | |
| 908 | if (!isEqual(S1: LISLR, S2: TrackedLR)) { |
| 909 | dbgs() << "\nGCNUpwardRPTracker error: Tracked and" |
| 910 | " LIS reported livesets mismatch:\n" |
| 911 | << print(LiveRegs: LISLR, MRI: *MRI); |
| 912 | reportMismatch(LISLR, TrackedLR, TRI: MRI->getTargetRegisterInfo()); |
| 913 | return false; |
| 914 | } |
| 915 | |
| 916 | auto LISPressure = getRegPressure(MRI: *MRI, LiveRegs: LISLR); |
| 917 | if (LISPressure != CurPressure) { |
| 918 | dbgs() << "GCNUpwardRPTracker error: Pressure sets different\nTracked: " |
| 919 | << print(RP: CurPressure) << "LIS rpt: " << print(RP: LISPressure); |
| 920 | return false; |
| 921 | } |
| 922 | return true; |
| 923 | } |
| 924 | |
| 925 | Printable llvm::print(const GCNRPTracker::LiveRegSet &LiveRegs, |
| 926 | const MachineRegisterInfo &MRI) { |
| 927 | return Printable([&LiveRegs, &MRI](raw_ostream &OS) { |
| 928 | const TargetRegisterInfo *TRI = MRI.getTargetRegisterInfo(); |
| 929 | for (unsigned I = 0, E = MRI.getNumVirtRegs(); I != E; ++I) { |
| 930 | Register Reg = Register::index2VirtReg(Index: I); |
| 931 | auto It = LiveRegs.find(Val: Reg); |
| 932 | if (It != LiveRegs.end() && It->second.any()) |
| 933 | OS << ' ' << printReg(Reg, TRI) << ':' << PrintLaneMask(LaneMask: It->second); |
| 934 | } |
| 935 | OS << '\n'; |
| 936 | }); |
| 937 | } |
| 938 | |
| 939 | void GCNRegPressure::dump() const { dbgs() << print(RP: *this); } |
| 940 | |
| 941 | static cl::opt<bool> UseDownwardTracker( |
| 942 | "amdgpu-print-rp-downward" , |
| 943 | cl::desc("Use GCNDownwardRPTracker for GCNRegPressurePrinter pass" ), |
| 944 | cl::init(Val: false), cl::Hidden); |
| 945 | |
| 946 | char llvm::GCNRegPressurePrinter::ID = 0; |
| 947 | char &llvm::GCNRegPressurePrinterID = GCNRegPressurePrinter::ID; |
| 948 | |
| 949 | INITIALIZE_PASS(GCNRegPressurePrinter, "amdgpu-print-rp" , "" , true, true) |
| 950 | |
| 951 | // Return lanemask of Reg's subregs that are live-through at [Begin, End] and |
| 952 | // are fully covered by Mask. |
| 953 | static LaneBitmask |
| 954 | getRegLiveThroughMask(const MachineRegisterInfo &MRI, const LiveIntervals &LIS, |
| 955 | Register Reg, SlotIndex Begin, SlotIndex End, |
| 956 | LaneBitmask Mask = LaneBitmask::getAll()) { |
| 957 | |
| 958 | auto IsInOneSegment = [Begin, End](const LiveRange &LR) -> bool { |
| 959 | auto *Segment = LR.getSegmentContaining(Idx: Begin); |
| 960 | return Segment && Segment->contains(I: End); |
| 961 | }; |
| 962 | |
| 963 | LaneBitmask LiveThroughMask; |
| 964 | const LiveInterval &LI = LIS.getInterval(Reg); |
| 965 | if (LI.hasSubRanges()) { |
| 966 | for (auto &SR : LI.subranges()) { |
| 967 | if ((SR.LaneMask & Mask) == SR.LaneMask && IsInOneSegment(SR)) |
| 968 | LiveThroughMask |= SR.LaneMask; |
| 969 | } |
| 970 | } else { |
| 971 | LaneBitmask RegMask = MRI.getMaxLaneMaskForVReg(Reg); |
| 972 | if ((RegMask & Mask) == RegMask && IsInOneSegment(LI)) |
| 973 | LiveThroughMask = RegMask; |
| 974 | } |
| 975 | |
| 976 | return LiveThroughMask; |
| 977 | } |
| 978 | |
| 979 | bool GCNRegPressurePrinter::runOnMachineFunction(MachineFunction &MF) { |
| 980 | const MachineRegisterInfo &MRI = MF.getRegInfo(); |
| 981 | const TargetRegisterInfo *TRI = MRI.getTargetRegisterInfo(); |
| 982 | LiveIntervals &LIS = getAnalysis<LiveIntervalsWrapperPass>().getLIS(); |
| 983 | |
| 984 | auto &OS = dbgs(); |
| 985 | |
| 986 | // Leading spaces are important for YAML syntax. |
| 987 | #define PFX " " |
| 988 | |
| 989 | OS << "---\nname: " << MF.getName() << "\nbody: |\n" ; |
| 990 | |
| 991 | auto printRP = [](const GCNRegPressure &RP) { |
| 992 | return Printable([&RP](raw_ostream &OS) { |
| 993 | OS << format(PFX " %-5d" , Vals: RP.getSGPRNum()) |
| 994 | << format(Fmt: " %-5d" , Vals: RP.getVGPRNum(UnifiedVGPRFile: false)); |
| 995 | }); |
| 996 | }; |
| 997 | |
| 998 | auto ReportLISMismatchIfAny = [&](const GCNRPTracker::LiveRegSet &TrackedLR, |
| 999 | const GCNRPTracker::LiveRegSet &LISLR) { |
| 1000 | if (LISLR != TrackedLR) { |
| 1001 | OS << PFX " mis LIS: " << llvm::print(LiveRegs: LISLR, MRI) |
| 1002 | << reportMismatch(LISLR, TrackedLR, TRI, PFX " " ); |
| 1003 | } |
| 1004 | }; |
| 1005 | |
| 1006 | // Register pressure before and at an instruction (in program order). |
| 1007 | SmallVector<std::pair<GCNRegPressure, GCNRegPressure>, 16> RP; |
| 1008 | |
| 1009 | for (auto &MBB : MF) { |
| 1010 | RP.clear(); |
| 1011 | RP.reserve(N: MBB.size()); |
| 1012 | |
| 1013 | OS << PFX; |
| 1014 | MBB.printName(os&: OS); |
| 1015 | OS << ":\n" ; |
| 1016 | |
| 1017 | SlotIndex MBBStartSlot = LIS.getSlotIndexes()->getMBBStartIdx(mbb: &MBB); |
| 1018 | SlotIndex MBBLastSlot = LIS.getSlotIndexes()->getMBBLastIdx(MBB: &MBB); |
| 1019 | |
| 1020 | GCNRPTracker::LiveRegSet LiveIn, LiveOut; |
| 1021 | GCNRegPressure RPAtMBBEnd; |
| 1022 | |
| 1023 | if (UseDownwardTracker) { |
| 1024 | if (MBB.empty()) { |
| 1025 | LiveIn = LiveOut = getLiveRegs(SI: MBBStartSlot, LIS, MRI); |
| 1026 | RPAtMBBEnd = getRegPressure(MRI, LiveRegs&: LiveIn); |
| 1027 | } else { |
| 1028 | GCNDownwardRPTracker RPT(LIS); |
| 1029 | RPT.reset(MI: MBB.front(), End: MBB.end()); |
| 1030 | |
| 1031 | LiveIn = RPT.getLiveRegs(); |
| 1032 | |
| 1033 | while (!RPT.advanceBeforeNext()) { |
| 1034 | GCNRegPressure RPBeforeMI = RPT.getPressure(); |
| 1035 | RPT.resetMaxPressure(); |
| 1036 | RPT.advanceToNext(); |
| 1037 | RP.emplace_back(Args&: RPBeforeMI, Args: RPT.getMaxPressure()); |
| 1038 | } |
| 1039 | |
| 1040 | LiveOut = RPT.getLiveRegs(); |
| 1041 | RPAtMBBEnd = RPT.getPressure(); |
| 1042 | } |
| 1043 | } else { |
| 1044 | GCNUpwardRPTracker RPT(LIS); |
| 1045 | RPT.reset(MRI, SI: MBBLastSlot); |
| 1046 | |
| 1047 | LiveOut = RPT.getLiveRegs(); |
| 1048 | RPAtMBBEnd = RPT.getPressure(); |
| 1049 | |
| 1050 | for (auto &MI : reverse(C&: MBB)) { |
| 1051 | RPT.resetMaxPressure(); |
| 1052 | RPT.recede(MI); |
| 1053 | if (!MI.isDebugInstr()) |
| 1054 | RP.emplace_back(Args: RPT.getPressure(), Args: RPT.getMaxPressure()); |
| 1055 | } |
| 1056 | |
| 1057 | LiveIn = RPT.getLiveRegs(); |
| 1058 | } |
| 1059 | |
| 1060 | OS << PFX " Live-in: " << llvm::print(LiveRegs: LiveIn, MRI); |
| 1061 | if (!UseDownwardTracker) |
| 1062 | ReportLISMismatchIfAny(LiveIn, getLiveRegs(SI: MBBStartSlot, LIS, MRI)); |
| 1063 | |
| 1064 | OS << PFX " SGPR VGPR\n" ; |
| 1065 | int I = 0; |
| 1066 | for (auto &MI : MBB) { |
| 1067 | if (!MI.isDebugInstr()) { |
| 1068 | auto &[RPBeforeInstr, RPAtInstr] = |
| 1069 | RP[UseDownwardTracker ? I : (RP.size() - 1 - I)]; |
| 1070 | ++I; |
| 1071 | OS << printRP(RPBeforeInstr) << '\n' << printRP(RPAtInstr) << " " ; |
| 1072 | } else |
| 1073 | OS << PFX " " ; |
| 1074 | MI.print(OS); |
| 1075 | } |
| 1076 | OS << printRP(RPAtMBBEnd) << '\n'; |
| 1077 | |
| 1078 | OS << PFX " Live-out:" << llvm::print(LiveRegs: LiveOut, MRI); |
| 1079 | if (UseDownwardTracker) |
| 1080 | ReportLISMismatchIfAny(LiveOut, getLiveRegs(SI: MBBLastSlot, LIS, MRI)); |
| 1081 | |
| 1082 | GCNRPTracker::LiveRegSet LiveThrough; |
| 1083 | for (auto [Reg, Mask] : LiveIn) { |
| 1084 | LaneBitmask MaskIntersection = Mask & LiveOut.lookup(Val: Reg); |
| 1085 | if (MaskIntersection.any()) { |
| 1086 | LaneBitmask LTMask = getRegLiveThroughMask( |
| 1087 | MRI, LIS, Reg, Begin: MBBStartSlot, End: MBBLastSlot, Mask: MaskIntersection); |
| 1088 | if (LTMask.any()) |
| 1089 | LiveThrough[Reg] = LTMask; |
| 1090 | } |
| 1091 | } |
| 1092 | OS << PFX " Live-thr:" << llvm::print(LiveRegs: LiveThrough, MRI); |
| 1093 | OS << printRP(getRegPressure(MRI, LiveRegs&: LiveThrough)) << '\n'; |
| 1094 | } |
| 1095 | OS << "...\n" ; |
| 1096 | return false; |
| 1097 | |
| 1098 | #undef PFX |
| 1099 | } |
| 1100 | |
| 1101 | #if !defined(NDEBUG) || defined(LLVM_ENABLE_DUMP) |
| 1102 | LLVM_DUMP_METHOD void llvm::dumpMaxRegPressure(MachineFunction &MF, |
| 1103 | GCNRegPressure::RegKind Kind, |
| 1104 | LiveIntervals &LIS, |
| 1105 | const MachineLoopInfo *MLI) { |
| 1106 | |
| 1107 | const MachineRegisterInfo &MRI = MF.getRegInfo(); |
| 1108 | const TargetRegisterInfo *TRI = MRI.getTargetRegisterInfo(); |
| 1109 | auto &OS = dbgs(); |
| 1110 | const char *RegName = GCNRegPressure::getName(Kind); |
| 1111 | |
| 1112 | unsigned MaxNumRegs = 0; |
| 1113 | const MachineInstr *MaxPressureMI = nullptr; |
| 1114 | GCNUpwardRPTracker RPT(LIS); |
| 1115 | for (const MachineBasicBlock &MBB : MF) { |
| 1116 | RPT.reset(MRI, LIS.getSlotIndexes()->getMBBEndIdx(&MBB).getPrevSlot()); |
| 1117 | for (const MachineInstr &MI : reverse(MBB)) { |
| 1118 | RPT.recede(MI); |
| 1119 | unsigned NumRegs = RPT.getMaxPressure().getNumRegs(Kind); |
| 1120 | if (NumRegs > MaxNumRegs) { |
| 1121 | MaxNumRegs = NumRegs; |
| 1122 | MaxPressureMI = &MI; |
| 1123 | } |
| 1124 | } |
| 1125 | } |
| 1126 | |
| 1127 | SlotIndex MISlot = LIS.getInstructionIndex(*MaxPressureMI); |
| 1128 | |
| 1129 | // Max pressure can occur at either the early-clobber or register slot. |
| 1130 | // Choose the maximum liveset between both slots. This is ugly but this is |
| 1131 | // diagnostic code. |
| 1132 | SlotIndex ECSlot = MISlot.getRegSlot(true); |
| 1133 | SlotIndex RSlot = MISlot.getRegSlot(false); |
| 1134 | GCNRPTracker::LiveRegSet ECLiveSet = getLiveRegs(ECSlot, LIS, MRI, Kind); |
| 1135 | GCNRPTracker::LiveRegSet RLiveSet = getLiveRegs(RSlot, LIS, MRI, Kind); |
| 1136 | unsigned ECNumRegs = getRegPressure(MRI, ECLiveSet).getNumRegs(Kind); |
| 1137 | unsigned RNumRegs = getRegPressure(MRI, RLiveSet).getNumRegs(Kind); |
| 1138 | GCNRPTracker::LiveRegSet *LiveSet = |
| 1139 | ECNumRegs > RNumRegs ? &ECLiveSet : &RLiveSet; |
| 1140 | SlotIndex MaxPressureSlot = ECNumRegs > RNumRegs ? ECSlot : RSlot; |
| 1141 | assert(getRegPressure(MRI, *LiveSet).getNumRegs(Kind) == MaxNumRegs); |
| 1142 | |
| 1143 | // Split live registers into single-def and multi-def sets. |
| 1144 | GCNRegPressure SDefPressure, MDefPressure; |
| 1145 | SmallVector<Register, 16> SDefRegs, MDefRegs; |
| 1146 | for (auto [Reg, LaneMask] : *LiveSet) { |
| 1147 | assert(GCNRegPressure::getRegKind(Reg, MRI) == Kind); |
| 1148 | LiveInterval &LI = LIS.getInterval(Reg); |
| 1149 | if (LI.getNumValNums() == 1 || |
| 1150 | (LI.hasSubRanges() && |
| 1151 | llvm::all_of(LI.subranges(), [](const LiveInterval::SubRange &SR) { |
| 1152 | return SR.getNumValNums() == 1; |
| 1153 | }))) { |
| 1154 | SDefPressure.inc(Reg, LaneBitmask::getNone(), LaneMask, MRI); |
| 1155 | SDefRegs.push_back(Reg); |
| 1156 | } else { |
| 1157 | MDefPressure.inc(Reg, LaneBitmask::getNone(), LaneMask, MRI); |
| 1158 | MDefRegs.push_back(Reg); |
| 1159 | } |
| 1160 | } |
| 1161 | unsigned SDefNumRegs = SDefPressure.getNumRegs(Kind); |
| 1162 | unsigned MDefNumRegs = MDefPressure.getNumRegs(Kind); |
| 1163 | assert(SDefNumRegs + MDefNumRegs == MaxNumRegs); |
| 1164 | |
| 1165 | auto printLoc = [&](const MachineBasicBlock *MBB, SlotIndex SI) { |
| 1166 | return Printable([&, MBB, SI](raw_ostream &OS) { |
| 1167 | OS << SI << ':' << printMBBReference(*MBB); |
| 1168 | if (MLI) |
| 1169 | if (const MachineLoop *ML = MLI->getLoopFor(MBB)) |
| 1170 | OS << " (LoopHdr " << printMBBReference(*ML->getHeader()) |
| 1171 | << ", Depth " << ML->getLoopDepth() << ")" ; |
| 1172 | }); |
| 1173 | }; |
| 1174 | |
| 1175 | auto PrintRegInfo = [&](Register Reg, LaneBitmask LiveMask) { |
| 1176 | GCNRegPressure RegPressure; |
| 1177 | RegPressure.inc(Reg, LaneBitmask::getNone(), LiveMask, MRI); |
| 1178 | OS << " " << printReg(Reg, TRI) << ':' |
| 1179 | << TRI->getRegClassName(MRI.getRegClass(Reg)) << ", LiveMask " |
| 1180 | << PrintLaneMask(LiveMask) << " (" << RegPressure.getNumRegs(Kind) << ' ' |
| 1181 | << RegName << "s)\n" ; |
| 1182 | |
| 1183 | // Use std::map to sort def/uses by SlotIndex. |
| 1184 | std::map<SlotIndex, const MachineInstr *> Instrs; |
| 1185 | for (const MachineInstr &MI : MRI.reg_nodbg_instructions(Reg)) { |
| 1186 | Instrs[LIS.getInstructionIndex(MI).getRegSlot()] = &MI; |
| 1187 | } |
| 1188 | |
| 1189 | for (const auto &[SI, MI] : Instrs) { |
| 1190 | OS << " " ; |
| 1191 | if (MI->definesRegister(Reg, TRI)) |
| 1192 | OS << "def " ; |
| 1193 | if (MI->readsRegister(Reg, TRI)) |
| 1194 | OS << "use " ; |
| 1195 | OS << printLoc(MI->getParent(), SI) << ": " << *MI; |
| 1196 | } |
| 1197 | }; |
| 1198 | |
| 1199 | OS << "\n*** Register pressure info (" << RegName << "s) for " << MF.getName() |
| 1200 | << " ***\n" ; |
| 1201 | OS << "Max pressure is " << MaxNumRegs << ' ' << RegName << "s at " |
| 1202 | << printLoc(MaxPressureMI->getParent(), MaxPressureSlot) << ": " |
| 1203 | << *MaxPressureMI; |
| 1204 | |
| 1205 | OS << "\nLive registers with single definition (" << SDefNumRegs << ' ' |
| 1206 | << RegName << "s):\n" ; |
| 1207 | |
| 1208 | // Sort SDefRegs by number of uses (smallest first) |
| 1209 | llvm::sort(SDefRegs, [&](Register A, Register B) { |
| 1210 | return std::distance(MRI.use_nodbg_begin(A), MRI.use_nodbg_end()) < |
| 1211 | std::distance(MRI.use_nodbg_begin(B), MRI.use_nodbg_end()); |
| 1212 | }); |
| 1213 | |
| 1214 | for (const Register Reg : SDefRegs) { |
| 1215 | PrintRegInfo(Reg, LiveSet->lookup(Reg)); |
| 1216 | } |
| 1217 | |
| 1218 | OS << "\nLive registers with multiple definitions (" << MDefNumRegs << ' ' |
| 1219 | << RegName << "s):\n" ; |
| 1220 | for (const Register Reg : MDefRegs) { |
| 1221 | PrintRegInfo(Reg, LiveSet->lookup(Reg)); |
| 1222 | } |
| 1223 | } |
| 1224 | #endif |
| 1225 | |
| 1226 | unsigned llvm::estimateGreedyVGPRPressure( |
| 1227 | MachineBasicBlock::const_iterator RegionBegin, |
| 1228 | MachineBasicBlock::const_iterator RegionEnd, |
| 1229 | const GCNRPTracker::LiveRegSet &LiveIns, const LiveIntervals &LIS, |
| 1230 | const MachineRegisterInfo &MRI, const SIRegisterInfo &TRI) { |
| 1231 | |
| 1232 | SetVector<const LiveInterval *> IntervalSet; |
| 1233 | IntervalSet.reserve(Size: LiveIns.size()); |
| 1234 | |
| 1235 | auto checkAndCollect = [&](Register VReg) { |
| 1236 | if (!VReg.isVirtual() || !LIS.hasInterval(Reg: VReg)) |
| 1237 | return; |
| 1238 | |
| 1239 | const TargetRegisterClass *RC = MRI.getRegClass(Reg: VReg); |
| 1240 | if (!TRI.hasVGPRs(RC)) |
| 1241 | return; |
| 1242 | |
| 1243 | const LiveInterval &LI = LIS.getInterval(Reg: VReg); |
| 1244 | IntervalSet.insert(X: &LI); |
| 1245 | }; |
| 1246 | |
| 1247 | // Collect live-ins. |
| 1248 | for (const auto &[RegNum, LaneMask] : LiveIns) { |
| 1249 | checkAndCollect(Register(RegNum)); |
| 1250 | } |
| 1251 | |
| 1252 | // Collect defs in region. |
| 1253 | for (MachineBasicBlock::const_iterator I = RegionBegin; I != RegionEnd; ++I) { |
| 1254 | for (const MachineOperand &MO : I->operands()) { |
| 1255 | if (!MO.isReg() || !MO.isDef()) |
| 1256 | continue; |
| 1257 | checkAndCollect(MO.getReg()); |
| 1258 | } |
| 1259 | } |
| 1260 | |
| 1261 | SmallVector<const LiveInterval *> Intervals = IntervalSet.takeVector(); |
| 1262 | llvm::sort(C&: Intervals, Comp: [](const LiveInterval *LHS, const LiveInterval *RHS) { |
| 1263 | return LHS->beginIndex() < RHS->beginIndex(); |
| 1264 | }); |
| 1265 | |
| 1266 | LiveIntervalUnion::Allocator Alloc; |
| 1267 | std::vector<LiveIntervalUnion> RegFile; |
| 1268 | unsigned MaxRegsUsed = 0; |
| 1269 | |
| 1270 | // Simulate greedy register allocation, assuming an unlimited number of |
| 1271 | // physical registers. |
| 1272 | for (const LiveInterval *LI : Intervals) { |
| 1273 | const TargetRegisterClass *RC = MRI.getRegClass(Reg: LI->reg()); |
| 1274 | unsigned Width = |
| 1275 | std::max<unsigned>(a: 1, b: TRI.getRegSizeInBits(RC: *RC).getFixedValue() / 32); |
| 1276 | unsigned Alignment = |
| 1277 | std::max<unsigned>(a: 1, b: TRI.getRegClassAlignmentNumBits(RC) / 32); |
| 1278 | |
| 1279 | unsigned Start = 0; |
| 1280 | while (true) { |
| 1281 | unsigned End = Start + Width; |
| 1282 | if (RegFile.size() < End) |
| 1283 | RegFile.resize(new_size: End, x: LiveIntervalUnion(Alloc)); |
| 1284 | |
| 1285 | bool Fits = true; |
| 1286 | for (unsigned Idx = Start; Idx < End; Idx++) { |
| 1287 | LiveIntervalUnion::Query Q(*LI, RegFile[Idx]); |
| 1288 | if (Q.checkInterference()) { |
| 1289 | Start = alignTo(Value: Idx + 1, Align: Alignment); |
| 1290 | Fits = false; |
| 1291 | break; |
| 1292 | } |
| 1293 | } |
| 1294 | |
| 1295 | if (Fits) { |
| 1296 | for (unsigned Idx = Start; Idx < End; Idx++) |
| 1297 | RegFile[Idx].unify(VirtReg: *LI, Range: *LI); |
| 1298 | MaxRegsUsed = std::max(a: MaxRegsUsed, b: End); |
| 1299 | break; |
| 1300 | } |
| 1301 | } |
| 1302 | } |
| 1303 | |
| 1304 | return MaxRegsUsed; |
| 1305 | } |
| 1306 | |