LLVM 24.0.0git
MachineSMEABIPass.cpp
Go to the documentation of this file.
1//===- MachineSMEABIPass.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// This pass implements the SME ABI requirements for ZA state. This includes
10// implementing the lazy (and agnostic) ZA state save schemes around calls.
11//
12//===----------------------------------------------------------------------===//
13//
14// This pass works by collecting instructions that require ZA to be in a
15// specific state (e.g., "ACTIVE" or "SAVED") and inserting the necessary state
16// transitions to ensure ZA is in the required state before instructions. State
17// transitions represent actions such as setting up or restoring a lazy save.
18// Certain points within a function may also have predefined states independent
19// of any instructions, for example, a "shared_za" function is always entered
20// and exited in the "ACTIVE" state.
21//
22// To handle ZA state across control flow, we make use of edge bundling. This
23// assigns each block an "incoming" and "outgoing" edge bundle (representing
24// incoming and outgoing edges). Initially, these are unique to each block;
25// then, in the process of forming bundles, the outgoing bundle of a block is
26// joined with the incoming bundle of all successors. The result is that each
27// bundle can be assigned a single ZA state, which ensures the state required by
28// all a blocks' successors is the same, and that each basic block will always
29// be entered with the same ZA state. This eliminates the need for splitting
30// edges to insert state transitions or "phi" nodes for ZA states.
31//
32// See below for a simple example of edge bundling.
33//
34// The following shows a conditionally executed basic block (BB1):
35//
36// if (cond)
37// BB1
38// BB2
39//
40// Initial Bundles Joined Bundles
41//
42// ┌──0──┐ ┌──0──┐
43// │ BB0 │ │ BB0 │
44// └──1──┘ └──1──┘
45// ├───────┐ ├───────┐
46// ▼ │ ▼ │
47// ┌──2──┐ │ ─────► ┌──1──┐ │
48// │ BB1 │ ▼ │ BB1 │ ▼
49// └──3──┘ ┌──4──┐ └──1──┘ ┌──1──┐
50// └───►4 BB2 │ └───►1 BB2 │
51// └──5──┘ └──2──┘
52//
53// On the left are the initial per-block bundles, and on the right are the
54// joined bundles (which are the result of the EdgeBundles analysis).
55
56#include "AArch64InstrInfo.h"
58#include "AArch64Subtarget.h"
69
70using namespace llvm;
71
72#define DEBUG_TYPE "aarch64-machine-sme-abi"
73
74namespace {
75
76// Note: For agnostic ZA, we assume the function is always entered/exited in the
77// "ACTIVE" state -- this _may_ not be the case (since OFF is also a
78// possibility, but for the purpose of placing ZA saves/restores, that does not
79// matter).
80enum ZAState : uint8_t {
81 // Any/unknown state (not valid)
82 ANY = 0,
83
84 // ZA is in use and active (i.e. within the accumulator)
85 ACTIVE,
86
87 // ZA is active, but ZT0 has been saved.
88 // This handles the edge case of sharedZA && !sharesZT0.
89 ACTIVE_ZT0_SAVED,
90
91 // A ZA save has been set up or committed (i.e. ZA is dormant or off)
92 // If the function uses ZT0 it must also be saved.
93 LOCAL_SAVED,
94
95 // ZA has been committed to the lazy save buffer of the current function.
96 // If the function uses ZT0 it must also be saved.
97 // ZA is off.
98 LOCAL_COMMITTED,
99
100 // The ZA/ZT0 state on entry to the function.
101 ENTRY,
102
103 // ZA is off.
104 OFF,
105
106 // The number of ZA states (not a valid state)
107 NUM_ZA_STATE
108};
109
110/// A bitmask enum to record live physical registers that the "emit*" routines
111/// may need to preserve. Note: This only tracks registers we may clobber.
112enum LiveRegs : uint8_t {
113 None = 0,
114 NZCV = 1 << 0,
115 W0 = 1 << 1,
116 W0_HI = 1 << 2,
117 X0 = W0 | W0_HI,
118 LLVM_MARK_AS_BITMASK_ENUM(/* LargestValue = */ W0_HI)
119};
120
121/// Holds the virtual registers live physical registers have been saved to.
122struct PhysRegSave {
123 LiveRegs PhysLiveRegs;
124 Register StatusFlags = Register();
125 Register X0Save = Register();
126};
127
128/// Contains the needed ZA state (and live registers) at an instruction. That is
129/// the state ZA must be in _before_ "InsertPt".
130struct InstInfo {
131 ZAState NeededState{ZAState::ANY};
133 LiveRegs PhysLiveRegs = LiveRegs::None;
134};
135
136/// Contains the needed ZA state for each instruction in a block. Instructions
137/// that do not require a ZA state are not recorded.
138struct BlockInfo {
140 ZAState FixedEntryState{ZAState::ANY};
141 ZAState DesiredIncomingState{ZAState::ANY};
142 ZAState DesiredOutgoingState{ZAState::ANY};
143 LiveRegs PhysLiveRegsAtEntry = LiveRegs::None;
144 LiveRegs PhysLiveRegsAtExit = LiveRegs::None;
145};
146
147/// Contains the needed ZA state information for all blocks within a function.
148struct FunctionInfo {
150 std::optional<MachineBasicBlock::iterator> AfterSMEProloguePt;
151 LiveRegs PhysLiveRegsAfterSMEPrologue = LiveRegs::None;
152};
153
154/// State/helpers that is only needed when emitting code to handle
155/// saving/restoring ZA.
156class EmitContext {
157public:
158 EmitContext() = default;
159
160 /// Get or create a TPIDR2 block in \p MF.
161 int getTPIDR2Block(MachineFunction &MF) {
162 if (TPIDR2BlockFI)
163 return *TPIDR2BlockFI;
164 MachineFrameInfo &MFI = MF.getFrameInfo();
165 TPIDR2BlockFI = MFI.CreateStackObject(16, Align(16), false);
166 return *TPIDR2BlockFI;
167 }
168
169 /// Get or create agnostic ZA buffer pointer in \p MF.
170 Register getAgnosticZABufferPtr(MachineFunction &MF) {
171 if (AgnosticZABufferPtr.isValid())
172 return AgnosticZABufferPtr;
173 Register BufferPtr =
174 MF.getInfo<AArch64FunctionInfo>()->getEarlyAllocSMESaveBuffer();
175 AgnosticZABufferPtr =
176 BufferPtr.isValid()
177 ? BufferPtr
178 : MF.getRegInfo().createVirtualRegister(&AArch64::GPR64RegClass);
179 return AgnosticZABufferPtr;
180 }
181
182 int getZT0SaveSlot(MachineFunction &MF) {
183 if (ZT0SaveFI)
184 return *ZT0SaveFI;
185 MachineFrameInfo &MFI = MF.getFrameInfo();
186 ZT0SaveFI = MFI.CreateSpillStackObject(64, Align(16));
187 return *ZT0SaveFI;
188 }
189
190 /// Returns true if the function must allocate a ZA save buffer on entry. This
191 /// will be the case if, at any point in the function, a ZA save was emitted.
192 bool needsSaveBuffer() const {
193 assert(!(TPIDR2BlockFI && AgnosticZABufferPtr) &&
194 "Cannot have both a TPIDR2 block and agnostic ZA buffer");
195 return TPIDR2BlockFI || AgnosticZABufferPtr.isValid();
196 }
197
198private:
199 std::optional<int> ZT0SaveFI;
200 std::optional<int> TPIDR2BlockFI;
201 Register AgnosticZABufferPtr = Register();
202};
203
204StringRef getZAStateString(ZAState State) {
205#define MAKE_CASE(V) \
206 case V: \
207 return #V;
208 switch (State) {
209 MAKE_CASE(ZAState::ANY)
210 MAKE_CASE(ZAState::ACTIVE)
211 MAKE_CASE(ZAState::ACTIVE_ZT0_SAVED)
212 MAKE_CASE(ZAState::LOCAL_SAVED)
213 MAKE_CASE(ZAState::LOCAL_COMMITTED)
214 MAKE_CASE(ZAState::ENTRY)
215 MAKE_CASE(ZAState::OFF)
216 default:
217 llvm_unreachable("Unexpected ZAState");
218 }
219#undef MAKE_CASE
220}
221
222static bool isZAorZTRegOp(const TargetRegisterInfo &TRI,
223 const MachineOperand &MO) {
224 if (!MO.isReg() || !MO.getReg().isPhysical())
225 return false;
226 return any_of(TRI.subregs_inclusive(MO.getReg()), [](const MCPhysReg &SR) {
227 return AArch64::MPR128RegClass.contains(SR) ||
228 AArch64::ZTRRegClass.contains(SR);
229 });
230}
231
232/// Returns the required ZA state needed before \p MI and an iterator pointing
233/// to where any code required to change the ZA state should be inserted.
234static std::pair<ZAState, MachineBasicBlock::iterator>
235getInstNeededZAState(const TargetRegisterInfo &TRI, MachineInstr &MI,
236 SMEAttrs SMEFnAttrs) {
238
239 // Note: InOutZAUsePseudo, RequiresZASavePseudo, and RequiresZT0SavePseudo are
240 // intended to mark the position immediately before a call. Due to
241 // SelectionDAG constraints, these markers occur after the ADJCALLSTACKDOWN,
242 // so we use std::prev(InsertPt) to get the position before the call.
243
244 if (MI.getOpcode() == AArch64::InOutZAUsePseudo)
245 return {ZAState::ACTIVE, std::prev(InsertPt)};
246
247 // Note: If we need to save both ZA and ZT0 we use RequiresZASavePseudo.
248 if (MI.getOpcode() == AArch64::RequiresZASavePseudo)
249 return {ZAState::LOCAL_SAVED, std::prev(InsertPt)};
250
251 // If we only need to save ZT0 there's two cases to consider:
252 // 1. The function has ZA state (that we don't need to save).
253 // - In this case we switch to the "ACTIVE_ZT0_SAVED" state.
254 // This only saves ZT0.
255 // 2. The function does not have ZA state
256 // - In this case we switch to "LOCAL_COMMITTED" state.
257 // This saves ZT0 and turns ZA off.
258 if (MI.getOpcode() == AArch64::RequiresZT0SavePseudo) {
259 return {SMEFnAttrs.hasZAState() ? ZAState::ACTIVE_ZT0_SAVED
260 : ZAState::LOCAL_COMMITTED,
261 std::prev(InsertPt)};
262 }
263
264 if (MI.isReturn()) {
265 bool ZAOffAtReturn = SMEFnAttrs.hasPrivateZAInterface();
266 return {ZAOffAtReturn ? ZAState::OFF : ZAState::ACTIVE, InsertPt};
267 }
268
269 for (auto &MO : MI.operands()) {
270 if (isZAorZTRegOp(TRI, MO))
271 return {ZAState::ACTIVE, InsertPt};
272 }
273
274 return {ZAState::ANY, InsertPt};
275}
276
277struct MachineSMEABI : public MachineFunctionPass {
278 inline static char ID = 0;
279
280 MachineSMEABI(CodeGenOptLevel OptLevel = CodeGenOptLevel::Default)
281 : MachineFunctionPass(ID), OptLevel(OptLevel) {}
282
283 bool runOnMachineFunction(MachineFunction &MF) override;
284
285 StringRef getPassName() const override { return "Machine SME ABI pass"; }
286
287 void getAnalysisUsage(AnalysisUsage &AU) const override {
288 AU.setPreservesCFG();
295 }
296
297 /// Collects the needed ZA state (and live registers) before each instruction
298 /// within the machine function.
299 FunctionInfo collectNeededZAStates(SMEAttrs SMEFnAttrs);
300
301 /// Assigns each edge bundle a ZA state based on the desired states of
302 /// incoming and outgoing blocks in the bundle.
303 SmallVector<ZAState> assignBundleZAStates(const EdgeBundles &Bundles,
304 const FunctionInfo &FnInfo);
305
306 /// Inserts code to handle changes between ZA states within the function.
307 /// E.g., ACTIVE -> LOCAL_SAVED will insert code required to save ZA.
308 void insertStateChanges(EmitContext &, const FunctionInfo &FnInfo,
309 const EdgeBundles &Bundles,
310 ArrayRef<ZAState> BundleStates);
311
312 void addSMELibCall(MachineInstrBuilder &MIB, RTLIB::Libcall LC,
313 CallingConv::ID ExpectedCC);
314
315 void emitZT0SaveRestore(EmitContext &, MachineBasicBlock &MBB,
316 MachineBasicBlock::iterator MBBI, bool IsSave);
317
318 // Emission routines for private and shared ZA functions (using lazy saves).
319 void emitSMEPrologue(MachineBasicBlock &MBB,
321 void emitRestoreLazySave(EmitContext &, MachineBasicBlock &MBB,
323 LiveRegs PhysLiveRegs);
324 void emitSetupLazySave(EmitContext &, MachineBasicBlock &MBB,
326 void emitAllocateLazySaveBuffer(EmitContext &, MachineBasicBlock &MBB,
329 bool ClearTPIDR2, bool On);
330
331 // Emission routines for agnostic ZA functions.
332 void emitSetupFullZASave(MachineBasicBlock &MBB,
334 LiveRegs PhysLiveRegs);
335 // Emit a "full" ZA save or restore. It is "full" in the sense that this
336 // function will emit a call to __arm_sme_save or __arm_sme_restore, which
337 // handles saving and restoring both ZA and ZT0.
338 void emitFullZASaveRestore(EmitContext &, MachineBasicBlock &MBB,
340 LiveRegs PhysLiveRegs, bool IsSave);
341 void emitAllocateFullZASaveBuffer(EmitContext &, MachineBasicBlock &MBB,
343 LiveRegs PhysLiveRegs);
344
345 /// Attempts to find an insertion point before \p Inst where the status flags
346 /// are not live. If \p Inst is `Block.Insts.end()` a point before the end of
347 /// the block is found.
348 std::pair<MachineBasicBlock::iterator, LiveRegs>
349 findStateChangeInsertionPoint(MachineBasicBlock &MBB, const BlockInfo &Block,
351 void emitStateChange(EmitContext &, MachineBasicBlock &MBB,
352 MachineBasicBlock::iterator MBBI, ZAState From,
353 ZAState To, LiveRegs PhysLiveRegs);
354
355 // Helpers for switching between lazy/full ZA save/restore routines.
356 void emitZASave(EmitContext &Context, MachineBasicBlock &MBB,
358 if (AFI->getSMEFnAttrs().hasAgnosticZAInterface())
359 return emitFullZASaveRestore(Context, MBB, MBBI, PhysLiveRegs,
360 /*IsSave=*/true);
361 return emitSetupLazySave(Context, MBB, MBBI);
362 }
363 void emitZARestore(EmitContext &Context, MachineBasicBlock &MBB,
365 if (AFI->getSMEFnAttrs().hasAgnosticZAInterface())
366 return emitFullZASaveRestore(Context, MBB, MBBI, PhysLiveRegs,
367 /*IsSave=*/false);
368 return emitRestoreLazySave(Context, MBB, MBBI, PhysLiveRegs);
369 }
370 void emitAllocateZASaveBuffer(EmitContext &Context, MachineBasicBlock &MBB,
372 LiveRegs PhysLiveRegs) {
373 if (AFI->getSMEFnAttrs().hasAgnosticZAInterface())
374 return emitAllocateFullZASaveBuffer(Context, MBB, MBBI, PhysLiveRegs);
375 return emitAllocateLazySaveBuffer(Context, MBB, MBBI);
376 }
377
378 /// Collects the reachable calls from \p MBBI marked with \p Marker. This is
379 /// intended to be used to emit lazy save remarks. Note: This stops at the
380 /// first marked call along any path.
381 void collectReachableMarkedCalls(const MachineBasicBlock &MBB,
384 unsigned Marker) const;
385
386 void emitCallSaveRemarks(const MachineBasicBlock &MBB,
388 unsigned Marker, StringRef RemarkName,
389 StringRef SaveName) const;
390
391 void emitError(const Twine &Message) {
392 LLVMContext &Context = MF->getFunction().getContext();
393 Context.emitError(MF->getName() + ": " + Message);
394 }
395
396 /// Save live physical registers to virtual registers.
397 PhysRegSave createPhysRegSave(LiveRegs PhysLiveRegs, MachineBasicBlock &MBB,
399 /// Restore physical registers from a save of their previous values.
400 void restorePhyRegSave(const PhysRegSave &RegSave, MachineBasicBlock &MBB,
402
403private:
405
406 MachineFunction *MF = nullptr;
407 const AArch64Subtarget *Subtarget = nullptr;
408 const AArch64RegisterInfo *TRI = nullptr;
409 const AArch64FunctionInfo *AFI = nullptr;
410 const AArch64InstrInfo *TII = nullptr;
411 const LibcallLoweringInfo *LLI = nullptr;
412
414 MachineRegisterInfo *MRI = nullptr;
415 MachineLoopInfo *MLI = nullptr;
416};
417
418static LiveRegs getPhysLiveRegs(LiveRegUnits const &LiveUnits) {
419 LiveRegs PhysLiveRegs = LiveRegs::None;
420 if (!LiveUnits.available(AArch64::NZCV))
421 PhysLiveRegs |= LiveRegs::NZCV;
422 // We have to track W0 and X0 separately as otherwise things can get
423 // confused if we attempt to preserve X0 but only W0 was defined.
424 if (!LiveUnits.available(AArch64::W0))
425 PhysLiveRegs |= LiveRegs::W0;
426 if (!LiveUnits.available(AArch64::W0_HI))
427 PhysLiveRegs |= LiveRegs::W0_HI;
428 return PhysLiveRegs;
429}
430
431static void setPhysLiveRegs(LiveRegUnits &LiveUnits, LiveRegs PhysLiveRegs) {
432 if (PhysLiveRegs & LiveRegs::NZCV)
433 LiveUnits.addReg(AArch64::NZCV);
434 if (PhysLiveRegs & LiveRegs::W0)
435 LiveUnits.addReg(AArch64::W0);
436 if (PhysLiveRegs & LiveRegs::W0_HI)
437 LiveUnits.addReg(AArch64::W0_HI);
438}
439
440[[maybe_unused]] bool isCallStartOpcode(unsigned Opc) {
441 switch (Opc) {
442 case AArch64::BLR:
443 case AArch64::BLRA:
444 case AArch64::TLSDESC_CALLSEQ:
445 case AArch64::TLSDESC_AUTH_CALLSEQ:
446 case AArch64::ADJCALLSTACKDOWN:
447 return true;
448 default:
449 return false;
450 }
451}
452
453FunctionInfo MachineSMEABI::collectNeededZAStates(SMEAttrs SMEFnAttrs) {
454 assert((SMEFnAttrs.hasAgnosticZAInterface() || SMEFnAttrs.hasZT0State() ||
455 SMEFnAttrs.hasZAState()) &&
456 "Expected function to have ZA/ZT0 state!");
457
459 LiveRegs PhysLiveRegsAfterSMEPrologue = LiveRegs::None;
460 std::optional<MachineBasicBlock::iterator> AfterSMEProloguePt;
461
462 for (MachineBasicBlock &MBB : *MF) {
463 BlockInfo &Block = Blocks[MBB.getNumber()];
464
465 if (MBB.isEntryBlock()) {
466 // Entry block:
467 Block.FixedEntryState = ZAState::ENTRY;
468 } else if (MBB.isEHPad()) {
469 // EH entry block:
470 Block.FixedEntryState = ZAState::LOCAL_COMMITTED;
471 }
472
473 LiveRegUnits LiveUnits(*TRI);
474 LiveUnits.addLiveOuts(MBB);
475
476 Block.PhysLiveRegsAtExit = getPhysLiveRegs(LiveUnits);
477 auto FirstTerminatorInsertPt = MBB.getFirstTerminator();
478 auto FirstNonPhiInsertPt = MBB.getFirstNonPHI();
479 for (MachineInstr &MI : reverse(MBB)) {
480 if (MI.isDebugInstr())
481 continue;
482
484 LiveUnits.stepBackward(MI);
485 LiveRegs PhysLiveRegs = getPhysLiveRegs(LiveUnits);
486 // The SMEStateAllocPseudo marker is added to a function if the save
487 // buffer was allocated in SelectionDAG. It marks the end of the
488 // allocation -- which is a safe point for this pass to insert any TPIDR2
489 // block setup.
490 if (MI.getOpcode() == AArch64::SMEStateAllocPseudo) {
491 AfterSMEProloguePt = MBBI;
492 PhysLiveRegsAfterSMEPrologue = PhysLiveRegs;
493 }
494 // Note: We treat Agnostic ZA as inout_za with an alternate save/restore.
495 auto [NeededState, InsertPt] = getInstNeededZAState(*TRI, MI, SMEFnAttrs);
496 assert((InsertPt == MBBI || isCallStartOpcode(InsertPt->getOpcode())) &&
497 "Unexpected state change insertion point!");
498 if (MBBI == FirstTerminatorInsertPt)
499 Block.PhysLiveRegsAtExit = PhysLiveRegs;
500 if (MBBI == FirstNonPhiInsertPt)
501 Block.PhysLiveRegsAtEntry = PhysLiveRegs;
502 if (NeededState != ZAState::ANY)
503 Block.Insts.push_back({NeededState, InsertPt, PhysLiveRegs});
504 }
505
506 // Reverse vector (as we had to iterate backwards for liveness).
507 std::reverse(Block.Insts.begin(), Block.Insts.end());
508
509 // Record the desired states on entry/exit of this block. These are the
510 // states that would not incur a state transition.
511 if (!Block.Insts.empty()) {
512 Block.DesiredIncomingState = Block.Insts.front().NeededState;
513 Block.DesiredOutgoingState = Block.Insts.back().NeededState;
514 }
515 }
516
517 return FunctionInfo{std::move(Blocks), AfterSMEProloguePt,
518 PhysLiveRegsAfterSMEPrologue};
519}
520
521/// Assigns each edge bundle a ZA state based on the desired states of incoming
522/// and outgoing blocks in the bundle.
524MachineSMEABI::assignBundleZAStates(const EdgeBundles &Bundles,
525 const FunctionInfo &FnInfo) {
526 SmallVector<ZAState> BundleStates(Bundles.getNumBundles());
527 for (unsigned I = 0, E = Bundles.getNumBundles(); I != E; ++I) {
528 std::optional<ZAState> BundleState;
529 for (unsigned BlockID : Bundles.getBlocks(I)) {
530 const BlockInfo &Block = FnInfo.Blocks[BlockID];
531 // Check if the block is an incoming block in the bundle. Note: We skip
532 // Block.FixedEntryState != ANY to ignore EH pads (which are only
533 // reachable via exceptions).
534 if (Block.FixedEntryState != ZAState::ANY ||
535 Bundles.getBundle(BlockID, /*Out=*/false) != I)
536 continue;
537
538 // Pick a state that matches all incoming blocks. Fall back to "ACTIVE" if
539 // any incoming state doesn't match. This will hoist the state from
540 // incoming blocks to outgoing blocks.
541 if (!BundleState)
542 BundleState = Block.DesiredIncomingState;
543 else if (BundleState != Block.DesiredIncomingState)
544 BundleState = ZAState::ACTIVE;
545 }
546
547 if (!BundleState || BundleState == ZAState::ANY)
548 BundleState = ZAState::ACTIVE;
549
550 BundleStates[I] = *BundleState;
551 }
552
553 return BundleStates;
554}
555
556std::pair<MachineBasicBlock::iterator, LiveRegs>
557MachineSMEABI::findStateChangeInsertionPoint(
558 MachineBasicBlock &MBB, const BlockInfo &Block,
560 LiveRegs PhysLiveRegs;
562 if (Inst != Block.Insts.end()) {
563 InsertPt = Inst->InsertPt;
564 PhysLiveRegs = Inst->PhysLiveRegs;
565 } else {
566 InsertPt = MBB.getFirstTerminator();
567 PhysLiveRegs = Block.PhysLiveRegsAtExit;
568 }
569
570 if (PhysLiveRegs == LiveRegs::None)
571 return {InsertPt, PhysLiveRegs}; // Nothing to do (no live regs).
572
573 // Find the previous state change. We can not move before this point.
574 MachineBasicBlock::iterator PrevStateChangeI;
575 if (Inst == Block.Insts.begin()) {
576 PrevStateChangeI = MBB.begin();
577 } else {
578 // Note: `std::prev(Inst)` is the previous InstInfo. We only create an
579 // InstInfo object for instructions that require a specific ZA state, so the
580 // InstInfo is the site of the previous state change in the block (which can
581 // be several MIs earlier).
582 PrevStateChangeI = std::prev(Inst)->InsertPt;
583 }
584
585 // Note: LiveUnits will only accurately track X0 and NZCV.
586 LiveRegUnits LiveUnits(*TRI);
587 setPhysLiveRegs(LiveUnits, PhysLiveRegs);
588 auto BestCandidate = std::make_pair(InsertPt, PhysLiveRegs);
589 for (MachineBasicBlock::iterator I = InsertPt; I != PrevStateChangeI; --I) {
590 if (I->isDebugInstr())
591 continue;
592
593 // Don't move before/into a call (which may have a state change before it).
594 if (I->getOpcode() == TII->getCallFrameDestroyOpcode() || I->isCall())
595 break;
596 LiveUnits.stepBackward(*I);
597 LiveRegs CurrentPhysLiveRegs = getPhysLiveRegs(LiveUnits);
598 // Find places where NZCV is available, but keep looking for locations where
599 // both NZCV and X0 are available, which can avoid some copies.
600 if (!(CurrentPhysLiveRegs & LiveRegs::NZCV))
601 BestCandidate = {I, CurrentPhysLiveRegs};
602 if (CurrentPhysLiveRegs == LiveRegs::None)
603 break;
604 }
605 return BestCandidate;
606}
607
608void MachineSMEABI::insertStateChanges(EmitContext &Context,
609 const FunctionInfo &FnInfo,
610 const EdgeBundles &Bundles,
611 ArrayRef<ZAState> BundleStates) {
612 for (MachineBasicBlock &MBB : *MF) {
613 const BlockInfo &Block = FnInfo.Blocks[MBB.getNumber()];
614 ZAState InState = BundleStates[Bundles.getBundle(MBB.getNumber(),
615 /*Out=*/false)];
616
617 ZAState CurrentState = Block.FixedEntryState;
618 if (CurrentState == ZAState::ANY)
619 CurrentState = InState;
620
621 for (auto &Inst : Block.Insts) {
622 if (CurrentState != Inst.NeededState) {
623 auto [InsertPt, PhysLiveRegs] =
624 findStateChangeInsertionPoint(MBB, Block, &Inst);
625 emitStateChange(Context, MBB, InsertPt, CurrentState, Inst.NeededState,
626 PhysLiveRegs);
627 CurrentState = Inst.NeededState;
628 }
629 }
630
631 if (MBB.succ_empty())
632 continue;
633
634 ZAState OutState =
635 BundleStates[Bundles.getBundle(MBB.getNumber(), /*Out=*/true)];
636 if (CurrentState != OutState) {
637 auto [InsertPt, PhysLiveRegs] =
638 findStateChangeInsertionPoint(MBB, Block, Block.Insts.end());
639 emitStateChange(Context, MBB, InsertPt, CurrentState, OutState,
640 PhysLiveRegs);
641 }
642 }
643}
644
647 if (MBB.empty())
648 return DebugLoc();
649 return MBBI != MBB.end() ? MBBI->getDebugLoc() : MBB.back().getDebugLoc();
650}
651
652/// Finds the first call (as determined by MachineInstr::isCall()) starting from
653/// \p MBBI in \p MBB marked with \p Marker (which is a marker opcode such as
654/// RequiresZASavePseudo). If a marked call is found, it is pushed to \p Calls
655/// and the function returns true.
656static bool findMarkedCall(const MachineBasicBlock &MBB,
659 unsigned Marker, unsigned CallDestroyOpcode) {
660 auto IsMarker = [&](auto &MI) { return MI.getOpcode() == Marker; };
661 auto MarkerInst = std::find_if(MBBI, MBB.end(), IsMarker);
662 if (MarkerInst == MBB.end())
663 return false;
665 while (++I != MBB.end()) {
666 if (I->isCall() || I->getOpcode() == CallDestroyOpcode)
667 break;
668 }
669 if (I != MBB.end() && I->isCall())
670 Calls.push_back(&*I);
671 // Note: This function always returns true if a "Marker" was found.
672 return true;
673}
674
675void MachineSMEABI::collectReachableMarkedCalls(
676 const MachineBasicBlock &StartMBB,
678 SmallVectorImpl<const MachineInstr *> &Calls, unsigned Marker) const {
679 assert(Marker == AArch64::InOutZAUsePseudo ||
680 Marker == AArch64::RequiresZASavePseudo ||
681 Marker == AArch64::RequiresZT0SavePseudo);
682 unsigned CallDestroyOpcode = TII->getCallFrameDestroyOpcode();
683 if (findMarkedCall(StartMBB, StartInst, Calls, Marker, CallDestroyOpcode))
684 return;
685
688 StartMBB.succ_rend());
689 while (!Worklist.empty()) {
690 const MachineBasicBlock *MBB = Worklist.pop_back_val();
691 auto [_, Inserted] = Visited.insert(MBB);
692 if (!Inserted)
693 continue;
694
695 if (!findMarkedCall(*MBB, MBB->begin(), Calls, Marker, CallDestroyOpcode))
696 Worklist.append(MBB->succ_rbegin(), MBB->succ_rend());
697 }
698}
699
700static StringRef getCalleeName(const MachineInstr &CallInst) {
701 assert(CallInst.isCall() && "expected a call");
702 for (const MachineOperand &MO : CallInst.operands()) {
703 if (MO.isSymbol())
704 return MO.getSymbolName();
705 if (MO.isGlobal())
706 return MO.getGlobal()->getName();
707 }
708 return {};
709}
710
711void MachineSMEABI::emitCallSaveRemarks(const MachineBasicBlock &MBB,
713 DebugLoc DL, unsigned Marker,
714 StringRef RemarkName,
715 StringRef SaveName) const {
716 auto SaveRemark = [&](DebugLoc DL, const MachineBasicBlock &MBB) {
717 return MachineOptimizationRemarkAnalysis("sme", RemarkName, DL, &MBB);
718 };
719 StringRef StateName = Marker == AArch64::RequiresZT0SavePseudo ? "ZT0" : "ZA";
720 ORE->emit([&] {
721 return SaveRemark(DL, MBB) << SaveName << " of " << StateName
722 << " emitted in '" << MF->getName() << "'";
723 });
724 if (!ORE->allowExtraAnalysis("sme"))
725 return;
726 SmallVector<const MachineInstr *> CallsRequiringSaves;
727 collectReachableMarkedCalls(MBB, MBBI, CallsRequiringSaves, Marker);
728 for (const MachineInstr *CallInst : CallsRequiringSaves) {
729 auto R = SaveRemark(CallInst->getDebugLoc(), *CallInst->getParent());
730 R << "call";
731 if (StringRef CalleeName = getCalleeName(*CallInst); !CalleeName.empty())
732 R << " to '" << CalleeName << "'";
733 R << " requires " << StateName << " save";
734 ORE->emit(R);
735 }
736}
737
738void MachineSMEABI::emitSetupLazySave(EmitContext &Context,
742
743 emitCallSaveRemarks(MBB, MBBI, DL, AArch64::RequiresZASavePseudo,
744 "SMELazySaveZA", "lazy save");
745
746 // Get pointer to TPIDR2 block.
747 Register TPIDR2 = MRI->createVirtualRegister(&AArch64::GPR64spRegClass);
748 Register TPIDR2Ptr = MRI->createVirtualRegister(&AArch64::GPR64RegClass);
749 BuildMI(MBB, MBBI, DL, TII->get(AArch64::ADDXri), TPIDR2)
750 .addFrameIndex(Context.getTPIDR2Block(*MF))
751 .addImm(0)
752 .addImm(0);
753 BuildMI(MBB, MBBI, DL, TII->get(TargetOpcode::COPY), TPIDR2Ptr)
754 .addReg(TPIDR2);
755 // Set TPIDR2_EL0 to point to TPIDR2 block.
756 BuildMI(MBB, MBBI, DL, TII->get(AArch64::MSR))
757 .addImm(AArch64SysReg::TPIDR2_EL0)
758 .addReg(TPIDR2Ptr);
759}
760
761PhysRegSave MachineSMEABI::createPhysRegSave(LiveRegs PhysLiveRegs,
764 DebugLoc DL) {
765 PhysRegSave RegSave{PhysLiveRegs};
766 if (PhysLiveRegs & LiveRegs::NZCV) {
767 RegSave.StatusFlags = MRI->createVirtualRegister(&AArch64::GPR64RegClass);
768 BuildMI(MBB, MBBI, DL, TII->get(AArch64::MRS), RegSave.StatusFlags)
769 .addImm(AArch64SysReg::NZCV)
770 .addReg(AArch64::NZCV, RegState::Implicit)
771 .setOperandDead(2); // implicit-def $nzcv
772 }
773 // Note: Preserving X0 is "free" as this is before register allocation, so
774 // the register allocator is still able to optimize these copies.
775 if (PhysLiveRegs & LiveRegs::W0) {
776 RegSave.X0Save = MRI->createVirtualRegister(PhysLiveRegs & LiveRegs::W0_HI
777 ? &AArch64::GPR64RegClass
778 : &AArch64::GPR32RegClass);
779 BuildMI(MBB, MBBI, DL, TII->get(TargetOpcode::COPY), RegSave.X0Save)
780 .addReg(PhysLiveRegs & LiveRegs::W0_HI ? AArch64::X0 : AArch64::W0);
781 }
782 return RegSave;
783}
784
785void MachineSMEABI::restorePhyRegSave(const PhysRegSave &RegSave,
788 DebugLoc DL) {
789 if (RegSave.StatusFlags.isValid())
790 BuildMI(MBB, MBBI, DL, TII->get(AArch64::MSR))
791 .addImm(AArch64SysReg::NZCV)
792 .addReg(RegSave.StatusFlags)
793 .addReg(AArch64::NZCV, RegState::ImplicitDefine);
794
795 if (RegSave.X0Save.isValid())
796 BuildMI(MBB, MBBI, DL, TII->get(TargetOpcode::COPY),
797 RegSave.PhysLiveRegs & LiveRegs::W0_HI ? AArch64::X0 : AArch64::W0)
798 .addReg(RegSave.X0Save);
799}
800
801void MachineSMEABI::addSMELibCall(MachineInstrBuilder &MIB, RTLIB::Libcall LC,
802 CallingConv::ID ExpectedCC) {
803 RTLIB::LibcallImpl LCImpl = LLI->getLibcallImpl(LC);
804 if (LCImpl == RTLIB::Unsupported)
805 emitError("cannot lower SME ABI (SME routines unsupported)");
808 if (CC != ExpectedCC)
809 emitError("invalid calling convention for SME routine: '" + ImplName + "'");
810 // FIXME: This assumes the ImplName StringRef is null-terminated.
811 MIB.addExternalSymbol(ImplName.data());
812 MIB.addRegMask(TRI->getCallPreservedMask(*MF, CC));
813}
814
815void MachineSMEABI::emitRestoreLazySave(EmitContext &Context,
818 LiveRegs PhysLiveRegs) {
820 Register TPIDR2EL0 = MRI->createVirtualRegister(&AArch64::GPR64RegClass);
821 Register TPIDR2 = AArch64::X0;
822
823 // TODO: Emit these within the restore MBB to prevent unnecessary saves.
824 PhysRegSave RegSave = createPhysRegSave(PhysLiveRegs, MBB, MBBI, DL);
825
826 // Enable ZA.
827 BuildMI(MBB, MBBI, DL, TII->get(AArch64::MSRpstatesvcrImm1))
828 .addImm(AArch64SVCR::SVCRZA)
829 .addImm(1);
830 // Get current TPIDR2_EL0.
831 BuildMI(MBB, MBBI, DL, TII->get(AArch64::MRS), TPIDR2EL0)
832 .addImm(AArch64SysReg::TPIDR2_EL0)
833 .setOperandDead(2); // implicit-def $nzcv
834 // Get pointer to TPIDR2 block.
835 BuildMI(MBB, MBBI, DL, TII->get(AArch64::ADDXri), TPIDR2)
836 .addFrameIndex(Context.getTPIDR2Block(*MF))
837 .addImm(0)
838 .addImm(0);
839 // (Conditionally) restore ZA state.
840 auto RestoreZA = BuildMI(MBB, MBBI, DL, TII->get(AArch64::RestoreZAPseudo))
841 .addReg(TPIDR2EL0)
842 .addReg(TPIDR2);
843 addSMELibCall(
844 RestoreZA, RTLIB::SMEABI_TPIDR2_RESTORE,
846 // Zero TPIDR2_EL0.
847 BuildMI(MBB, MBBI, DL, TII->get(AArch64::MSR))
848 .addImm(AArch64SysReg::TPIDR2_EL0)
849 .addReg(AArch64::XZR);
850
851 restorePhyRegSave(RegSave, MBB, MBBI, DL);
852}
853
854void MachineSMEABI::emitZAMode(MachineBasicBlock &MBB,
856 bool ClearTPIDR2, bool On) {
858
859 if (ClearTPIDR2)
860 BuildMI(MBB, MBBI, DL, TII->get(AArch64::MSR))
861 .addImm(AArch64SysReg::TPIDR2_EL0)
862 .addReg(AArch64::XZR);
863
864 // Disable ZA.
865 BuildMI(MBB, MBBI, DL, TII->get(AArch64::MSRpstatesvcrImm1))
866 .addImm(AArch64SVCR::SVCRZA)
867 .addImm(On ? 1 : 0);
868}
869
870void MachineSMEABI::emitAllocateLazySaveBuffer(
871 EmitContext &Context, MachineBasicBlock &MBB,
873 MachineFrameInfo &MFI = MF->getFrameInfo();
875 Register SP = MRI->createVirtualRegister(&AArch64::GPR64RegClass);
876 Register SVL = MRI->createVirtualRegister(&AArch64::GPR64RegClass);
877 Register Buffer = AFI->getEarlyAllocSMESaveBuffer();
878
879 // Calculate SVL.
880 BuildMI(MBB, MBBI, DL, TII->get(AArch64::RDSVLI_XI), SVL).addImm(1);
881
882 // 1. Allocate the lazy save buffer.
883 if (!Buffer.isValid()) {
884 // TODO: On Windows, we allocate the lazy save buffer in SelectionDAG (so
885 // Buffer is valid). This is done to reuse the existing
886 // expansions (which can insert stack checks). This works, but it means we
887 // will always allocate the lazy save buffer (even if the function contains
888 // no lazy saves). If we want to handle Windows here, we'll need to
889 // implement something similar to LowerWindowsDYNAMIC_STACKALLOC.
890 assert(!Subtarget->isTargetWindows() &&
891 "Lazy ZA save is not yet supported on Windows");
892 Buffer = MRI->createVirtualRegister(&AArch64::GPR64RegClass);
893 // Get original stack pointer.
894 BuildMI(MBB, MBBI, DL, TII->get(TargetOpcode::COPY), SP)
895 .addReg(AArch64::SP);
896 // Allocate a lazy-save buffer object of the size given, normally SVL * SVL
897 BuildMI(MBB, MBBI, DL, TII->get(AArch64::MSUBXrrr), Buffer)
898 .addReg(SVL)
899 .addReg(SVL)
900 .addReg(SP);
901 BuildMI(MBB, MBBI, DL, TII->get(TargetOpcode::COPY), AArch64::SP)
902 .addReg(Buffer);
903 // We have just allocated a variable sized object, tell this to PEI.
904 MFI.CreateVariableSizedObject(Align(16), nullptr);
905 }
906
907 // 2. Setup the TPIDR2 block.
908 {
909 // Note: This case just needs to do `SVL << 48`. It is not implemented as we
910 // generally don't support big-endian SVE/SME.
911 if (!Subtarget->isLittleEndian())
913 "TPIDR2 block initialization is not supported on big-endian targets");
914
915 // Store buffer pointer and num_za_save_slices.
916 // Bytes 10-15 are implicitly zeroed.
917 BuildMI(MBB, MBBI, DL, TII->get(AArch64::STPXi))
918 .addReg(Buffer)
919 .addReg(SVL)
920 .addFrameIndex(Context.getTPIDR2Block(*MF))
921 .addImm(0);
922 }
923}
924
925static constexpr unsigned ZERO_ALL_ZA_MASK = 0b11111111;
926
927void MachineSMEABI::emitSMEPrologue(MachineBasicBlock &MBB,
930
931 bool ZeroZA = AFI->getSMEFnAttrs().isNewZA();
932 bool ZeroZT0 = AFI->getSMEFnAttrs().isNewZT0();
934 // Get current TPIDR2_EL0.
935 Register TPIDR2EL0 = MRI->createVirtualRegister(&AArch64::GPR64RegClass);
936 BuildMI(MBB, MBBI, DL, TII->get(AArch64::MRS))
937 .addReg(TPIDR2EL0, RegState::Define)
938 .addImm(AArch64SysReg::TPIDR2_EL0)
939 .setOperandDead(2); // implicit-def $nzcv
940 // If TPIDR2_EL0 is non-zero, commit the lazy save.
941 // NOTE: Functions that only use ZT0 don't need to zero ZA.
942 auto CommitZASave =
943 BuildMI(MBB, MBBI, DL, TII->get(AArch64::CommitZASavePseudo))
944 .addReg(TPIDR2EL0)
945 .addImm(ZeroZA)
946 .addImm(ZeroZT0);
947 addSMELibCall(
948 CommitZASave, RTLIB::SMEABI_TPIDR2_SAVE,
950 if (ZeroZA)
951 CommitZASave.addDef(AArch64::ZAB0, RegState::ImplicitDefine);
952 if (ZeroZT0)
953 CommitZASave.addDef(AArch64::ZT0, RegState::ImplicitDefine);
954 // Enable ZA (as ZA could have previously been in the OFF state).
955 BuildMI(MBB, MBBI, DL, TII->get(AArch64::MSRpstatesvcrImm1))
956 .addImm(AArch64SVCR::SVCRZA)
957 .addImm(1);
958 } else if (AFI->getSMEFnAttrs().hasSharedZAInterface()) {
959 if (ZeroZA)
960 BuildMI(MBB, MBBI, DL, TII->get(AArch64::ZERO_M))
962 .addDef(AArch64::ZAB0, RegState::ImplicitDefine);
963 if (ZeroZT0)
964 BuildMI(MBB, MBBI, DL, TII->get(AArch64::ZERO_T)).addDef(AArch64::ZT0);
965 }
966}
967
968void MachineSMEABI::emitFullZASaveRestore(EmitContext &Context,
971 LiveRegs PhysLiveRegs, bool IsSave) {
973
974 if (IsSave)
975 emitCallSaveRemarks(MBB, MBBI, DL, AArch64::RequiresZASavePseudo,
976 "SMEFullZASave", "full save");
977
978 PhysRegSave RegSave = createPhysRegSave(PhysLiveRegs, MBB, MBBI, DL);
979
980 // Copy the buffer pointer into X0.
981 Register BufferPtr = AArch64::X0;
982 BuildMI(MBB, MBBI, DL, TII->get(TargetOpcode::COPY), BufferPtr)
983 .addReg(Context.getAgnosticZABufferPtr(*MF));
984
985 // Call __arm_sme_save/__arm_sme_restore.
986 auto SaveRestoreZA = BuildMI(MBB, MBBI, DL, TII->get(AArch64::BL))
987 .addReg(BufferPtr, RegState::Implicit);
988 addSMELibCall(
989 SaveRestoreZA,
990 IsSave ? RTLIB::SMEABI_SME_SAVE : RTLIB::SMEABI_SME_RESTORE,
992
993 restorePhyRegSave(RegSave, MBB, MBBI, DL);
994}
995
996void MachineSMEABI::emitZT0SaveRestore(EmitContext &Context,
999 bool IsSave) {
1001
1002 // Note: This will report calls that _only_ need ZT0 saved. Call that save
1003 // both ZA and ZT0 will be under the SMELazySaveZA remark. This prevents
1004 // reporting the same calls twice.
1005 if (IsSave)
1006 emitCallSaveRemarks(MBB, MBBI, DL, AArch64::RequiresZT0SavePseudo,
1007 "SMEZT0Save", "spill");
1008
1009 Register ZT0Save = MRI->createVirtualRegister(&AArch64::GPR64spRegClass);
1010
1011 BuildMI(MBB, MBBI, DL, TII->get(AArch64::ADDXri), ZT0Save)
1012 .addFrameIndex(Context.getZT0SaveSlot(*MF))
1013 .addImm(0)
1014 .addImm(0);
1015
1016 if (IsSave) {
1017 BuildMI(MBB, MBBI, DL, TII->get(AArch64::STR_TX))
1018 .addReg(AArch64::ZT0)
1019 .addReg(ZT0Save);
1020 } else {
1021 BuildMI(MBB, MBBI, DL, TII->get(AArch64::LDR_TX), AArch64::ZT0)
1022 .addReg(ZT0Save);
1023 }
1024}
1025
1026void MachineSMEABI::emitAllocateFullZASaveBuffer(
1027 EmitContext &Context, MachineBasicBlock &MBB,
1029 // Buffer already allocated in SelectionDAG.
1030 if (AFI->getEarlyAllocSMESaveBuffer())
1031 return;
1032
1034 Register BufferPtr = Context.getAgnosticZABufferPtr(*MF);
1035 Register BufferSize = MRI->createVirtualRegister(&AArch64::GPR64RegClass);
1036
1037 PhysRegSave RegSave = createPhysRegSave(PhysLiveRegs, MBB, MBBI, DL);
1038
1039 // Calculate the SME state size.
1040 {
1041 auto SMEStateSize = BuildMI(MBB, MBBI, DL, TII->get(AArch64::BL))
1042 .addReg(AArch64::X0, RegState::ImplicitDefine);
1043 addSMELibCall(
1044 SMEStateSize, RTLIB::SMEABI_SME_STATE_SIZE,
1046 BuildMI(MBB, MBBI, DL, TII->get(TargetOpcode::COPY), BufferSize)
1047 .addReg(AArch64::X0);
1048 }
1049
1050 // Allocate a buffer object of the size given __arm_sme_state_size.
1051 {
1052 MachineFrameInfo &MFI = MF->getFrameInfo();
1053 BuildMI(MBB, MBBI, DL, TII->get(AArch64::SUBXrx64), AArch64::SP)
1054 .addReg(AArch64::SP)
1055 .addReg(BufferSize)
1057 BuildMI(MBB, MBBI, DL, TII->get(TargetOpcode::COPY), BufferPtr)
1058 .addReg(AArch64::SP);
1059
1060 // We have just allocated a variable sized object, tell this to PEI.
1061 MFI.CreateVariableSizedObject(Align(16), nullptr);
1062 }
1063
1064 restorePhyRegSave(RegSave, MBB, MBBI, DL);
1065}
1066
1067struct FromState {
1068 ZAState From;
1069
1070 constexpr uint8_t to(ZAState To) const {
1071 static_assert(NUM_ZA_STATE < 16, "expected ZAState to fit in 4-bits");
1072 return uint8_t(From) << 4 | uint8_t(To);
1073 }
1074};
1075
1076constexpr FromState transitionFrom(ZAState From) { return FromState{From}; }
1077
1078void MachineSMEABI::emitStateChange(EmitContext &Context,
1081 ZAState From, ZAState To,
1082 LiveRegs PhysLiveRegs) {
1083 // ZA not used.
1084 if (From == ZAState::ANY || To == ZAState::ANY)
1085 return;
1086
1087 // If we're exiting from the ENTRY state that means that the function has not
1088 // used ZA, so in the case of private ZA/ZT0 functions we can omit any set up.
1089 if (From == ZAState::ENTRY && To == ZAState::OFF)
1090 return;
1091
1092 // TODO: Avoid setting up the save buffer if there's no transition to
1093 // LOCAL_SAVED.
1094 if (From == ZAState::ENTRY) {
1095 assert(&MBB == &MBB.getParent()->front() &&
1096 "ENTRY state only valid in entry block");
1097 emitSMEPrologue(MBB, MBB.getFirstNonPHI());
1098 if (To == ZAState::ACTIVE)
1099 return; // Nothing more to do (ZA is active after the prologue).
1100
1101 // Note: "emitNewZAPrologue" zeros ZA, so we may need to setup a lazy save
1102 // if "To" is "ZAState::LOCAL_SAVED". It may be possible to improve this
1103 // case by changing the placement of the zero instruction.
1104 From = ZAState::ACTIVE;
1105 }
1106
1107 SMEAttrs SMEFnAttrs = AFI->getSMEFnAttrs();
1108 bool IsAgnosticZA = SMEFnAttrs.hasAgnosticZAInterface();
1109 bool HasZT0State = SMEFnAttrs.hasZT0State();
1110 bool HasZAState = IsAgnosticZA || SMEFnAttrs.hasZAState();
1111
1112 switch (transitionFrom(From).to(To)) {
1113 // This section handles: ACTIVE <-> ACTIVE_ZT0_SAVED
1114 case transitionFrom(ZAState::ACTIVE).to(ZAState::ACTIVE_ZT0_SAVED):
1115 emitZT0SaveRestore(Context, MBB, InsertPt, /*IsSave=*/true);
1116 break;
1117 case transitionFrom(ZAState::ACTIVE_ZT0_SAVED).to(ZAState::ACTIVE):
1118 emitZT0SaveRestore(Context, MBB, InsertPt, /*IsSave=*/false);
1119 break;
1120
1121 // This section handles: ACTIVE[_ZT0_SAVED] -> LOCAL_SAVED
1122 case transitionFrom(ZAState::ACTIVE).to(ZAState::LOCAL_SAVED):
1123 case transitionFrom(ZAState::ACTIVE_ZT0_SAVED).to(ZAState::LOCAL_SAVED):
1124 if (HasZT0State && From == ZAState::ACTIVE)
1125 emitZT0SaveRestore(Context, MBB, InsertPt, /*IsSave=*/true);
1126 if (HasZAState)
1127 emitZASave(Context, MBB, InsertPt, PhysLiveRegs);
1128 break;
1129
1130 // This section handles: ACTIVE -> LOCAL_COMMITTED
1131 case transitionFrom(ZAState::ACTIVE).to(ZAState::LOCAL_COMMITTED):
1132 // TODO: We could support ZA state here, but this transition is currently
1133 // only possible when we _don't_ have ZA state.
1134 assert(HasZT0State && !HasZAState && "Expect to only have ZT0 state.");
1135 emitZT0SaveRestore(Context, MBB, InsertPt, /*IsSave=*/true);
1136 emitZAMode(MBB, InsertPt, /*ClearTPIDR2=*/false, /*On=*/false);
1137 break;
1138
1139 // This section handles: LOCAL_COMMITTED -> (OFF|LOCAL_SAVED)
1140 case transitionFrom(ZAState::LOCAL_COMMITTED).to(ZAState::OFF):
1141 case transitionFrom(ZAState::LOCAL_COMMITTED).to(ZAState::LOCAL_SAVED):
1142 // These transitions are a no-op.
1143 break;
1144
1145 // This section handles: LOCAL_(SAVED|COMMITTED) -> ACTIVE[_ZT0_SAVED]
1146 case transitionFrom(ZAState::LOCAL_COMMITTED).to(ZAState::ACTIVE):
1147 case transitionFrom(ZAState::LOCAL_COMMITTED).to(ZAState::ACTIVE_ZT0_SAVED):
1148 case transitionFrom(ZAState::LOCAL_SAVED).to(ZAState::ACTIVE):
1149 case transitionFrom(ZAState::LOCAL_SAVED).to(ZAState::ACTIVE_ZT0_SAVED):
1150 if (HasZAState)
1151 emitZARestore(Context, MBB, InsertPt, PhysLiveRegs);
1152 else
1153 emitZAMode(MBB, InsertPt, /*ClearTPIDR2=*/false, /*On=*/true);
1154 if (HasZT0State && To == ZAState::ACTIVE)
1155 emitZT0SaveRestore(Context, MBB, InsertPt, /*IsSave=*/false);
1156 break;
1157
1158 // This section handles transitions to OFF (not previously covered)
1159 case transitionFrom(ZAState::ACTIVE).to(ZAState::OFF):
1160 case transitionFrom(ZAState::ACTIVE_ZT0_SAVED).to(ZAState::OFF):
1161 case transitionFrom(ZAState::LOCAL_SAVED).to(ZAState::OFF):
1162 assert(SMEFnAttrs.hasPrivateZAInterface() &&
1163 "Did not expect to turn ZA off in shared/agnostic ZA function");
1164 emitZAMode(MBB, InsertPt, /*ClearTPIDR2=*/From == ZAState::LOCAL_SAVED,
1165 /*On=*/false);
1166 break;
1167
1168 default:
1169 dbgs() << "Error: Transition from " << getZAStateString(From) << " to "
1170 << getZAStateString(To) << '\n';
1171 llvm_unreachable("Unimplemented state transition");
1172 }
1173}
1174
1175/// Returns true if private ZA setup can be elided. This occurs when there is
1176/// no instruction within the function that requires ZA to be active.
1177static bool canElidePrivateZASetup(const FunctionInfo &FnInfo) {
1178 for (const BlockInfo &BlockInfo : FnInfo.Blocks) {
1179 for (const InstInfo &InstInfo : BlockInfo.Insts) {
1180 if (InstInfo.NeededState == ZAState::ACTIVE ||
1181 InstInfo.NeededState == ZAState::ACTIVE_ZT0_SAVED)
1182 return false;
1183 }
1184 }
1185 return true;
1186}
1187
1188} // end anonymous namespace
1189
1190INITIALIZE_PASS(MachineSMEABI, "aarch64-machine-sme-abi", "Machine SME ABI",
1191 false, false)
1192
1193bool MachineSMEABI::runOnMachineFunction(MachineFunction &MF) {
1194 AFI = MF.getInfo<AArch64FunctionInfo>();
1195 SMEAttrs SMEFnAttrs = AFI->getSMEFnAttrs();
1196 if (!SMEFnAttrs.hasZAState() && !SMEFnAttrs.hasZT0State() &&
1197 !SMEFnAttrs.hasAgnosticZAInterface())
1198 return false;
1199
1200 Subtarget = &MF.getSubtarget<AArch64Subtarget>();
1201 if (!Subtarget->hasSME() && !SMEFnAttrs.hasAgnosticZAInterface())
1202 return false;
1203
1204 assert(MF.getRegInfo().isSSA() && "Expected to be run on SSA form!");
1205
1206 this->MF = &MF;
1207 ORE = &getAnalysis<MachineOptimizationRemarkEmitterPass>().getORE();
1208 LLI = &getAnalysis<LibcallLoweringInfoWrapper>().getLibcallLowering(
1209 *MF.getFunction().getParent(), *Subtarget);
1210 TII = Subtarget->getInstrInfo();
1211 TRI = Subtarget->getRegisterInfo();
1212 MRI = &MF.getRegInfo();
1213
1214 const EdgeBundles &Bundles =
1215 getAnalysis<EdgeBundlesWrapperLegacy>().getEdgeBundles();
1216
1217 FunctionInfo FnInfo = collectNeededZAStates(SMEFnAttrs);
1218
1219 if (SMEFnAttrs.hasPrivateZAInterface() && canElidePrivateZASetup(FnInfo))
1220 return false;
1221
1222 SmallVector<ZAState> BundleStates = assignBundleZAStates(Bundles, FnInfo);
1223
1224 EmitContext Context;
1225 insertStateChanges(Context, FnInfo, Bundles, BundleStates);
1226
1227 if (Context.needsSaveBuffer()) {
1228 if (FnInfo.AfterSMEProloguePt) {
1229 // Note: With inline stack probes the AfterSMEProloguePt may not be in the
1230 // entry block (due to the probing loop).
1231 MachineBasicBlock::iterator MBBI = *FnInfo.AfterSMEProloguePt;
1232 emitAllocateZASaveBuffer(Context, *MBBI->getParent(), MBBI,
1233 FnInfo.PhysLiveRegsAfterSMEPrologue);
1234 } else {
1235 MachineBasicBlock &EntryBlock = MF.front();
1236 emitAllocateZASaveBuffer(
1237 Context, EntryBlock, EntryBlock.getFirstNonPHI(),
1238 FnInfo.Blocks[EntryBlock.getNumber()].PhysLiveRegsAtEntry);
1239 }
1240 }
1241
1242 return true;
1243}
1244
1246 return new MachineSMEABI(OptLevel);
1247}
static constexpr unsigned ZERO_ALL_ZA_MASK
assert(UImm &&(UImm !=~static_cast< T >(0)) &&"Invalid immediate!")
MachineBasicBlock & MBB
MachineBasicBlock MachineBasicBlock::iterator DebugLoc DL
MachineBasicBlock MachineBasicBlock::iterator MBBI
static GCRegistry::Add< CoreCLRGC > E("coreclr", "CoreCLR-compatible GC")
const HexagonInstrInfo * TII
#define _
IRTranslator LLVM IR MI
This file implements the LivePhysRegs utility for tracking liveness of physical registers.
#define ENTRY(ASMNAME, ENUM)
#define I(x, y, z)
Definition MD5.cpp:57
===- MachineOptimizationRemarkEmitter.h - Opt Diagnostics -*- C++ -*-—===//
#define MAKE_CASE(V)
Register const TargetRegisterInfo * TRI
Promote Memory to Register
Definition Mem2Reg.cpp:110
if(PassOpts->AAPipeline)
#define INITIALIZE_PASS(passName, arg, name, cfg, analysis)
Definition PassSupport.h:56
Func MI getDebugLoc()))
This file defines the SmallVector class.
AArch64FunctionInfo - This class is derived from MachineFunctionInfo and contains private AArch64-spe...
Represent the analysis usage information of a pass.
AnalysisUsage & addPreservedID(const void *ID)
AnalysisUsage & addRequired()
LLVM_ABI void setPreservesCFG()
This function should be called by the pass, iff they do not:
Definition Pass.cpp:278
Represent a constant reference to an array (0 or more elements consecutively in memory),...
Definition ArrayRef.h:40
This class represents a function call, abstracting a target machine's calling convention.
A debug info location.
Definition DebugLoc.h:126
ArrayRef< unsigned > getBlocks(unsigned Bundle) const
getBlocks - Return an array of blocks that are connected to Bundle.
Definition EdgeBundles.h:53
unsigned getBundle(unsigned N, bool Out) const
getBundle - Return the ingoing (Out = false) or outgoing (Out = true) bundle number for basic block N
Definition EdgeBundles.h:47
unsigned getNumBundles() const
getNumBundles - Return the total number of bundles in the CFG.
Definition EdgeBundles.h:50
FunctionPass class - This class is used to implement most global optimizations.
Definition Pass.h:314
LLVMContext & getContext() const
getContext - Return a reference to the LLVMContext associated with this function.
Definition Function.cpp:356
const DebugLoc & getDebugLoc() const
Return the debug location for this node as a DebugLoc.
This is an important class for using LLVM in a threaded context.
Definition LLVMContext.h:68
LLVM_ABI void emitError(const Instruction *I, const Twine &ErrorStr)
emitError - Emit an error message to the currently installed error handler with optional location inf...
Tracks which library functions to use for a particular subtarget or function.
CallingConv::ID getLibcallImplCallingConv(RTLIB::LibcallImpl Call) const
Get the CallingConv that should be used for the specified libcall.
RTLIB::LibcallImpl getLibcallImpl(RTLIB::Libcall Call) const
Return the lowering's selection of implementation call for Call.
A set of register units used to track register liveness.
bool available(MCRegister Reg) const
Returns true if no part of physical register Reg is live.
void addReg(MCRegister Reg)
Adds register units covered by physical register Reg.
LLVM_ABI void stepBackward(const MachineInstr &MI)
Updates liveness when stepping backwards over the instruction MI.
LLVM_ABI void addLiveOuts(const MachineBasicBlock &MBB)
Adds registers living out of block MBB.
MachineInstrBundleIterator< const MachineInstr > const_iterator
int getNumber() const
MachineBasicBlocks are uniquely numbered at the function level, unless they're not in a MachineFuncti...
LLVM_ABI iterator getFirstNonPHI()
Returns a pointer to the first instruction in this block that is not a PHINode instruction.
succ_reverse_iterator succ_rbegin()
MachineInstrBundleIterator< MachineInstr > iterator
succ_reverse_iterator succ_rend()
The MachineFrameInfo class represents an abstract stack frame until prolog/epilog code is inserted.
LLVM_ABI int CreateStackObject(uint64_t Size, Align Alignment, bool isSpillSlot, const AllocaInst *Alloca=nullptr, uint8_t ID=0)
Create a new statically sized stack object, returning a nonnegative identifier to represent it.
LLVM_ABI int CreateSpillStackObject(uint64_t Size, Align Alignment, TargetStackID::Value StackID=TargetStackID::Default)
Create a new statically sized stack object that represents a spill slot, returning a nonnegative iden...
LLVM_ABI int CreateVariableSizedObject(Align Alignment, const AllocaInst *Alloca)
Notify the MachineFrameInfo object that a variable sized object has been created.
MachineFunctionPass - This class adapts the FunctionPass interface to allow convenient creation of pa...
void getAnalysisUsage(AnalysisUsage &AU) const override
getAnalysisUsage - Subclasses that override getAnalysisUsage must call this.
StringRef getName() const
getName - Return the name of the corresponding LLVM function.
MachineFrameInfo & getFrameInfo()
getFrameInfo - Return the frame info object for the current function.
MachineRegisterInfo & getRegInfo()
getRegInfo - Return information about the registers currently in use.
Function & getFunction()
Return the LLVM function that this machine code represents.
unsigned getNumBlockIDs() const
getNumBlockIDs - Return the number of MBB ID's allocated.
Ty * getInfo()
getInfo - Keep track of various per-function pieces of information for backends that would like to do...
const MachineInstrBuilder & addExternalSymbol(const char *FnName, unsigned TargetFlags=0) const
const MachineInstrBuilder & setOperandDead(unsigned OpIdx) const
const MachineInstrBuilder & addReg(Register RegNo, RegState Flags={}, unsigned SubReg=0) const
Add a new virtual register operand.
const MachineInstrBuilder & addImm(int64_t Val) const
Add a new immediate operand.
const MachineInstrBuilder & addFrameIndex(int Idx) const
const MachineInstrBuilder & addRegMask(const uint32_t *Mask) const
const MachineInstrBuilder & addDef(Register RegNo, RegState Flags={}, unsigned SubReg=0) const
Add a virtual register definition operand.
Representation of each machine instruction.
MachineOperand class - Representation of each machine instruction operand.
const GlobalValue * getGlobal() const
bool isReg() const
isReg - Tests if this is a MO_Register operand.
bool isSymbol() const
isSymbol - Tests if this is a MO_ExternalSymbol operand.
bool isGlobal() const
isGlobal - Tests if this is a MO_GlobalAddress operand.
const char * getSymbolName() const
Register getReg() const
getReg - Returns the register number.
Diagnostic information for optimization analysis remarks.
LLVM_ABI void emit(DiagnosticInfoOptimizationBase &OptDiag)
Emit an optimization remark.
bool allowExtraAnalysis(StringRef PassName) const
Whether we allow for extra compile-time budget to perform more analysis to be more informative.
MachineRegisterInfo - Keep track of information for virtual and physical registers,...
LLVM_ABI Register createVirtualRegister(const TargetRegisterClass *RegClass, StringRef Name="")
createVirtualRegister - Create and return a new virtual register in the function with the specified r...
Wrapper class representing virtual and physical registers.
Definition Register.h:20
constexpr bool isValid() const
Definition Register.h:112
constexpr bool isPhysical() const
Return true if the specified register number is in the physical register namespace.
Definition Register.h:83
SMEAttrs is a utility class to parse the SME ACLE attributes on functions.
bool hasAgnosticZAInterface() const
bool hasPrivateZAInterface() const
bool hasSharedZAInterface() const
std::pair< iterator, bool > insert(PtrType Ptr)
Inserts Ptr if and only if there is no element in the container equal to Ptr.
SmallPtrSet - This class implements a set which is optimized for holding SmallSize or less elements.
This class consists of common code factored out of the SmallVector class to reduce code duplication b...
typename SuperClass::const_iterator const_iterator
void append(ItTy in_start, ItTy in_end)
Add the specified range to the end of the SmallVector.
void push_back(const T &Elt)
This is a 'vector' (really, a variable-sized array), optimized for the case when the array is small.
Represent a constant reference to a string, i.e.
Definition StringRef.h:56
constexpr bool empty() const
Check if the string is empty.
Definition StringRef.h:141
constexpr const char * data() const
Get a pointer to the start of the string (which may not be null terminated).
Definition StringRef.h:138
TargetRegisterInfo base class - We assume that the target defines a static array of TargetRegisterDes...
Twine - A lightweight data structure for efficiently representing the concatenation of temporary valu...
Definition Twine.h:82
op_range operands()
Definition User.h:267
LLVM_ABI StringRef getName() const
Return a constant reference to the value's name.
Definition Value.cpp:319
const ParentTy * getParent() const
Definition ilist_node.h:34
#define llvm_unreachable(msg)
Marks that the current location is not supposed to be reachable.
static unsigned getArithExtendImm(AArch64_AM::ShiftExtendType ET, unsigned Imm)
getArithExtendImm - Encode the extend type and shift amount for an arithmetic instruction: imm: 3-bit...
unsigned ID
LLVM IR allows to use arbitrary numbers as calling convention identifiers.
Definition CallingConv.h:24
@ AArch64_SME_ABI_Support_Routines_PreserveMost_From_X0
Preserve X0-X13, X19-X29, SP, Z0-Z31, P0-P15.
@ AArch64_SME_ABI_Support_Routines_PreserveMost_From_X1
Preserve X1-X15, X19-X29, SP, Z0-Z31, P0-P15.
This is an optimization pass for GlobalISel generic memory operations.
MachineInstrBuilder BuildMI(MachineFunction &MF, const MIMetadata &MIMD, const MCInstrDesc &MCID)
Builder interface. Specify how to create the initial instruction itself.
@ Implicit
Not emitted register (e.g. carry, or temporary result).
@ Define
Register definition.
FunctionPass * createMachineSMEABIPass(CodeGenOptLevel)
LLVM_ABI char & MachineDominatorsID
MachineDominators - This pass is a machine dominators analysis pass.
LLVM_ABI void reportFatalInternalError(Error Err)
Report a fatal error that indicates a bug in LLVM.
Definition Error.cpp:173
LLVM_ABI char & MachineLoopInfoID
MachineLoopInfo - This pass is a loop analysis pass.
bool any_of(R &&range, UnaryPredicate P)
Provide wrappers to std::any_of which take ranges instead of having to pass begin/end explicitly.
Definition STLExtras.h:1762
auto reverse(ContainerTy &&C)
Definition STLExtras.h:408
LLVM_ABI raw_ostream & dbgs()
dbgs() - This returns a reference to a raw_ostream for debugging messages.
Definition Debug.cpp:209
CodeGenOptLevel
Code generation optimization level.
Definition CodeGen.h:227
@ Default
-O2, -Os, -Oz
Definition CodeGen.h:230
uint16_t MCPhysReg
An unsigned integer type large enough to represent all physical registers, but not necessarily virtua...
Definition MCRegister.h:21
This struct is a compact representation of a valid (non-zero power of two) alignment.
Definition Alignment.h:39
static StringRef getLibcallImplName(RTLIB::LibcallImpl CallImpl)
Get the libcall routine name for the specified libcall implementation.