From 526eb8bb48bfe5845238dce2629f9bfb83552134 Mon Sep 17 00:00:00 2001 From: yanghaoran29 Date: Wed, 29 Jul 2026 16:49:35 +0800 Subject: [PATCH] Update: batch A5 scheduler progress publication Publish last_task_alive every 16 local advances while work remains, and force the drained tail so reclamation cannot stall on a partial batch. Reset the publication shadow on ring reuse and cover threshold and drain behavior in A5 unit tests. --- .../runtime/scheduler/pto_scheduler.h | 13 ++++- .../runtime/shared/pto_runtime2_init.cpp | 2 + tests/ut/cpp/a5/test_wiring.cpp | 50 +++++++++++++++++++ 3 files changed, 63 insertions(+), 2 deletions(-) diff --git a/src/a5/runtime/tensormap_and_ringbuffer/runtime/scheduler/pto_scheduler.h b/src/a5/runtime/tensormap_and_ringbuffer/runtime/scheduler/pto_scheduler.h index d5b74aa46..572dc2807 100644 --- a/src/a5/runtime/tensormap_and_ringbuffer/runtime/scheduler/pto_scheduler.h +++ b/src/a5/runtime/tensormap_and_ringbuffer/runtime/scheduler/pto_scheduler.h @@ -442,6 +442,9 @@ struct PTO2SchedulerState { // --- Cache Line 0: ring pointer (read-only) + hot path (read-write) --- PTO2SharedMemoryRingHeader *ring; int32_t last_task_alive; + // Shared-memory publication trails local reclamation by at most 15 + // live tasks while the ring has pending work. + int32_t last_published_to_sm{0}; std::atomic advance_lock; // multi-thread CAS // --- Cache Line 1+: Orch-side wiring dep_pool --- @@ -461,7 +464,13 @@ struct PTO2SchedulerState { void reset_for_reuse(void *sm_dev_base, int32_t ring_id, std::atomic *orch_err); void destroy(); - void sync_to_sm() { ring->fc.last_task_alive.store(last_task_alive, std::memory_order_release); } + void sync_to_sm(bool force = false) { + constexpr int32_t PUBLISH_INTERVAL_K = 16; + if (last_task_alive == last_published_to_sm) return; + if (!force && last_task_alive - last_published_to_sm < PUBLISH_INTERVAL_K) return; + ring->fc.last_task_alive.store(last_task_alive, std::memory_order_release); + last_published_to_sm = last_task_alive; + } #if SIMPLER_DFX void publish_dep_pool_snapshot() { @@ -497,7 +506,7 @@ struct PTO2SchedulerState { ring->get_slot_state_by_task_id(id).reset_for_reuse(); } - sync_to_sm(); + sync_to_sm(last_task_alive == current_task_index); } } ring_sched_states[PTO2_MAX_RING_DEPTH]; diff --git a/src/a5/runtime/tensormap_and_ringbuffer/runtime/shared/pto_runtime2_init.cpp b/src/a5/runtime/tensormap_and_ringbuffer/runtime/shared/pto_runtime2_init.cpp index e9a3d8a2e..d3530dcd2 100644 --- a/src/a5/runtime/tensormap_and_ringbuffer/runtime/shared/pto_runtime2_init.cpp +++ b/src/a5/runtime/tensormap_and_ringbuffer/runtime/shared/pto_runtime2_init.cpp @@ -90,6 +90,7 @@ bool PTO2SchedulerState::RingSchedState::init_data_from_layout(void *sm_dev_base // arithmetic, no SM load. ring = pto2_sm_layout::ring_header_addr(sm_dev_base, ring_id); last_task_alive = 0; + last_published_to_sm = 0; advance_lock.store(0, std::memory_order_relaxed); #if SIMPLER_DFX dep_pool_snapshot_tail.store(1, std::memory_order_relaxed); @@ -109,6 +110,7 @@ void PTO2SchedulerState::RingSchedState::reset_for_reuse( ) { ring = pto2_sm_layout::ring_header_addr(sm_dev_base, ring_id); last_task_alive = 0; + last_published_to_sm = 0; advance_lock.store(0, std::memory_order_relaxed); dep_pool.reset_for_reuse(orch_err); #if SIMPLER_DFX diff --git a/tests/ut/cpp/a5/test_wiring.cpp b/tests/ut/cpp/a5/test_wiring.cpp index 24b25a663..adc531a90 100644 --- a/tests/ut/cpp/a5/test_wiring.cpp +++ b/tests/ut/cpp/a5/test_wiring.cpp @@ -375,6 +375,56 @@ TEST_F(WiringTest, AdvanceRingPointersScansConsumed) { EXPECT_EQ(ring->fc.last_task_alive.load(), 3); } +TEST_F(WiringTest, AdvanceRingPointersBatchesSharedMemoryPublication) { + auto &rss = sched.ring_sched_states[0]; + auto *ring = rss.ring; + + ring->fc.current_task_index.store(18, std::memory_order_release); + for (int i = 0; i < 17; i++) { + ring->get_slot_state_by_task_id(i).task_state.store(PTO2_TASK_CONSUMED); + } + + rss.advance_ring_pointers(); + EXPECT_EQ(rss.last_task_alive, 17); + EXPECT_EQ(ring->fc.last_task_alive.load(), 17); + + ring->fc.last_task_alive.store(0); + rss.last_task_alive = 0; + rss.last_published_to_sm = 0; + for (int advances : {1, 15, 16, 17}) { + for (int i = 0; i < advances; i++) { + ring->get_slot_state_by_task_id(i).task_state.store(PTO2_TASK_CONSUMED); + } + ring->get_slot_state_by_task_id(advances).task_state.store(PTO2_TASK_COMPLETED); + rss.advance_ring_pointers(); + EXPECT_EQ(ring->fc.last_task_alive.load(), advances >= 16 ? advances : 0); + + ring->fc.last_task_alive.store(0); + rss.last_task_alive = 0; + rss.last_published_to_sm = 0; + } +} + +TEST_F(WiringTest, AdvanceRingPointersPublishesDrainedTail) { + auto &rss = sched.ring_sched_states[0]; + auto *ring = rss.ring; + + for (int advances : {1, 15, 16, 17}) { + ring->fc.current_task_index.store(advances, std::memory_order_release); + for (int i = 0; i < advances; i++) { + ring->get_slot_state_by_task_id(i).task_state.store(PTO2_TASK_CONSUMED); + } + + rss.advance_ring_pointers(); + EXPECT_EQ(rss.last_task_alive, advances); + EXPECT_EQ(ring->fc.last_task_alive.load(), advances); + + ring->fc.last_task_alive.store(0); + rss.last_task_alive = 0; + rss.last_published_to_sm = 0; + } +} + TEST_F(WiringTest, AdvanceRingPointersStopsAtNonConsumed) { auto &rss = sched.ring_sched_states[0]; auto *ring = rss.ring;