Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 5 additions & 2 deletions native/csrc/bindings.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -764,11 +764,14 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
.def("staging_cap", &ring_py::RingEnginePy::staging_cap)
.def("task_cap", &ring_py::RingEnginePy::task_cap)
.def("payload_tensor", &ring_py::RingEnginePy::payload_tensor)
// Safety-net surface (eager only). available_capacity() and
// reserve_one() are CPU-only and fast -- no GIL release needed.
// Safety-net surface (eager only). available_capacity(),
// available_task_slots() and reserve_one() are CPU-only and fast --
// no GIL release needed.
// flush_and_wait() blocks on cudaStreamSynchronize + drain flush --
// GIL released so other Python threads aren't blocked.
.def("available_capacity", &ring_py::RingEnginePy::available_capacity)
.def("available_task_slots",
&ring_py::RingEnginePy::available_task_slots)
.def("reserve_one",
&ring_py::RingEnginePy::reserve_one,
py::arg("nbytes"))
Expand Down
38 changes: 29 additions & 9 deletions native/csrc/ring/ring_engine_py.cu
Original file line number Diff line number Diff line change
Expand Up @@ -631,22 +631,29 @@ at::Tensor RingEnginePy::payload_tensor() const {
//
// Thread safety of the check-and-reserve pattern used by the safety net:
//
// if nbytes <= available_capacity():
// if nbytes <= available_capacity() and available_task_slots() > 0:
// reserve_one(nbytes)
//
// Both halves are needed. A step with more hooks than task entries gets
// STEP_OVERSIZED and runs through the safety net with plenty of payload
// room, so a bytes-only check lets the (task_cap+1)-th producer publish
// over slot 0 while its READY word is still unread.
//
// The main thread (this thread) is the only writer of cpu_payload_head_
// (it advances only through reserve / reserve_one calls). The drain
// thread only ever advances cpu_payload_tail_committed_ forward as it
// frees ring space. Between the check and the reserve:
// and cpu_task_head_ (they advance only through reserve / reserve_one
// calls). The drain thread only ever advances the committed tails
// forward as it frees ring space. Between the check and the reserve:
// - tail may move forward (drain freed more): actual available at
// reserve time is >= what we observed.
// - head is unchanged (single-threaded writer).
// So the check's "fits" decision remains valid at reserve time. No extra
// locking around the pair is required.
// locking around the pair is required. The same holds for the task
// head and tail.
//
// Within available_capacity(), the two accessor calls happen under
// separate mutex acquires (drain.cpu_payload_head() and
// drain.cpu_payload_tail_committed() each take mgmt_mu_ internally).
// Within available_capacity() (and likewise available_task_slots()), the
// two accessor calls happen under separate mutex acquires
// (drain.cpu_payload_head() and drain.cpu_payload_tail_committed() each
// take mgmt_mu_ internally).
// The observed snapshot is non-atomic: if drain advances tail between
// the two reads, available_observed = pcap - head + tail_later, which
// is >= the true available at the time of the head read. That is, the
Expand All @@ -661,11 +668,24 @@ uint64_t RingEnginePy::available_capacity() const {
return pcap - (drain.cpu_payload_head() - drain.cpu_payload_tail_committed());
}

uint64_t RingEnginePy::available_task_slots() const {
auto& drain = impl_->engine.drain_thread();
const uint64_t tcap = impl_->engine.task_cap();
return tcap - (drain.cpu_task_head() - drain.cpu_task_tail_committed());
}

// Per-hook reservation: claim nbytes of payload + 1 task entry for an
// upcoming producer kernel launch. Caller must have checked
// available_capacity() first. drain.reserve takes mgmt_mu_ internally.
// available_capacity() and available_task_slots() first. The task check
// is repeated here because producers never read the tails: an entry past
// task_cap silently overwrites an unread slot rather than failing.
// drain.reserve takes mgmt_mu_ internally.
void RingEnginePy::reserve_one(uint64_t nbytes) {
refuse_on_record_ring(impl_->record_mode, "legacy per-hook reservation");
if (available_task_slots() == 0) {
throw std::logic_error(
"reserve_one: no free task-ring entry; flush_and_wait first");
}
impl_->engine.drain_thread().reserve(
ring::align_up(nbytes, ring::PAYLOAD_ALIGN), 1);
}
Expand Down
8 changes: 7 additions & 1 deletion native/csrc/ring/ring_engine_py.h
Original file line number Diff line number Diff line change
Expand Up @@ -207,10 +207,16 @@ class RingEnginePy {
// pending drain. CPU-only read.
uint64_t available_capacity() const;

// Free task-ring entries not currently reserved and not pending
// drain. CPU-only read.
uint64_t available_task_slots() const;

// Per-hook reservation: claim `nbytes` of payload ring + 1 task entry
// for an upcoming producer kernel launch. Used by the safety net
// when force_eager is on and the spec is dynamic-shape. Advances
// cpu_payload_head/cpu_task_head atomically.
// cpu_payload_head/cpu_task_head atomically. Throws std::logic_error,
// reserving nothing, when no task entry is free: the producer would
// overwrite an unconsumed READY word. Callers flush_and_wait() first.
void reserve_one(uint64_t nbytes);

// Synchronise the current CUDA stream + force drain to process all
Expand Down
3 changes: 2 additions & 1 deletion src/dmi/configuration/estimate.py
Original file line number Diff line number Diff line change
Expand Up @@ -818,7 +818,8 @@ def check_ring_fit(
f"The busiest rank fires {over_task_cap} hooks in one step, "
f"above the {task_entries} task entries the ring was configured "
"with. prepare_step returns STEP_OVERSIZED regardless of bytes "
"and the adapter falls back to eager CPU-direct dispatch -- "
"and the adapter falls back to eager per-hook dispatch, which "
"syncs and flushes the ring each time its task entries fill -- "
"capture keeps working, but the serving path pays for it. "
"Raise ring task entries, narrow the layer range, or deselect "
"observations."
Expand Down
9 changes: 7 additions & 2 deletions src/dmi/hooks/point.py
Original file line number Diff line number Diff line change
Expand Up @@ -317,9 +317,14 @@ def forward(self, x: Tensor) -> Tensor:
# min is cached on the transport and shared by every hook --
# querying the pair per hook would repeat it for the whole
# active set on the first eager forward.
# A free task entry is needed too: a step with more hooks
# than task entries lands here with the payload ring nearly
# empty, and reserving past task_cap makes a producer
# overwrite an unread slot. A flush frees both.
effective_cap = transport.effective_cap
if transport_bytes <= min(engine.available_capacity(),
effective_cap):
if (transport_bytes <= min(engine.available_capacity(),
effective_cap)
and engine.available_task_slots() > 0):
engine.reserve_one(nbytes)
dispatch_producer(ring_payload, x_cont, strip_t, strip_rb,
self._ring_hook_type, self._ring_hook_id)
Expand Down
3 changes: 2 additions & 1 deletion src/dmi/transport/ring.py
Original file line number Diff line number Diff line change
Expand Up @@ -226,7 +226,8 @@ def __init__(self, ring_engine: Any) -> None:

# When True, HookPoint.forward takes the runtime safety-net branch
# instead of the fast path:
# 1. fits in current slack -> reserve_one + ring
# 1. fits in current slack (bytes AND a free task entry)
# -> reserve_one + ring
# 2. fits after flushing the ring -> flush_and_wait + reserve_one + ring
# 3. single tensor > ring -> flush_and_wait + submit_cpu_direct
# Owned by adaptor_base.before_forward (per-batch reassignment based
Expand Down
29 changes: 29 additions & 0 deletions tests/native/ring/test_ring_engine.cu
Original file line number Diff line number Diff line change
Expand Up @@ -224,6 +224,34 @@ static void test_native_reservation_uses_transport_alignment() {
EXPECT(before - engine.available_capacity() == 32);
}

static void test_reserve_one_refuses_without_a_free_task_slot() {
banner("reserve_one refuses once every task entry is reserved");
// Producers never read the tails, so a reservation past task_cap would
// publish over an unread READY word. Payload room is plentiful here:
// the task ring is the only limit, as for a STEP_OVERSIZED step whose
// hook count exceeds task_ring_entries.
ring_py::RingEnginePy engine(make_py_config(), ring_py::SubmitFn{});
engine.init();
const uint64_t task_cap = engine.task_cap();
EXPECT(engine.available_task_slots() == task_cap);
for (uint64_t i = 0; i < task_cap; ++i) {
engine.reserve_one(16);
}
EXPECT(engine.available_task_slots() == 0);

const uint64_t bytes = engine.available_capacity();
EXPECT(bytes >= 16);
bool refused = false;
try {
engine.reserve_one(16);
} catch (const std::logic_error& error) {
refused = std::strstr(error.what(), "task-ring entry") != nullptr;
}
EXPECT(refused);
EXPECT(engine.available_task_slots() == 0);
EXPECT(engine.available_capacity() == bytes);
}

template <typename Fn>
static bool refuses_as_legacy_on_record_ring(Fn&& call) {
try {
Expand Down Expand Up @@ -784,6 +812,7 @@ int main() {
std::printf("test_ring_engine (current drain pipeline)\n");
test_ring_geometry_requires_payload_alignment();
test_native_reservation_uses_transport_alignment();
test_reserve_one_refuses_without_a_free_task_slot();
test_record_ring_refuses_every_legacy_producer_entry();
test_static_force_flush();
test_prefix_force_flush();
Expand Down
3 changes: 3 additions & 0 deletions tests/test_hook_point_eager_cap_cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,6 +37,9 @@ def payload_tensor(self) -> torch.Tensor:
def available_capacity(self) -> int:
return self.available

def available_task_slots(self) -> int:
return 1 # task capacity is not under test here

def payload_cap(self) -> int:
self.payload_cap_calls += 1
return self.capacity
Expand Down
108 changes: 108 additions & 0 deletions tests/test_hook_point_eager_task_slots.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,108 @@
"""The eager safety net must not reserve more tasks than the task ring holds.

A legacy step with more firing hooks than ``task_ring_entries`` gets
``STEP_OVERSIZED`` from ``prepare_step`` (after a flush, so the ring is
empty), and the adapter sets ``force_eager``. Each hook then took the safety
net in ``HookPoint.forward``, which checked only payload BYTES before
``reserve_one`` -- and ``reserve_one`` advanced the task head with no
task-capacity check. With plenty of payload room, hook ``task_cap`` reserved
sequence ``task_cap``, whose producer release-stores into slot 0 while hook
0's READY word there is still unread (producers never read the tails). The
drain then either pairs the wrong size with hook 0's TensorMeta or clears the
new word and waits at that sequence forever while ``flush_and_wait`` reports
success. Found by TLA+ model checking of the payload ring.

The fake below keeps a legacy ring's head/tail counters. ``HookPoint.forward``
only takes the eager branch for a CUDA tensor, so a CPU tensor subclass poses
as one (as in tests/test_record_runtime.py) and ``dispatch_producer`` is
monkeypatched; the transport is a real ``RingTransport``.
"""
from __future__ import annotations

import pytest
import torch

from dmi.hooks.specs import align_up_py

pytestmark = pytest.mark.cpu


class _FakeCudaTensor(torch.Tensor):
@property
def is_cuda(self) -> bool:
return True


class _LegacyRingEngine:
"""Payload and task heads/tails of a legacy ring; flush drains both."""

def __init__(self, task_cap: int, payload_cap: int):
self.task_cap = task_cap
self.capacity = payload_cap
self.task_head = self.task_tail = 0
self.payload_head = self.payload_tail = 0
self.events: list[str] = []
self.max_outstanding_tasks = 0

def payload_tensor(self) -> torch.Tensor:
return torch.empty(64, dtype=torch.uint8)

def payload_cap(self) -> int:
return self.capacity

def staging_cap(self) -> int:
return self.capacity

def available_capacity(self) -> int:
return self.capacity - (self.payload_head - self.payload_tail)

def available_task_slots(self) -> int:
return self.task_cap - (self.task_head - self.task_tail)

def reserve_one(self, nbytes: int) -> None:
self.events.append("reserve")
self.payload_head += align_up_py(nbytes, 16)
self.task_head += 1
self.max_outstanding_tasks = max(
self.max_outstanding_tasks, self.task_head - self.task_tail)

def flush_and_wait(self) -> None:
self.events.append("flush")
self.task_tail = self.task_head
self.payload_tail = self.payload_head


def _eager_hook(monkeypatch, engine, dispatched):
from dmi.hooks.point import HookPoint
from dmi.transport import ring as ring_transport
from dmi.transport.ring import RingTransport

transport = RingTransport(engine)
transport.force_eager = True
monkeypatch.setattr(ring_transport, "_active_transport", transport)
monkeypatch.setattr(
"dmi.hooks.point.dispatch_producer",
lambda *args: dispatched.append(args),
)
hook = HookPoint()
hook._ring_hook_type = 1
hook._ring_hook_id = 2
hook._ring_payload = transport._ring_payload
return hook


def test_a_step_with_more_hooks_than_task_entries_flushes_before_overflowing(
monkeypatch):
task_cap = 4
engine = _LegacyRingEngine(task_cap=task_cap, payload_cap=1 << 20)
dispatched = []
hook = _eager_hook(monkeypatch, engine, dispatched)

value = torch.arange(32, dtype=torch.uint8).as_subclass(_FakeCudaTensor)
for _ in range(task_cap + 1):
hook(value)

assert len(dispatched) == task_cap + 1, "every hook still takes the ring"
assert engine.max_outstanding_tasks <= task_cap, (
"a reservation past task_cap overwrites an unread task slot")
assert engine.events == ["reserve"] * task_cap + ["flush", "reserve"]
3 changes: 3 additions & 0 deletions tests/test_producer_chunked_schema.py
Original file line number Diff line number Diff line change
Expand Up @@ -172,6 +172,9 @@ def payload_tensor(self) -> torch.Tensor:
def available_capacity(self) -> int:
return self.available

def available_task_slots(self) -> int:
return 1 # task capacity is not under test here

def payload_cap(self) -> int:
return self.capacity

Expand Down
Loading