[llvm] [AArch64][SME] Peephole optimization to push smstop za earlier (PR #226160)

Benjamin Maxwell via llvm-commits llvm-commits at lists.llvm.org
Thu Sep 24 06:24:14 PDT 2026


https://github.com/MacDue created https://github.com/llvm/llvm-project/pull/226160

This patch adds a simple peephole optimization that pushes ZAState::OFF earlier to avoid redundant saves at the ends of functions. An example is calling a private ZA function before returning in a function with new ZA/ZT0 state.

>From 919cff9fc6d6e655fc65f51c9226fab36e736d4c Mon Sep 17 00:00:00 2001
From: Benjamin Maxwell <benjamin.maxwell at arm.com>
Date: Thu, 24 Sep 2026 09:21:18 +0000
Subject: [PATCH 1/2] Precommit some extra tests

---
 .../CodeGen/AArch64/sme-lazy-save-call.ll     | 52 ++++++++++
 llvm/test/CodeGen/AArch64/sme-zt0-state.ll    | 99 +++++++++++++++++++
 2 files changed, 151 insertions(+)

diff --git a/llvm/test/CodeGen/AArch64/sme-lazy-save-call.ll b/llvm/test/CodeGen/AArch64/sme-lazy-save-call.ll
index bbdda5fa8f484f..3037f516642c2a 100644
--- a/llvm/test/CodeGen/AArch64/sme-lazy-save-call.ll
+++ b/llvm/test/CodeGen/AArch64/sme-lazy-save-call.ll
@@ -4,6 +4,7 @@
 declare void @private_za_callee()
 declare void @shared_za_callee() "aarch64_inout_za"
 declare void @preserves_za_callee() "aarch64_preserves_za"
+declare void @other_private_za_callee(i64)
 
 declare float @llvm.cos.f32(float)
 
@@ -395,3 +396,54 @@ define i64  @test_many_callee_arguments(
     i64 %0, i64 %1, i64 %2, i64 %3, i64 %4, i64 %5, i64 %6, i64 %7, i64 %8, i64 %9)
   ret i64 %ret
 }
