[llvm] [AArch64][SME] Use LibcallLoweringInfo in the MachineSMEABIPass (PR #177762)

via llvm-commits llvm-commits at lists.llvm.org
Sat Jan 24 04:48:05 PST 2026


llvmbot wrote:


<!--LLVM PR SUMMARY COMMENT-->

@llvm/pr-subscribers-backend-aarch64

Author: Benjamin Maxwell (MacDue)

<details>
<summary>Changes</summary>

This adds a new helper to add calls to SME routines (addSMELibCall) and check they are using the expected CC.

---
Full diff: https://github.com/llvm/llvm-project/pull/177762.diff


1 Files Affected:

- (modified) llvm/lib/Target/AArch64/MachineSMEABIPass.cpp (+43-30) 


``````````diff
diff --git a/llvm/lib/Target/AArch64/MachineSMEABIPass.cpp b/llvm/lib/Target/AArch64/MachineSMEABIPass.cpp
index 823c754a0ac05..74e47dac0ac05 100644
--- a/llvm/lib/Target/AArch64/MachineSMEABIPass.cpp
+++ b/llvm/lib/Target/AArch64/MachineSMEABIPass.cpp
@@ -305,6 +305,7 @@ struct MachineSMEABI : public MachineFunctionPass {
     AU.setPreservesCFG();
     AU.addRequired<EdgeBundlesWrapperLegacy>();
     AU.addRequired<MachineOptimizationRemarkEmitterPass>();
+    AU.addRequired<LibcallLoweringInfoWrapper>();
     AU.addPreservedID(MachineLoopInfoID);
     AU.addPreservedID(MachineDominatorsID);
     MachineFunctionPass::getAnalysisUsage(AU);
@@ -330,6 +331,9 @@ struct MachineSMEABI : public MachineFunctionPass {
   /// predecessors).
   void propagateDesiredStates(FunctionInfo &FnInfo, bool Forwards = true);
 
+  void addSMELibCall(MachineInstrBuilder &MIB, RTLIB::Libcall LC,
+                     CallingConv::ID ExpectedCC);
+
   void emitZT0SaveRestore(EmitContext &, MachineBasicBlock &MBB,
                           MachineBasicBlock::iterator MBBI, bool IsSave);
 
@@ -420,7 +424,8 @@ struct MachineSMEABI : public MachineFunctionPass {
   const AArch64Subtarget *Subtarget = nullptr;
   const AArch64RegisterInfo *TRI = nullptr;
   const AArch64FunctionInfo *AFI = nullptr;
-  const TargetInstrInfo *TII = nullptr;
+  const AArch64InstrInfo *TII = nullptr;
+  const LibcallLoweringInfo *LLI = nullptr;
 
   MachineOptimizationRemarkEmitter *ORE = nullptr;
   MachineRegisterInfo *MRI = nullptr;
@@ -876,11 +881,22 @@ void MachineSMEABI::restorePhyRegSave(const PhysRegSave &RegSave,
         .addReg(RegSave.X0Save);
 }
 
+void MachineSMEABI::addSMELibCall(MachineInstrBuilder &MIB, RTLIB::Libcall LC,
+                                  CallingConv::ID ExpectedCC) {
+  RTLIB::LibcallImpl LCImpl = LLI->getLibcallImpl(LC);
+  assert(LCImpl != RTLIB::Unsupported && "Expected SME routines to exist.");
+  CallingConv::ID CC = LLI->getLibcallImplCallingConv(LCImpl);
+  assert(CC == ExpectedCC && "Unexpected calling convention for SME rountine");
+  StringRef SymbolName = RTLIB::RuntimeLibcallsInfo::getLibcallImplName(LCImpl);
+  // FIXME: This assumes the SymbolName StringRef is null-terminated.
+  MIB.addExternalSymbol(SymbolName.data());
+  MIB.addRegMask(TRI->getCallPreservedMask(*MF, CC));
+}
+
 void MachineSMEABI::emitRestoreLazySave(EmitContext &Context,
                                         MachineBasicBlock &MBB,
                                         MachineBasicBlock::iterator MBBI,
                                         LiveRegs PhysLiveRegs) {
-  auto *TLI = Subtarget->getTargetLowering();
   DebugLoc DL = getDebugLoc(MBB, MBBI);
   Register TPIDR2EL0 = MRI->createVirtualRegister(&AArch64::GPR64RegClass);
   Register TPIDR2 = AArch64::X0;
@@ -901,11 +917,12 @@ void MachineSMEABI::emitRestoreLazySave(EmitContext &Context,
       .addImm(0)
       .addImm(0);
   // (Conditionally) restore ZA state.
-  BuildMI(MBB, MBBI, DL, TII->get(AArch64::RestoreZAPseudo))
-      .addReg(TPIDR2EL0)
-      .addReg(TPIDR2)
-      .addExternalSymbol(TLI->getLibcallName(RTLIB::SMEABI_TPIDR2_RESTORE))
-      .addRegMask(TRI->SMEABISupportRoutinesCallPreservedMaskFromX0());
+  auto RestoreZA = BuildMI(MBB, MBBI, DL, TII->get(AArch64::RestoreZAPseudo))
+                       .addReg(TPIDR2EL0)
+                       .addReg(TPIDR2);
+  addSMELibCall(
+      RestoreZA, RTLIB::SMEABI_TPIDR2_RESTORE,
+      CallingConv::AArch64_SME_ABI_Support_Routines_PreserveMost_From_X0);
   // Zero TPIDR2_EL0.
   BuildMI(MBB, MBBI, DL, TII->get(AArch64::MSR))
       .addImm(AArch64SysReg::TPIDR2_EL0)
@@ -989,7 +1006,6 @@ static constexpr unsigned ZERO_ALL_ZA_MASK = 0b11111111;
 
 void MachineSMEABI::emitSMEPrologue(MachineBasicBlock &MBB,
                                     MachineBasicBlock::iterator MBBI) {
-  auto *TLI = Subtarget->getTargetLowering();
   DebugLoc DL = getDebugLoc(MBB, MBBI);
 
   bool ZeroZA = AFI->getSMEFnAttrs().isNewZA();
@@ -1006,9 +1022,10 @@ void MachineSMEABI::emitSMEPrologue(MachineBasicBlock &MBB,
         BuildMI(MBB, MBBI, DL, TII->get(AArch64::CommitZASavePseudo))
             .addReg(TPIDR2EL0)
             .addImm(ZeroZA)
-            .addImm(ZeroZT0)
-            .addExternalSymbol(TLI->getLibcallName(RTLIB::SMEABI_TPIDR2_SAVE))
-            .addRegMask(TRI->SMEABISupportRoutinesCallPreservedMaskFromX0());
+            .addImm(ZeroZT0);
+    addSMELibCall(
+        CommitZASave, RTLIB::SMEABI_TPIDR2_SAVE,
+        CallingConv::AArch64_SME_ABI_Support_Routines_PreserveMost_From_X0);
     if (ZeroZA)
       CommitZASave.addDef(AArch64::ZAB0, RegState::ImplicitDefine);
     if (ZeroZT0)
@@ -1031,8 +1048,6 @@ void MachineSMEABI::emitFullZASaveRestore(EmitContext &Context,
                                           MachineBasicBlock &MBB,
                                           MachineBasicBlock::iterator MBBI,
                                           LiveRegs PhysLiveRegs, bool IsSave) {
-  auto *TLI = Subtarget->getTargetLowering();
-
   DebugLoc DL = getDebugLoc(MBB, MBBI);
 
   if (IsSave)
@@ -1047,13 +1062,12 @@ void MachineSMEABI::emitFullZASaveRestore(EmitContext &Context,
       .addReg(Context.getAgnosticZABufferPtr(*MF));
 
   // Call __arm_sme_save/__arm_sme_restore.
-  BuildMI(MBB, MBBI, DL, TII->get(AArch64::BL))
-      .addReg(BufferPtr, RegState::Implicit)
-      .addExternalSymbol(TLI->getLibcallName(
-          IsSave ? RTLIB::SMEABI_SME_SAVE : RTLIB::SMEABI_SME_RESTORE))
-      .addRegMask(TRI->getCallPreservedMask(
-          *MF,
-          CallingConv::AArch64_SME_ABI_Support_Routines_PreserveMost_From_X1));
+  auto SaveRestoreZA = BuildMI(MBB, MBBI, DL, TII->get(AArch64::BL))
+                           .addReg(BufferPtr, RegState::Implicit);
+  addSMELibCall(
+      SaveRestoreZA,
+      IsSave ? RTLIB::SMEABI_SME_SAVE : RTLIB::SMEABI_SME_RESTORE,
+      CallingConv::AArch64_SME_ABI_Support_Routines_PreserveMost_From_X1);
 
   restorePhyRegSave(RegSave, MBB, MBBI, DL);
 }
@@ -1103,14 +1117,11 @@ void MachineSMEABI::emitAllocateFullZASaveBuffer(
 
   // Calculate the SME state size.
   {
-    auto *TLI = Subtarget->getTargetLowering();
-    const AArch64RegisterInfo *TRI = Subtarget->getRegisterInfo();
-    BuildMI(MBB, MBBI, DL, TII->get(AArch64::BL))
-        .addExternalSymbol(TLI->getLibcallName(RTLIB::SMEABI_SME_STATE_SIZE))
-        .addReg(AArch64::X0, RegState::ImplicitDefine)
-        .addRegMask(TRI->getCallPreservedMask(
-            *MF, CallingConv::
-                     AArch64_SME_ABI_Support_Routines_PreserveMost_From_X1));
+    auto SMEStateSize = BuildMI(MBB, MBBI, DL, TII->get(AArch64::BL))
+                            .addReg(AArch64::X0, RegState::ImplicitDefine);
+    addSMELibCall(
+        SMEStateSize, RTLIB::SMEABI_SME_STATE_SIZE,
+        CallingConv::AArch64_SME_ABI_Support_Routines_PreserveMost_From_X1);
     BuildMI(MBB, MBBI, DL, TII->get(TargetOpcode::COPY), BufferSize)
         .addReg(AArch64::X0);
   }
@@ -1245,7 +1256,8 @@ INITIALIZE_PASS(MachineSMEABI, "aarch64-machine-sme-abi", "Machine SME ABI",
                 false, false)
 
 bool MachineSMEABI::runOnMachineFunction(MachineFunction &MF) {
-  if (!MF.getSubtarget<AArch64Subtarget>().hasSME())
+  Subtarget = &MF.getSubtarget<AArch64Subtarget>();
+  if (!Subtarget->hasSME())
     return false;
 
   AFI = MF.getInfo<AArch64FunctionInfo>();
@@ -1257,8 +1269,9 @@ bool MachineSMEABI::runOnMachineFunction(MachineFunction &MF) {
   assert(MF.getRegInfo().isSSA() && "Expected to be run on SSA form!");
 
   this->MF = &MF;
-  Subtarget = &MF.getSubtarget<AArch64Subtarget>();
   ORE = &getAnalysis<MachineOptimizationRemarkEmitterPass>().getORE();
+  LLI = &getAnalysis<LibcallLoweringInfoWrapper>().getLibcallLowering(
+      *MF.getFunction().getParent(), *Subtarget);
   TII = Subtarget->getInstrInfo();
   TRI = Subtarget->getRegisterInfo();
   MRI = &MF.getRegInfo();

``````````

</details>


https://github.com/llvm/llvm-project/pull/177762


More information about the llvm-commits mailing list