gqa(adr-0065): P4 — strict-FIFO RW hazard tracker
_RwHazardTracker (pe_scheduler): a composite registers its write set (rw_handles) as in-flight; a new composite whose read (op operands) or write handles intersect an in-flight write set waits on those composites' done events before admission (ADR-0065 D6.3 / DDD 8). _dispatch_composite calls admit() before feeding and retires on the done event. Legacy composites (rw_handles=()) never register and never block -> existing benches untouched (byte-equal). Strict FIFO is conservative (partial overlap waits for full drain); RW-aware reorder is deferred (A4). Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -23,6 +23,53 @@ if TYPE_CHECKING:
|
|||||||
from kernbench.topology.types import Node
|
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):
|
class PeSchedulerComponent(ComponentBase):
|
||||||
"""PE_SCHEDULER: sole dispatcher inside a PE (ADR-0014 D1, D6).
|
"""PE_SCHEDULER: sole dispatcher inside a PE (ADR-0014 D1, D6).
|
||||||
|
|
||||||
@@ -60,6 +107,7 @@ class PeSchedulerComponent(ComponentBase):
|
|||||||
self._ensure_dispatch_table()
|
self._ensure_dispatch_table()
|
||||||
self._pending_feeds: simpy.Store | None = None
|
self._pending_feeds: simpy.Store | None = None
|
||||||
self._pipeline_counter = 0
|
self._pipeline_counter = 0
|
||||||
|
self._hazard = _RwHazardTracker() # ADR-0065 D6.3 strict-FIFO RW
|
||||||
|
|
||||||
def start(self, env: simpy.Environment) -> None:
|
def start(self, env: simpy.Environment) -> None:
|
||||||
self._pending_feeds = simpy.Store(env)
|
self._pending_feeds = simpy.Store(env)
|
||||||
@@ -108,9 +156,17 @@ class PeSchedulerComponent(ComponentBase):
|
|||||||
def _dispatch_composite(
|
def _dispatch_composite(
|
||||||
self, env: simpy.Environment, pe_txn: Any, cmd: Any,
|
self, env: simpy.Environment, pe_txn: Any, cmd: Any,
|
||||||
) -> Generator:
|
) -> 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
|
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)
|
plan = self._generate_plan(cmd)
|
||||||
|
|
||||||
self._pipeline_counter += 1
|
self._pipeline_counter += 1
|
||||||
@@ -124,6 +180,13 @@ class PeSchedulerComponent(ComponentBase):
|
|||||||
assert self._pending_feeds is not None
|
assert self._pending_feeds is not None
|
||||||
yield self._pending_feeds.put((plan, ctx))
|
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:
|
def _feed_loop(self, env: simpy.Environment) -> Generator:
|
||||||
"""Single feeder process: FIFO command ordering (ADR-0014 D6).
|
"""Single feeder process: FIFO command ordering (ADR-0014 D6).
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
Reference in New Issue
Block a user