+
+define void @no_lazy_save_for_private_return_arms(i1 %cond) "aarch64_new_za" nounwind {
+; CHECK-LABEL: no_lazy_save_for_private_return_arms:
+; CHECK:       // %bb.0: // %entry
+; CHECK-NEXT:    stp x29, x30, [sp, #-32]! // 16-byte Folded Spill
+; CHECK-NEXT:    stp x20, x19, [sp, #16] // 16-byte Folded Spill
+; CHECK-NEXT:    mov x29, sp
+; CHECK-NEXT:    sub sp, sp, #16
+; CHECK-NEXT:    rdsvl x8, #1
+; CHECK-NEXT:    mov x9, sp
+; CHECK-NEXT:    msub x9, x8, x8, x9
+; CHECK-NEXT:    mov sp, x9
+; CHECK-NEXT:    stp x9, x8, [x29, #-16]
+; CHECK-NEXT:    mrs x8, TPIDR2_EL0
+; CHECK-NEXT:    cbz x8, .LBB10_2
+; CHECK-NEXT:  // %bb.1: // %entry
+; CHECK-NEXT:    bl __arm_tpidr2_save
+; CHECK-NEXT:    msr TPIDR2_EL0, xzr
+; CHECK-NEXT:    zero {za}
+; CHECK-NEXT:  .LBB10_2: // %entry
+; CHECK-NEXT:    smstart za
+; CHECK-NEXT:    mov w20, w0
+; CHECK-NEXT:    bl shared_za_callee
+; CHECK-NEXT:    sub x8, x29, #16
+; CHECK-NEXT:    msr TPIDR2_EL0, x8
+; CHECK-NEXT:    tbz w20, #0, .LBB10_4
+; CHECK-NEXT:  // %bb.3: // %left
+; CHECK-NEXT:    bl private_za_callee
+; CHECK-NEXT:    b .LBB10_5
+; CHECK-NEXT:  .LBB10_4: // %right
+; CHECK-NEXT:    mov w0, #42 // =0x2a
+; CHECK-NEXT:    bl other_private_za_callee
+; CHECK-NEXT:  .LBB10_5: // %common.ret
+; CHECK-NEXT:    msr TPIDR2_EL0, xzr
+; CHECK-NEXT:    smstop za
+; CHECK-NEXT:    mov sp, x29
+; CHECK-NEXT:    ldp x20, x19, [sp, #16] // 16-byte Folded Reload
+; CHECK-NEXT:    ldp x29, x30, [sp], #32 // 16-byte Folded Reload
+; CHECK-NEXT:    ret
+entry:
+  call void @shared_za_callee()
+  br i1 %cond, label %left, label %right
+
+left:
+  call void @private_za_callee()
+  ret void
+
+right:
+  call void @other_private_za_callee(i64 42)
+  ret void
+}
diff --git a/llvm/test/CodeGen/AArch64/sme-zt0-state.ll b/llvm/test/CodeGen/AArch64/sme-zt0-state.ll
index 7ef3262e5811c3..788ca7d932b2cc 100644
--- a/llvm/test/CodeGen/AArch64/sme-zt0-state.ll
+++ b/llvm/test/CodeGen/AArch64/sme-zt0-state.ll
@@ -328,3 +328,102 @@ define void @za_zt0_private_za_to_shared_za(ptr %callee) "aarch64_inout_za" "aar
   call void %callee() "aarch64_inout_za";
   ret void;
 }
+
+define void @no_need_to_save_zt0(ptr %callee) "aarch64_new_za" "aarch64_new_zt0" nounwind {
+; CHECK-LABEL: no_need_to_save_zt0:
+; CHECK:       // %bb.0:
+; CHECK-NEXT:    sub sp, sp, #80
+; CHECK-NEXT:    str x30, [sp, #64] // 8-byte Spill
+; CHECK-NEXT:    mrs x8, TPIDR2_EL0
+; CHECK-NEXT:    cbz x8, .LBB15_2
+; CHECK-NEXT:  // %bb.1:
+; CHECK-NEXT:    bl __arm_tpidr2_save
+; CHECK-NEXT:    msr TPIDR2_EL0, xzr
+; CHECK-NEXT:    zero {za}
+; CHECK-NEXT:    zero { zt0 }
+; CHECK-NEXT:  .LBB15_2:
+; CHECK-NEXT:    smstart za
+; CHECK-NEXT:    mov x8, sp
+; CHECK-NEXT:    str zt0, [x8]
+; CHECK-NEXT:    blr x0
+; CHECK-NEXT:    smstop za
+; CHECK-NEXT:    ldr x30, [sp, #64] // 8-byte Reload
+; CHECK-NEXT:    add sp, sp, #80
+; CHECK-NEXT:    ret
+  call void %callee() "aarch64_inout_za"
+  ret void;
+}
+
+define void @no_need_to_save_zt0_after_call(ptr %callee) "aarch64_new_za" "aarch64_new_zt0" nounwind {
+; CHECK-LABEL: no_need_to_save_zt0_after_call:
+; CHECK:       // %bb.0:
+; CHECK-NEXT:    sub sp, sp, #80
+; CHECK-NEXT:    stp x30, x19, [sp, #64] // 16-byte Folded Spill
+; CHECK-NEXT:    mrs x8, TPIDR2_EL0
+; CHECK-NEXT:    cbz x8, .LBB16_2
+; CHECK-NEXT:  // %bb.1:
+; CHECK-NEXT:    bl __arm_tpidr2_save
+; CHECK-NEXT:    msr TPIDR2_EL0, xzr
+; CHECK-NEXT:    zero {za}
+; CHECK-NEXT:    zero { zt0 }
+; CHECK-NEXT:  .LBB16_2:
+; CHECK-NEXT:    smstart za
+; CHECK-NEXT:    mov x19, x0
+; CHECK-NEXT:    blr x0
+; CHECK-NEXT:    mov x8, sp
+; CHECK-NEXT:    str zt0, [x8]
+; CHECK-NEXT:    blr x19
+; CHECK-NEXT:    smstop za
+; CHECK-NEXT:    ldp x30, x19, [sp, #64] // 16-byte Folded Reload
+; CHECK-NEXT:    add sp, sp, #80
+; CHECK-NEXT:    ret
+  call void %callee() "aarch64_inout_za" "aarch64_inout_zt0"
+  call void %callee() "aarch64_inout_za"
+  ret void;
+}
+
+define void @no_restore_of_dead_zt0(ptr %private, ptr %shared_za) "aarch64_new_za" "aarch64_new_zt0" nounwind {
+; CHECK-LABEL: no_restore_of_dead_zt0:
+; CHECK:       // %bb.0:
+; CHECK-NEXT:    stp x29, x30, [sp, #-32]! // 16-byte Folded Spill
+; CHECK-NEXT:    str x19, [sp, #16] // 8-byte Spill
+; CHECK-NEXT:    mov x29, sp
+; CHECK-NEXT:    sub sp, sp, #80
+; CHECK-NEXT:    rdsvl x8, #1
+; CHECK-NEXT:    mov x9, sp
+; CHECK-NEXT:    msub x9, x8, x8, x9
+; CHECK-NEXT:    mov sp, x9
+; CHECK-NEXT:    stp x9, x8, [x29, #-80]
+; CHECK-NEXT:    mrs x8, TPIDR2_EL0
+; CHECK-NEXT:    cbz x8, .LBB17_2
+; CHECK-NEXT:  // %bb.1:
+; CHECK-NEXT:    bl __arm_tpidr2_save
+; CHECK-NEXT:    msr TPIDR2_EL0, xzr
+; CHECK-NEXT:    zero {za}
+; CHECK-NEXT:    zero { zt0 }
+; CHECK-NEXT:  .LBB17_2:
+; CHECK-NEXT:    smstart za
+; CHECK-NEXT:    sub x8, x29, #64
+; CHECK-NEXT:    sub x9, x29, #80
+; CHECK-NEXT:    mov x19, x1
+; CHECK-NEXT:    str zt0, [x8]
+; CHECK-NEXT:    msr TPIDR2_EL0, x9
+; CHECK-NEXT:    blr x0
+; CHECK-NEXT:    smstart za
+; CHECK-NEXT:    mrs x8, TPIDR2_EL0
+; CHECK-NEXT:    sub x0, x29, #80
+; CHECK-NEXT:    cbnz x8, .LBB17_4
+; CHECK-NEXT:  // %bb.3:
+; CHECK-NEXT:    bl __arm_tpidr2_restore
+; CHECK-NEXT:  .LBB17_4:
+; CHECK-NEXT:    msr TPIDR2_EL0, xzr
+; CHECK-NEXT:    blr x19
+; CHECK-NEXT:    smstop za
+; CHECK-NEXT:    mov sp, x29
+; CHECK-NEXT:    ldr x19, [sp, #16] // 8-byte Reload
+; CHECK-NEXT:    ldp x29, x30, [sp], #32 // 16-byte Folded Reload
+; CHECK-NEXT:    ret
+  call void %private()
+  call void %shared_za() "aarch64_inout_za"
+  ret void
+}

>From 0a3e3688dbaf5d71c7745cb60bd6755061ad6cee Mon Sep 17 00:00:00 2001
From: Benjamin Maxwell <benjamin.maxwell at arm.com>
Date: Thu, 24 Sep 2026 09:27:25 +0000
Subject: [PATCH 2/2] [AArch64][SME] Peephole optimization to push `smstop za`
 earlier

This patch adds a simple peephole optimization that pushes ZAState::OFF
earlier to avoid redundant saves at the ends of functions. An example is
calling a private ZA function before returning in a function with new
ZA/ZT0 state.
---
 llvm/lib/Target/AArch64/MachineSMEABIPass.cpp | 57 +++++++++++++++++--
 .../CodeGen/AArch64/sme-lazy-save-call.ll     |  5 +-
 .../test/CodeGen/AArch64/sme-za-exceptions.ll |  5 +-
 .../sme-za-function-with-many-blocks.ll       |  2 -
 llvm/test/CodeGen/AArch64/sme-zt0-state.ll    | 16 ++----
 5 files changed, 59 insertions(+), 26 deletions(-)

diff --git a/llvm/lib/Target/AArch64/MachineSMEABIPass.cpp b/llvm/lib/Target/AArch64/MachineSMEABIPass.cpp
index 5dd62d1d66278a..04b95419daaf16 100644
--- a/llvm/lib/Target/AArch64/MachineSMEABIPass.cpp
+++ b/llvm/lib/Target/AArch64/MachineSMEABIPass.cpp
@@ -219,6 +219,11 @@ StringRef getZAStateString(ZAState State) {
 #undef MAKE_CASE
 }
 
+/// Returns true if \p State could requires ZA to be on.
+static bool isRequiresZAOn(ZAState State) {
+  return State == ZAState::ACTIVE || State == ZAState::ACTIVE_ZT0_SAVED;
+}
+
 static bool isZAorZTRegOp(const TargetRegisterInfo &TRI,
                           const MachineOperand &MO) {
   if (!MO.isReg() || !MO.getReg().isPhysical())
@@ -298,6 +303,9 @@ struct MachineSMEABI : public MachineFunctionPass {
   /// within the machine function.
   FunctionInfo collectNeededZAStates(SMEAttrs SMEFnAttrs);
 
+  /// Simple peephole optimizations to remove redundant saves and restores.
+  void peepholeOptimizeStateChanges(FunctionInfo &FnInfo);
+
   /// Assigns each edge bundle a ZA state based on the desired states of
   /// incoming and outgoing blocks in the bundle.
   SmallVector<ZAState> assignBundleZAStates(const EdgeBundles &Bundles,
@@ -518,6 +526,43 @@ FunctionInfo MachineSMEABI::collectNeededZAStates(SMEAttrs SMEFnAttrs) {
                       PhysLiveRegsAfterSMEPrologue};
 }
 
+void MachineSMEABI::peepholeOptimizeStateChanges(FunctionInfo &FnInfo) {
+  if (OptLevel == CodeGenOptLevel::None)
+    return;
+
+  for (BlockInfo &Block : FnInfo.Blocks) {
+    if (Block.Insts.size() <= 1 ||
+        Block.Insts.back().NeededState != ZAState::OFF)
+      continue;
+
+    unsigned DeadTransitions = 0;
+    for (unsigned I = Block.Insts.size() - 1; I >= 1; --I) {
+      assert(Block.Insts[I].NeededState == ZAState::OFF);
+      InstInfo &PreviousInst = Block.Insts[I - 1];
+
+      if (!isRequiresZAOn(PreviousInst.NeededState)) {
+        // If the previous state requires a save and the current state is OFF,
+        // set the previous state to OFF to avoid a redundant save.
+        PreviousInst.NeededState = ZAState::OFF;
+        ++DeadTransitions;
+      } else if (PreviousInst.NeededState == ZAState::ACTIVE_ZT0_SAVED &&
+                 ((I >= 2 &&
+                   Block.Insts[I - 2].NeededState == ZAState::ACTIVE) ||
+                  (I == 1 && Block.FixedEntryState == ZAState::ENTRY))) {
+        // If the previous state is ZT0 saved, the current state is OFF, and the
+        // state before that is ACTIVE. Then set the previous state to ACTIVE to
+        // avoid a redundant ZT0 save.
+        PreviousInst.NeededState = ZAState::ACTIVE;
+      }
+
+      if (PreviousInst.NeededState != ZAState::OFF)
+        break; // No more folds possible.
+    }
+    Block.DesiredIncomingState = Block.Insts.front().NeededState;
+    Block.Insts.pop_back_n(DeadTransitions);
+  }
+}
+
 /// Assigns each edge bundle a ZA state based on the desired states of incoming
 /// and outgoing blocks in the bundle.
 SmallVector<ZAState>
@@ -1174,8 +1219,7 @@ void MachineSMEABI::emitStateChange(EmitContext &Context,
 static bool canElidePrivateZASetup(const FunctionInfo &FnInfo) {
   for (const BlockInfo &BlockInfo : FnInfo.Blocks) {
     for (const InstInfo &InstInfo : BlockInfo.Insts) {
-      if (InstInfo.NeededState == ZAState::ACTIVE ||
-          InstInfo.NeededState == ZAState::ACTIVE_ZT0_SAVED)
+      if (isRequiresZAOn(InstInfo.NeededState))
         return false;
     }
   }
@@ -1213,8 +1257,13 @@ bool MachineSMEABI::runOnMachineFunction(MachineFunction &MF) {
 
   FunctionInfo FnInfo = collectNeededZAStates(SMEFnAttrs);
 
-  if (SMEFnAttrs.hasPrivateZAInterface() && canElidePrivateZASetup(FnInfo))
-    return false;
+  if (SMEFnAttrs.hasPrivateZAInterface()) {
+    if (canElidePrivateZASetup(FnInfo))
+      return false;
+
+    // If we couldn't elide ZA setup, try to push ZAState::OFF earlier.
+    peepholeOptimizeStateChanges(FnInfo);
+  }
 
   SmallVector<ZAState> BundleStates = assignBundleZAStates(Bundles, FnInfo);
 
diff --git a/llvm/test/CodeGen/AArch64/sme-lazy-save-call.ll b/llvm/test/CodeGen/AArch64/sme-lazy-save-call.ll
index 3037f516642c2a..4834f0020aa5d9 100644
--- a/llvm/test/CodeGen/AArch64/sme-lazy-save-call.ll
+++ b/llvm/test/CodeGen/AArch64/sme-lazy-save-call.ll
@@ -198,11 +198,8 @@ define void @test_lazy_save_mixed_shared_and_private_callees() "aarch64_new_za"
 ; CHECK-NEXT:    msr TPIDR2_EL0, xzr
 ; CHECK-NEXT:    bl shared_za_callee
 ; CHECK-NEXT:    bl preserves_za_callee
-; CHECK-NEXT:    sub x8, x29, #16
-; CHECK-NEXT:    msr TPIDR2_EL0, x8
-; CHECK-NEXT:    bl private_za_callee
-; CHECK-NEXT:    msr TPIDR2_EL0, xzr
 ; CHECK-NEXT:    smstop za
+; CHECK-NEXT:    bl private_za_callee
 ; CHECK-NEXT:    mov sp, x29
 ; CHECK-NEXT:    ldr x19, [sp, #16] // 8-byte Reload
 ; CHECK-NEXT:    ldp x29, x30, [sp], #32 // 16-byte Folded Reload
diff --git a/llvm/test/CodeGen/AArch64/sme-za-exceptions.ll b/llvm/test/CodeGen/AArch64/sme-za-exceptions.ll
index 6eb4de449aaa6c..6878e0e76583b7 100644
--- a/llvm/test/CodeGen/AArch64/sme-za-exceptions.ll
+++ b/llvm/test/CodeGen/AArch64/sme-za-exceptions.ll
@@ -263,11 +263,8 @@ define void @try_catch_shared_za_callee() "aarch64_new_za" personality ptr @__gx
 ; CHECK-NEXT:  .LBB2_6: // %catch
 ; CHECK-NEXT:    msr TPIDR2_EL0, xzr
 ; CHECK-NEXT:    bl noexcept_shared_za_call
-; CHECK-NEXT:    sub x8, x29, #16
-; CHECK-NEXT:    msr TPIDR2_EL0, x8
-; CHECK-NEXT:    bl __cxa_end_catch
-; CHECK-NEXT:    msr TPIDR2_EL0, xzr
 ; CHECK-NEXT:    smstop za
+; CHECK-NEXT:    bl __cxa_end_catch
 ; CHECK-NEXT:    b .LBB2_3
   invoke void @shared_za_call() #4
           to label %exit unwind label %catch
diff --git a/llvm/test/CodeGen/AArch64/sme-za-function-with-many-blocks.ll b/llvm/test/CodeGen/AArch64/sme-za-function-with-many-blocks.ll
index 01a1746866f4f3..27353b5f016418 100644
--- a/llvm/test/CodeGen/AArch64/sme-za-function-with-many-blocks.ll
+++ b/llvm/test/CodeGen/AArch64/sme-za-function-with-many-blocks.ll
@@ -21,8 +21,6 @@ define void @matmul(ptr %0, ptr %1, i64 %2, i64 %3, i64 %4, i64 %5, i64 %6, ptr
 ; CHECK-LABEL: matmul:
 ; CHECK:      zero {za}
 ; CHECK-NOT:  TPIDR2_EL0
-; CHECK:      msr TPIDR2_EL0, x{{.*}}
-; CHECK-NOT:  .LBB{{.*}}
 ; CHECK:      bl printMemrefF32
   %22 = insertvalue { ptr, ptr, i64, [2 x i64], [2 x i64] } poison, ptr %14, 0
   %23 = insertvalue { ptr, ptr, i64, [2 x i64], [2 x i64] } %22, ptr %15, 1
diff --git a/llvm/test/CodeGen/AArch64/sme-zt0-state.ll b/llvm/test/CodeGen/AArch64/sme-zt0-state.ll
index 788ca7d932b2cc..e24cc9c5eb0705 100644
--- a/llvm/test/CodeGen/AArch64/sme-zt0-state.ll
+++ b/llvm/test/CodeGen/AArch64/sme-zt0-state.ll
@@ -332,8 +332,7 @@ define void @za_zt0_private_za_to_shared_za(ptr %callee) "aarch64_inout_za" "aar
 define void @no_need_to_save_zt0(ptr %callee) "aarch64_new_za" "aarch64_new_zt0" nounwind {
 ; CHECK-LABEL: no_need_to_save_zt0:
 ; CHECK:       // %bb.0:
-; CHECK-NEXT:    sub sp, sp, #80
-; CHECK-NEXT:    str x30, [sp, #64] // 8-byte Spill
+; CHECK-NEXT:    str x30, [sp, #-16]! // 8-byte Folded Spill
 ; CHECK-NEXT:    mrs x8, TPIDR2_EL0
 ; CHECK-NEXT:    cbz x8, .LBB15_2
 ; CHECK-NEXT:  // %bb.1:
@@ -343,12 +342,9 @@ define void @no_need_to_save_zt0(ptr %callee) "aarch64_new_za" "aarch64_new_zt0"
 ; CHECK-NEXT:    zero { zt0 }
 ; CHECK-NEXT:  .LBB15_2:
 ; CHECK-NEXT:    smstart za
-; CHECK-NEXT:    mov x8, sp
-; CHECK-NEXT:    str zt0, [x8]
 ; CHECK-NEXT:    blr x0
 ; CHECK-NEXT:    smstop za
-; CHECK-NEXT:    ldr x30, [sp, #64] // 8-byte Reload
-; CHECK-NEXT:    add sp, sp, #80
+; CHECK-NEXT:    ldr x30, [sp], #16 // 8-byte Folded Reload
 ; CHECK-NEXT:    ret
   call void %callee() "aarch64_inout_za"
   ret void;
@@ -357,8 +353,7 @@ define void @no_need_to_save_zt0(ptr %callee) "aarch64_new_za" "aarch64_new_zt0"
 define void @no_need_to_save_zt0_after_call(ptr %callee) "aarch64_new_za" "aarch64_new_zt0" nounwind {
 ; CHECK-LABEL: no_need_to_save_zt0_after_call:
 ; CHECK:       // %bb.0:
-; CHECK-NEXT:    sub sp, sp, #80
-; CHECK-NEXT:    stp x30, x19, [sp, #64] // 16-byte Folded Spill
+; CHECK-NEXT:    stp x30, x19, [sp, #-16]! // 16-byte Folded Spill
 ; CHECK-NEXT:    mrs x8, TPIDR2_EL0
 ; CHECK-NEXT:    cbz x8, .LBB16_2
 ; CHECK-NEXT:  // %bb.1:
@@ -370,12 +365,9 @@ define void @no_need_to_save_zt0_after_call(ptr %callee) "aarch64_new_za" "aarch
 ; CHECK-NEXT:    smstart za
 ; CHECK-NEXT:    mov x19, x0
 ; CHECK-NEXT:    blr x0
-; CHECK-NEXT:    mov x8, sp
-; CHECK-NEXT:    str zt0, [x8]
 ; CHECK-NEXT:    blr x19
 ; CHECK-NEXT:    smstop za
-; CHECK-NEXT:    ldp x30, x19, [sp, #64] // 16-byte Folded Reload
-; CHECK-NEXT:    add sp, sp, #80
+; CHECK-NEXT:    ldp x30, x19, [sp], #16 // 16-byte Folded Reload
 ; CHECK-NEXT:    ret
   call void %callee() "aarch64_inout_za" "aarch64_inout_zt0"
   call void %callee() "aarch64_inout_za"



More information about the llvm-commits mailing list