diff --git a/src/kernbench/components/builtin/pe_scheduler.py b/src/kernbench/components/builtin/pe_scheduler.py index b72f4e8..40a2780 100644 --- a/src/kernbench/components/builtin/pe_scheduler.py +++ b/src/kernbench/components/builtin/pe_scheduler.py @@ -23,6 +23,53 @@ if TYPE_CHECKING: from kernbench.topology.types import Node +class _RwHazardTracker: + """Strict-FIFO cross-composite RW hazard tracker (ADR-0065 D6.3 / DDD §8). + + A composite is registered as in-flight by its *write* set + (``rw_handles``). A new composite whose read **or** write handles + intersect any in-flight write set waits on those composites' done events + before being admitted. Composites with no ``rw_handles`` (legacy) are + never registered and never block — existing benches are unaffected. + + Strict FIFO is conservative: even partial overlap waits for the prior + overlapping composite to fully drain (RW-aware reorder is ADR-0065 A4, + deferred). + """ + + def __init__(self) -> None: + # (completion_id, write-handle id set, done_event) + self._inflight: list[tuple[str, set, Any]] = [] + + @staticmethod + def _conflict_ids(cmd: Any) -> set: + """Handle ids this composite reads (op operands) or writes (rw).""" + ids = {h.id for h in cmd.rw_handles} + for op in cmd.ops: + for h in op.operands.values(): + hid = getattr(h, "id", None) + if hid is not None: + ids.add(hid) + return ids + + def admit(self, env: Any, cmd: Any, done_event: Any) -> Generator: + """Block until no in-flight composite's write set overlaps this + composite's read/write set, then register this composite's writes.""" + cids = self._conflict_ids(cmd) + while True: + blocking = [ev for _cid, rw, ev in self._inflight if rw & cids] + if not blocking: + break + yield blocking[0] + if cmd.rw_handles: + self._inflight.append( + (cmd.completion.id, {h.id for h in cmd.rw_handles}, done_event) + ) + + def retire(self, completion_id: str) -> None: + self._inflight = [t for t in self._inflight if t[0] != completion_id] + + class PeSchedulerComponent(ComponentBase): """PE_SCHEDULER: sole dispatcher inside a PE (ADR-0014 D1, D6). @@ -60,6 +107,7 @@ class PeSchedulerComponent(ComponentBase): self._ensure_dispatch_table() self._pending_feeds: simpy.Store | None = None self._pipeline_counter = 0 + self._hazard = _RwHazardTracker() # ADR-0065 D6.3 strict-FIFO RW def start(self, env: simpy.Environment) -> None: self._pending_feeds = simpy.Store(env) @@ -108,9 +156,17 @@ class PeSchedulerComponent(ComponentBase): def _dispatch_composite( self, env: simpy.Environment, pe_txn: Any, cmd: Any, ) -> Generator: - """Generate plan and enqueue to feeder. Non-blocking (ADR-0014 D6).""" + """Generate plan and enqueue to feeder. Non-blocking (ADR-0014 D6). + + Strict-FIFO RW gate (ADR-0065 D6.3): wait for overlapping in-flight + composites, register this one's write set, retire on completion. + Legacy composites (no rw_handles) pass through immediately. + """ from kernbench.components.builtin.pe_types import PipelineContext + yield from self._hazard.admit(env, cmd, pe_txn.done) + env.process(self._retire_on_done(env, pe_txn.done, cmd.completion.id)) + plan = self._generate_plan(cmd) self._pipeline_counter += 1 @@ -124,6 +180,13 @@ class PeSchedulerComponent(ComponentBase): assert self._pending_feeds is not None yield self._pending_feeds.put((plan, ctx)) + def _retire_on_done( + self, env: simpy.Environment, done_event: Any, completion_id: str, + ) -> Generator: + """Retire a composite from the RW hazard tracker on completion.""" + yield done_event + self._hazard.retire(completion_id) + def _feed_loop(self, env: simpy.Environment) -> Generator: """Single feeder process: FIFO command ordering (ADR-0014 D6). diff --git a/tests/test_pe_scheduler_strict_fifo.py b/tests/test_pe_scheduler_strict_fifo.py new file mode 100644 index 0000000..797cdff --- /dev/null +++ b/tests/test_pe_scheduler_strict_fifo.py @@ -0,0 +1,153 @@ +"""Phase 1 spec tests for ADR-0065 P4 — strict-FIFO RW hazard tracker. + +PE_SCHEDULER serializes composites that share read/write handles +(ADR-0065 D6.3 / DDD §8): a new composite whose read or write handles +intersect an in-flight composite's ``rw_handles`` waits until all earlier +overlapping composites complete. Non-overlapping composites proceed. + +These unit-test the ``_RwHazardTracker`` directly with a SimPy env. Legacy +composites (``rw_handles == ()``) never register and never block, so the +existing benches' Stage sequences are untouched. + +Phase 1 (this commit): tests only. FAIL until P4: + - ``_RwHazardTracker`` does not exist in ``pe_scheduler``. +""" +from __future__ import annotations + +import simpy + +from kernbench.common.pe_commands import ( + CompletionHandle, + CompositeCmd, + OpSpec, + Scope, + TensorHandle, +) + + +def _h(addr: int) -> TensorHandle: + return TensorHandle(id=f"h{addr}", addr=addr, shape=(8,), dtype="f16", + nbytes=16, space="tcm") + + +def _cmd(cid: str, rw=(), reads=()) -> CompositeCmd: + """A composite that writes ``rw`` (rw_handles) and reads ``reads`` + (as MATH-op operands).""" + if reads: + ops = tuple( + OpSpec(kind="math", scope=Scope.KERNEL, operands={"src": r}, out=r) + for r in reads + ) + else: + ops = (OpSpec(kind="math", scope=Scope.KERNEL, operands={}, out=None),) + return CompositeCmd(completion=CompletionHandle(id=cid), ops=ops, + rw_handles=tuple(rw)) + + +def _tracker(): + from kernbench.components.builtin.pe_scheduler import _RwHazardTracker + + return _RwHazardTracker() + + +def test_admit_immediate_when_no_inflight(): + env = simpy.Environment() + tr = _tracker() + state: dict = {} + + def driver(): + yield from tr.admit(env, _cmd("a", rw=(_h(1),)), env.event()) + state["t"] = env.now + + env.process(driver()) + env.run() + assert state["t"] == 0 + + +def test_admit_blocks_until_overlapping_inflight_retires(): + env = simpy.Environment() + tr = _tracker() + O = _h(1) + eA = env.event() + state: dict = {} + + def b_proc(): + # B reads O, which in-flight A writes (RAW) → must block. + yield from tr.admit(env, _cmd("b", reads=(O,)), env.event()) + state["B"] = env.now + + def driver(): + yield from tr.admit(env, _cmd("a", rw=(O,)), eA) + state["A"] = env.now + env.process(b_proc()) + yield env.timeout(10) + tr.retire("a") + eA.succeed() # A completes at t=10 + yield env.timeout(1) + + env.process(driver()) + env.run() + assert state["A"] == 0 + assert state["B"] >= 10, f"B admitted at {state.get('B')}, must wait for A" + + +def test_admit_no_block_when_disjoint(): + env = simpy.Environment() + tr = _tracker() + O, X = _h(1), _h(2) + state: dict = {} + + def driver(): + yield from tr.admit(env, _cmd("a", rw=(O,)), env.event()) + yield from tr.admit(env, _cmd("b", reads=(X,)), env.event()) + state["B"] = env.now + + env.process(driver()) + env.run() + assert state["B"] == 0, "disjoint composite must not block" + + +def test_write_after_write_overlap_blocks(): + """Two composites both writing the same handle (WAW) serialize.""" + env = simpy.Environment() + tr = _tracker() + O = _h(1) + eA = env.event() + state: dict = {} + + def b_proc(): + yield from tr.admit(env, _cmd("b", rw=(O,)), env.event()) + state["B"] = env.now + + def driver(): + yield from tr.admit(env, _cmd("a", rw=(O,)), eA) + env.process(b_proc()) + yield env.timeout(5) + tr.retire("a") + eA.succeed() + yield env.timeout(1) + + env.process(driver()) + env.run() + assert state["B"] >= 5 + + +def test_legacy_composite_never_registers(): + """A composite with empty rw_handles must not block a later composite + that touches the same input handles (legacy composites carry no RW + metadata → no serialization, byte-equal behavior).""" + env = simpy.Environment() + tr = _tracker() + O = _h(1) + state: dict = {} + + def driver(): + # A is "legacy": reads O but declares no rw_handles → not registered. + yield from tr.admit(env, _cmd("a", reads=(O,)), env.event()) + # B reads O too; since A registered nothing, B must not block. + yield from tr.admit(env, _cmd("b", reads=(O,)), env.event()) + state["B"] = env.now + + env.process(driver()) + env.run() + assert state["B"] == 0