gqa(adr-0064/0065): flat-ops CompositeCmd (P1) + structural dispatch cost (ADR-0064 Rev2); promote ADR-0064

ADR-0065 P1: CompositeCmd -> flat ordered ops list (drop legacy op/a/b/out_addr fields); OpSpec.operands dict + out handle. Meaning-preserving (op_log byte-equal); pe_scheduler + op_log read the head op.

ADR-0064 Rev2: replace Rev1 per-op cost table with structural FIXED + logical_bytes*R formula. logical_bytes on every PeCommand; new common/pe_cost_model.py; cost centralized in TLContext._emit (load/recv_async charge explicitly); pe_cpu/kernel_runner wire the per-PE model + clock. D7: cap exceeded -> ValueError (no auto-segmentation). Remove Rev1 cpu_issue_cost.py + its tests. No goldens churn.

Promote ADR-0064 Rev2 Proposed->Accepted (docs/adr/ + docs/adr-ko/); amend D7 (error not segmentation) + record P1-before-P0 ordering in ADR-0064/0065 Migration notes (EN+KO).

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
2026-06-10 19:18:04 -07:00
parent 79ddb12b42
commit 47e2c78c66
18 changed files with 703 additions and 550 deletions
+4 -2
View File
@@ -56,12 +56,14 @@ def test_composite_epilogue_roundtrip():
("scale", Scope.OUTPUT_TILE),
]
# Single-op call (no epilogue) keeps the legacy code path: ops stays empty.
# Single-op call (no epilogue): flat-ops still emits one head op
# (ADR-0065 D1 — every composite carries a head op, no legacy empty path).
tl2 = TLContext()
tl2.composite(op="gemm", a=a, b=b, out_ptr=0x2000)
cmd2 = tl2._commands[-1]
assert isinstance(cmd2, CompositeCmd)
assert cmd2.ops == ()
assert len(cmd2.ops) == 1
assert cmd2.ops[0].kind == "gemm"
@pytest.mark.parametrize("bad,match", [
+127
View File
@@ -0,0 +1,127 @@
"""Phase 1 spec tests for ADR-0065 P1 — flat-ops ``CompositeCmd`` refactor.
P1 restructures ``CompositeCmd`` from the legacy
``(op, a, b, out_addr, out_nbytes, math_op, ops)`` shape into a flat ordered
op list ``(completion, ops, rw_handles, data_op)`` (ADR-0065 D1/D2). The
user-facing ``tl.composite(op=..., a, b, epilogue=[...])`` API is preserved
(D6.4) and lowers internally to the flat shape.
This is a *meaning-preserving* refactor: every existing bench's op_log must
stay byte-equal (the existing integration suite is the regression net). The
tests here pin the new lowering *shape* only.
Phase 1 (this commit): tests only. All FAIL until the P2 production refactor:
- current ``CompositeCmd`` requires ``op/a/b/out_addr`` (no flat-only ctor),
- current ``OpSpec.operands`` is a tuple and there is no ``out`` field.
"""
from __future__ import annotations
from kernbench.common.pe_commands import CompositeCmd, OpSpec, Scope, TensorHandle
from kernbench.triton_emu.tl_context import TLContext
_LEGACY_FIELDS = ("op", "a", "b", "out_addr", "out_nbytes", "math_op")
def _tl() -> TLContext:
return TLContext(pe_id=0, num_programs=1, dispatch_cycles=0)
def _composites(tl: TLContext) -> list[CompositeCmd]:
return [c for c in tl.commands if isinstance(c, CompositeCmd)]
# ── plain GEMM composite (no epilogue) ───────────────────────────────
def test_plain_gemm_composite_lowers_to_single_head_op():
"""A composite with no epilogue still emits exactly one flat head op."""
tl = _tl()
a = tl.ref(0x1000, shape=(8, 16), dtype="f16")
b = tl.ref(0x2000, shape=(16, 8), dtype="f16")
tl.composite("gemm", a, b, out_ptr=0x3000)
comps = _composites(tl)
assert len(comps) == 1, f"expected 1 composite; got {len(comps)}"
cmd = comps[0]
assert len(cmd.ops) == 1, f"plain gemm must have 1 head op; got {len(cmd.ops)}"
head = cmd.ops[0]
assert head.kind == "gemm"
assert head.operands == {"a": a, "b": b}, head.operands
assert head.out is not None and head.out.addr == 0x3000, head.out
assert cmd.rw_handles == ()
def test_composite_has_no_legacy_fields():
"""Legacy ``op/a/b/out_addr/out_nbytes/math_op`` are gone from the cmd."""
tl = _tl()
a = tl.ref(0x1000, shape=(8, 16), dtype="f16")
b = tl.ref(0x2000, shape=(16, 8), dtype="f16")
tl.composite("gemm", a, b, out_ptr=0x3000)
cmd = _composites(tl)[0]
for f in _LEGACY_FIELDS:
assert not hasattr(cmd, f), f"legacy field {f!r} must be removed"
# ── GEMM + epilogue ──────────────────────────────────────────────────
def test_gemm_composite_with_epilogue_lowers_to_flat_ops():
"""``epilogue=[relu, bias]`` lowers to flat ops after the head, with the
epilogue operand migrated from positional tuple to a named dict."""
tl = _tl()
a = tl.ref(0x1000, shape=(8, 16), dtype="f16")
b = tl.ref(0x2000, shape=(16, 8), dtype="f16")
bias = tl.ref(0x4000, shape=(8, 8), dtype="f16")
tl.composite(
"gemm", a, b, out_ptr=0x3000,
epilogue=[{"op": "relu"}, {"op": "bias", "bias": bias}],
)
cmd = _composites(tl)[0]
assert [o.kind for o in cmd.ops] == ["gemm", "relu", "bias"]
assert cmd.ops[0].scope == Scope.OUTPUT_TILE # head scope preserved
assert cmd.ops[1].kind == "relu"
assert cmd.ops[1].scope == Scope.OUTPUT_TILE
# bias handle migrated to a named dict keyed by the EPILOGUE_OPS field.
assert cmd.ops[2].operands == {"bias": bias}, cmd.ops[2].operands
# ── MATH composite ───────────────────────────────────────────────────
def test_math_composite_lowers_to_flat_ops():
"""``op="math", math_op="exp"`` lowers to a single head op carrying the
math op kind in ``extra`` and its input as a named operand."""
tl = _tl()
a = tl.ref(0x1000, shape=(8, 16), dtype="f16")
tl.composite("math", a, math_op="exp", out_ptr=0x3000)
cmd = _composites(tl)[0]
assert len(cmd.ops) == 1
head = cmd.ops[0]
assert head.kind == "math"
assert head.extra.get("math_op") == "exp", head.extra
assert head.operands == {"a": a}, head.operands
assert head.out is not None and head.out.addr == 0x3000
# ── dataclass shape (unit) ───────────────────────────────────────────
def test_opspec_accepts_dict_operands_and_out_handle():
h = TensorHandle(id="t1", addr=0x10, shape=(4, 4), dtype="f16", nbytes=32)
op = OpSpec(kind="gemm", scope=Scope.OUTPUT_TILE, operands={"a": h}, out=h)
assert op.operands == {"a": h}
assert op.out is h
def test_compositecmd_flat_ctor_and_rw_handles_default():
from kernbench.common.pe_commands import CompletionHandle
h = TensorHandle(id="t1", addr=0x10, shape=(4, 4), dtype="f16", nbytes=32)
head = OpSpec(kind="math", scope=Scope.KERNEL, operands={"a": h}, out=h)
cmd = CompositeCmd(completion=CompletionHandle(id="c1"), ops=(head,))
assert cmd.ops == (head,)
assert cmd.rw_handles == ()
assert cmd.data_op is True
+268
View File
@@ -0,0 +1,268 @@
"""Phase 1 spec tests for ADR-0064 Revision 2 — structural CPU dispatch cost.
Rev2 replaces the Rev1 per-op-type cost table (``cpu_issue_cost.py``,
``DEFAULT_CPU_ISSUE_COST``) with a structural formula
dispatch_cycles(cmd) = FIXED_PER_CMD + cmd.logical_bytes * R
where ``logical_bytes`` is a structural property of each Pe command
(ADR-0064 D2). Tests are written against the *formula* (D1) and the byte
rule (D2), not specific absolute latencies, so they survive recalibration.
Assumed post-Phase-2 surface:
- ``kernbench.common.pe_cost_model`` exports ``PeCostModel``
(fields: ``fixed_per_cmd_cycles``, ``byte_cycles_recip``,
``max_composite_logical_bytes``; method ``dispatch_cycles(lb)``) and
``DEFAULT_PE_COST_MODEL`` (FIXED=40, R=0.0625, cap=1024).
- ``OpSpec``/``CompositeCmd``/``DmaReadCmd``/… expose ``logical_bytes``.
- ``TLContext(cost_model=...)`` charges ``FIXED + lb*R`` (÷ clock) as a
``PeCpuOverheadCmd`` before each dispatched command; ``tl.cycles(n)``
and ``WaitCmd`` are bypassed (D5).
- Per the user-revised D7: a composite whose ``logical_bytes`` exceeds
``max_composite_logical_bytes`` raises a validation error at emit time
(no auto-segmentation).
Phase 1 (this commit): tests only. All FAIL until Phase 2:
- ``pe_cost_model`` module does not exist yet,
- ``logical_bytes`` properties do not exist,
- ``TLContext`` has no ``cost_model`` kwarg.
"""
from __future__ import annotations
import pytest
from kernbench.common.pe_commands import (
CompletionHandle,
CompositeCmd,
OpSpec,
PeCpuOverheadCmd,
Scope,
TensorHandle,
)
from kernbench.triton_emu.tl_context import TLContext
def _h(addr: int = 0x10, shape: tuple[int, ...] = (64, 64)) -> TensorHandle:
return TensorHandle(id=f"h{addr:x}", addr=addr, shape=shape,
dtype="f16", nbytes=2 * shape[0] * shape[1])
def _gemm_opspec(out_h: TensorHandle) -> OpSpec:
"""The ADR-0064 D3 worked example: a single-OpSpec GEMM head."""
a, b = _h(0x1000), _h(0x2000)
return OpSpec(
kind="gemm", scope=Scope.OUTPUT_TILE,
operands={"a": a, "b": b}, out=out_h,
extra={"m": 64, "k": 64, "n": 64},
)
# ── PeCostModel defaults (D3) ────────────────────────────────────────
def test_cost_model_defaults():
from kernbench.common.pe_cost_model import DEFAULT_PE_COST_MODEL
m = DEFAULT_PE_COST_MODEL
assert m.fixed_per_cmd_cycles == 40
assert m.byte_cycles_recip == 0.0625 # = 16 bytes/cycle
assert m.max_composite_logical_bytes == 1024
def test_dispatch_cycles_formula():
"""dispatch_cycles(lb) == FIXED + lb*R (ADR-0064 D1)."""
from kernbench.common.pe_cost_model import PeCostModel
m = PeCostModel(fixed_per_cmd_cycles=40, byte_cycles_recip=0.0625)
assert m.dispatch_cycles(54) == pytest.approx(43.375) # D3 anchor
assert m.dispatch_cycles(0) == 40
# ── logical_bytes byte rule (D2 / D3 worked example) ─────────────────
def test_opspec_logical_bytes_single_gemm_is_40():
"""ADR-0064 D3: opcode1+scope1 + len1 + 2 handles16 + out8 + len1 +
m/k/n 12 = 40."""
op = _gemm_opspec(_h(0x3000))
assert op.logical_bytes == 40
def test_composite_logical_bytes_single_gemm_is_54():
"""ADR-0064 D3: framing4 + ops-len1 + GEMM 40 + rw-len1 + rw 8 = 54."""
out_h = _h(0x3000)
comp = CompositeCmd(
completion=CompletionHandle(id="c1"),
ops=(_gemm_opspec(out_h),),
rw_handles=(out_h,),
)
assert comp.logical_bytes == 54
def test_per_op_summation_no_dedup():
"""ADR-0064 D2 counting rule: each operand handle reference is counted
independently — the same handle in 3 OpSpecs + rw_handles is counted
4 times, not deduplicated."""
O = _h(0x9000, shape=(8, 8))
op = OpSpec(kind="mul", scope=Scope.KERNEL, operands={"s": O}, out=O)
# op.logical_bytes: 1+1 + 1+8(one operand) + 8(out) + 1+0(no extra) = 20
assert op.logical_bytes == 20
with_rw = CompositeCmd(
completion=CompletionHandle(id="c1"), ops=(op, op, op), rw_handles=(O,),
)
without_rw = CompositeCmd(
completion=CompletionHandle(id="c2"), ops=(op, op, op), rw_handles=(),
)
# rw_handles entry adds exactly 8 bytes; O is NOT deduplicated against
# the 3 operand references in ops.
assert with_rw.logical_bytes - without_rw.logical_bytes == 8
# Full identity: 4 + 1 + 3*20 + 1 + 8*1 = 74.
assert with_rw.logical_bytes == 74
def test_primitive_commands_have_logical_bytes():
"""Every dispatched Pe command exposes a positive structural size."""
from kernbench.common.pe_commands import (
CopyCmd, DmaReadCmd, DmaWriteCmd, GemmCmd, MathCmd,
)
a, b, out = _h(0x10), _h(0x20), _h(0x30)
assert DmaReadCmd(handle=a, src_addr=0x10, nbytes=128).logical_bytes > 0
assert DmaWriteCmd(handle=a, dst_addr=0x10, nbytes=128).logical_bytes > 0
assert GemmCmd(a=a, b=b, out=out, m=4, k=4, n=4).logical_bytes > 0
assert MathCmd(op="exp", inputs=(a,), out=out).logical_bytes > 0
assert CopyCmd(src=a, dst=b, nbytes=128).logical_bytes > 0
# ── formula wiring in TLContext (D1) ─────────────────────────────────
def test_tlcontext_charges_formula_per_command():
"""With a cost_model set, each dispatched command is preceded by a
PeCpuOverheadCmd carrying FIXED + lb*R cycles. Using R=0 isolates the
per-command FIXED term (= command-count signal) independent of lb."""
from kernbench.common.pe_cost_model import PeCostModel
tl = TLContext(
pe_id=0, num_programs=1,
cost_model=PeCostModel(fixed_per_cmd_cycles=100, byte_cycles_recip=0.0),
)
tl.load(0x1000, shape=(4, 4), dtype="f16")
overheads = [c for c in tl.commands if isinstance(c, PeCpuOverheadCmd)]
assert len(overheads) == 1
assert overheads[0].cycles == 100
def test_tl_cycles_bypasses_dispatch_cost():
"""ADR-0064 Test #4 (D5 bypass): manual tl.cycles(n) issues exactly n,
not n + dispatch_cycles."""
from kernbench.common.pe_cost_model import DEFAULT_PE_COST_MODEL
tl = TLContext(pe_id=0, num_programs=1, cost_model=DEFAULT_PE_COST_MODEL)
tl.cycles(7)
overheads = [c for c in tl.commands if isinstance(c, PeCpuOverheadCmd)]
assert overheads == [PeCpuOverheadCmd(cycles=7)]
# ── D7 (user-revised): hard cap → error, no segmentation ─────────────
def test_composite_over_cap_raises():
"""A composite whose logical_bytes exceeds max_composite_logical_bytes
raises a validation error at emit time (revised D7 — no auto-split)."""
from kernbench.common.pe_cost_model import PeCostModel
tl = TLContext(
pe_id=0, num_programs=1,
cost_model=PeCostModel(max_composite_logical_bytes=100),
)
a = tl.ref(0x1000, shape=(8, 16), dtype="f16")
b = tl.ref(0x2000, shape=(16, 8), dtype="f16")
with pytest.raises(ValueError, match="logical_bytes"):
tl.composite(
"gemm", a, b, out_ptr=0x3000,
epilogue=[{"op": "relu"}] * 50, # pushes logical_bytes past 100
)
def test_composite_within_cap_ok():
"""A normal composite (well under the default 1024 cap) emits fine."""
from kernbench.common.pe_cost_model import DEFAULT_PE_COST_MODEL
tl = TLContext(pe_id=0, num_programs=1, cost_model=DEFAULT_PE_COST_MODEL)
a = tl.ref(0x1000, shape=(8, 16), dtype="f16")
b = tl.ref(0x2000, shape=(16, 8), dtype="f16")
tl.composite("gemm", a, b, out_ptr=0x3000) # must not raise
comps = [c for c in tl.commands if isinstance(c, CompositeCmd)]
assert len(comps) == 1
assert comps[0].logical_bytes < 1024
# ── live PE_CPU wiring / path parity (ADR-0064 Test #7) ──────────────
def test_live_pe_cpu_applies_cost_model(monkeypatch):
"""The live PE_CPU path constructs TLContext with a PeCostModel (not the
removed Rev1 table). Restores the wiring assertion lost with Rev1 and
covers review item #4 / Test #7 (both live paths read the same model)."""
from pathlib import Path
from kernbench.common.pe_cost_model import PeCostModel
from kernbench.policy.address.phyaddr import PhysAddr
from kernbench.runtime_api.kernel import KernelLaunchMsg, KernelRef
from kernbench.sim_engine.engine import GraphEngine
from kernbench.sim_engine.transaction import Transaction
from kernbench.topology.builder import load_topology
from kernbench.triton_emu import tl_context as _tlc
from kernbench.triton_emu.registry import clear_registry, register_kernel
captured: list[dict] = []
real_init = _tlc.TLContext.__init__
def spy_init(self, *args, **kwargs):
captured.append(dict(kwargs))
real_init(self, *args, **kwargs)
monkeypatch.setattr(_tlc.TLContext, "__init__", spy_init)
clear_registry()
slice_bytes = 48 * (1 << 30) // 8
hbm_pa = PhysAddr.pe_hbm_addr(
sip_id=0, die_id=0, pe_id=0,
pe_local_hbm_offset=0x1000, slice_size_bytes=slice_bytes,
).encode()
def single_load_kernel(tl):
tl.load(hbm_pa, shape=(4, 4), dtype="f16")
register_kernel("test_rev2_cost_model_wiring", single_load_kernel)
topo = Path(__file__).parent.parent / "topology.yaml"
engine = GraphEngine(load_topology(topo))
pe_cpu_id = "sip0.cube0.pe0.pe_cpu"
done = engine._env.event()
txn = Transaction(
request=KernelLaunchMsg(
correlation_id="t", request_id="r",
kernel_ref=KernelRef(name="test_rev2_cost_model_wiring", kind="builtin"),
args=(),
),
path=[pe_cpu_id], step=0, nbytes=0, done=done,
)
def inject():
yield engine._components[pe_cpu_id]._inbox.put(txn)
yield done
engine._env.process(inject())
engine._env.run()
clear_registry()
pe_ctx_calls = [c for c in captured
if c.get("pe_id") == 0 and "cost_model" in c]
assert len(pe_ctx_calls) >= 1, (
f"live PE_CPU must construct TLContext with a cost_model kwarg; "
f"captures={captured}"
)
assert isinstance(pe_ctx_calls[-1]["cost_model"], PeCostModel)
-382
View File
@@ -1,382 +0,0 @@
"""Phase 1 spec tests for ADR-0064 (per-op-type CPU issue cost model).
Phase E lands the cost-table machinery and turns it on by default so the
hybrid's CPU-saturation lever (ADR-0060 §1) becomes measurable instead of
modelled away. Today every ``tl.*`` call goes through
``_emit_dispatch_overhead()`` with a single uniform ``dispatch_cycles``
scalar that is hard-coded to 0 on the live PE_CPU paths
(``pe_cpu.py:_execute_legacy`` and ``kernel_runner.py:run``).
These tests assume the post-Phase-2 surface:
- ``kernbench.common.cpu_issue_cost`` exports ``OpKind``,
``DEFAULT_CPU_ISSUE_COST`` (the ADR-0064 D1 table) and
``get_issue_cost(kind, table=None) -> int``.
- ``TLContext.__init__`` accepts ``issue_cost_table: dict[str, int] | None``.
When provided, ``_emit_dispatch_overhead(kind)`` looks up the per-kind
cost; when absent, it falls back to the uniform ``dispatch_cycles``
(back-compat with the existing ADR-0046 §D6 contract).
- Live PE_CPU paths (greenlet + legacy replay) construct TLContext with
``issue_cost_table=DEFAULT_CPU_ISSUE_COST``.
Phase 1 (this commit): tests only. All tests FAIL until Phase 2.
"""
from __future__ import annotations
from pathlib import Path
import pytest
from kernbench.common.pe_commands import (
CompositeCmd,
DmaReadCmd,
DmaWriteCmd,
GemmCmd,
MathCmd,
PeCpuOverheadCmd,
)
from kernbench.policy.address.phyaddr import PhysAddr
from kernbench.runtime_api.kernel import KernelLaunchMsg, KernelRef
from kernbench.sim_engine.engine import GraphEngine
from kernbench.sim_engine.transaction import Transaction
from kernbench.topology.builder import load_topology
from kernbench.triton_emu.registry import clear_registry, register_kernel
from kernbench.triton_emu.tl_context import TLContext, run_kernel
TOPOLOGY_PATH = Path(__file__).parent.parent / "topology.yaml"
def _engine():
return GraphEngine(load_topology(TOPOLOGY_PATH))
def _hbm_pa(sip: int = 0, cube: int = 0, pe_id: int = 0) -> int:
slice_bytes = 48 * (1 << 30) // 8
pa = PhysAddr.pe_hbm_addr(
sip_id=sip, die_id=cube, pe_id=pe_id,
pe_local_hbm_offset=0x1000, slice_size_bytes=slice_bytes,
)
return pa.encode()
# ── T1: default table shape ──────────────────────────────────────
def test_default_cost_table_has_expected_keys():
"""ADR-0064 D1 table — 8 keys with composite ≫ primitive ratio."""
from kernbench.common.cpu_issue_cost import DEFAULT_CPU_ISSUE_COST
expected_keys = {
"composite", "load", "store", "dot", "math",
"ipcq_send", "ipcq_recv", "copy_to",
}
assert set(DEFAULT_CPU_ISSUE_COST.keys()) == expected_keys, (
f"DEFAULT_CPU_ISSUE_COST keys must match ADR-0064 D1; "
f"got {set(DEFAULT_CPU_ISSUE_COST.keys())}"
)
# Composite is the lever — ratio against primitives must be ≥ 4×.
assert DEFAULT_CPU_ISSUE_COST["composite"] == 40
for primitive in ("load", "store", "dot", "math",
"ipcq_send", "ipcq_recv", "copy_to"):
assert DEFAULT_CPU_ISSUE_COST[primitive] == 5, (
f"primitive {primitive!r} default cost must be 5 ns"
)
# ── T2: get_issue_cost lookup ────────────────────────────────────
def test_get_issue_cost_lookup():
"""Helper returns table value; unknown kind returns 0 (no charge)."""
from kernbench.common.cpu_issue_cost import (
DEFAULT_CPU_ISSUE_COST,
get_issue_cost,
)
assert get_issue_cost("composite") == 40
assert get_issue_cost("load") == 5
assert get_issue_cost("unknown_kind") == 0
# Custom table override
custom = {"composite": 100, "load": 1}
assert get_issue_cost("composite", table=custom) == 100
assert get_issue_cost("load", table=custom) == 1
assert get_issue_cost("store", table=custom) == 0
# ── T3: TLContext consumes a passed cost table ───────────────────
def test_tlcontext_accepts_issue_cost_table():
"""TLContext(issue_cost_table=...) → tl.load emits the per-kind cycles."""
tl = TLContext(
pe_id=0, num_programs=1,
dispatch_cycles=0,
issue_cost_table={"load": 7, "store": 3, "dot": 2, "math": 2},
)
tl.load(0x1000, shape=(4, 4), dtype="f16")
overheads = [c for c in tl.commands if isinstance(c, PeCpuOverheadCmd)]
assert len(overheads) == 1, (
f"expected exactly one PeCpuOverheadCmd before the DmaReadCmd; "
f"got {len(overheads)} (cmds={[type(c).__name__ for c in tl.commands]})"
)
assert overheads[0].cycles == 7, (
f"load issue cost from table must be 7; got {overheads[0].cycles}"
)
# ── T4: composite ≫ primitive issue cost differential ────────────
def test_tlcontext_composite_vs_primitive_charges_differ():
"""Composite kernel charges once (40); primitive sequence charges per-op."""
from kernbench.common.cpu_issue_cost import DEFAULT_CPU_ISSUE_COST
# Composite kernel: 1 load + 1 composite = 5 + 40 = 45 cycles of issue cost.
tl_comp = TLContext(
pe_id=0, num_programs=1,
dispatch_cycles=0,
issue_cost_table=DEFAULT_CPU_ISSUE_COST,
)
a = tl_comp.load(0x1000, shape=(8, 16), dtype="f16")
b_ref = tl_comp.ref(0x2000, shape=(16, 8), dtype="f16")
tl_comp.composite("gemm", a, b_ref, out_ptr=0x3000)
comp_cycles = sum(
c.cycles for c in tl_comp.commands if isinstance(c, PeCpuOverheadCmd)
)
# Primitive kernel: 2 loads + 1 dot + 1 math (exp) = 5 + 5 + 5 + 5 = 20 cycles.
tl_prim = TLContext(
pe_id=0, num_programs=1,
dispatch_cycles=0,
issue_cost_table=DEFAULT_CPU_ISSUE_COST,
)
a = tl_prim.load(0x1000, shape=(8, 16), dtype="f16")
b = tl_prim.load(0x2000, shape=(16, 8), dtype="f16")
c_out = tl_prim.dot(a, b)
tl_prim.exp(c_out)
prim_cycles = sum(
c.cycles for c in tl_prim.commands if isinstance(c, PeCpuOverheadCmd)
)
# ADR-0064 ratio: composite issue cost >> primitive issue cost per op.
# For these specific kernels: composite=45 (5+40), primitive=20 (4×5).
assert comp_cycles == 45, f"composite kernel: expected 45, got {comp_cycles}"
assert prim_cycles == 20, f"primitive kernel: expected 20, got {prim_cycles}"
# Headline assertion: a single composite charges more than all 3
# post-load primitives combined (40 > 3×5) — the ADR-0060 §1 lever.
assert comp_cycles - 5 > 3 * 5, (
f"composite issue charge {comp_cycles - 5} must exceed "
f"3× primitive issue charge {3 * 5} (ADR-0064 ratio)"
)
# ── T5: cost table is purely additive (Q2 — no double-count) ─────
def test_cost_table_is_additive_only():
"""Q2 invariant: the cost table is additive on PE_CPU, NOT folded into
DMA/GEMM/MATH command shapes. Two TLContexts running the same kernel
with different cost tables must produce identical non-overhead command
sequences (same DmaReadCmd, GemmCmd, MathCmd, addrs, shapes, dtypes).
Only the PeCpuOverheadCmd ``cycles`` field is allowed to differ.
"""
from kernbench.common.cpu_issue_cost import DEFAULT_CPU_ISSUE_COST
def kernel(tl):
a = tl.load(0x1000, shape=(8, 16), dtype="f16")
b = tl.load(0x2000, shape=(16, 8), dtype="f16")
c = tl.dot(a, b)
d = tl.exp(c)
tl.store(0x3000, d)
tl_zero = TLContext(
pe_id=0, num_programs=1,
dispatch_cycles=0,
issue_cost_table={}, # empty table → every kind = 0 cost
)
run_kernel(kernel, tl_zero)
tl_default = TLContext(
pe_id=0, num_programs=1,
dispatch_cycles=0,
issue_cost_table=DEFAULT_CPU_ISSUE_COST,
)
run_kernel(kernel, tl_default)
# Strip PeCpuOverheadCmd from both streams; what remains must match.
non_overhead_zero = [
c for c in tl_zero.commands if not isinstance(c, PeCpuOverheadCmd)
]
non_overhead_default = [
c for c in tl_default.commands if not isinstance(c, PeCpuOverheadCmd)
]
assert len(non_overhead_zero) == len(non_overhead_default), (
f"non-overhead command count differs: "
f"{len(non_overhead_zero)} vs {len(non_overhead_default)}"
)
for a_cmd, b_cmd in zip(non_overhead_zero, non_overhead_default):
assert type(a_cmd) is type(b_cmd), (
f"non-overhead command type changed under cost table: "
f"{type(a_cmd).__name__} vs {type(b_cmd).__name__}"
)
# Total overhead under zero-table must be 0; under default must be > 0.
cycles_zero = sum(
c.cycles for c in tl_zero.commands if isinstance(c, PeCpuOverheadCmd)
)
cycles_default = sum(
c.cycles for c in tl_default.commands if isinstance(c, PeCpuOverheadCmd)
)
assert cycles_zero == 0, f"empty table must add 0 cycles; got {cycles_zero}"
assert cycles_default > 0, (
f"default table must add > 0 cycles; got {cycles_default}"
)
# ── T6: greenlet vs legacy replay use same cost table (review #6) ─
def test_greenlet_and_legacy_path_parity():
"""Both PE_CPU execution paths read the same cost table.
The greenlet path (kernel_runner.py:run) and the legacy replay path
(pe_cpu.py:_execute_legacy) must construct TLContext with the same
default cost table so the same kernel produces identical PE_CPU
overhead cycles via either route. This is ADR-0064 review item #6.
"""
from kernbench.common.cpu_issue_cost import DEFAULT_CPU_ISSUE_COST
# Build the command list for a representative kernel via TLContext
# using the default table — this is what both live paths should see.
def kernel(tl):
a = tl.load(0x1000, shape=(4, 4), dtype="f16")
b = tl.load(0x2000, shape=(4, 4), dtype="f16")
c = tl.dot(a, b)
tl.store(0x3000, c)
tl1 = TLContext(
pe_id=0, num_programs=1,
dispatch_cycles=0,
issue_cost_table=DEFAULT_CPU_ISSUE_COST,
)
run_kernel(kernel, tl1)
cycles1 = sum(
c.cycles for c in tl1.commands if isinstance(c, PeCpuOverheadCmd)
)
tl2 = TLContext(
pe_id=0, num_programs=1,
dispatch_cycles=0,
issue_cost_table=DEFAULT_CPU_ISSUE_COST,
)
run_kernel(kernel, tl2)
cycles2 = sum(
c.cycles for c in tl2.commands if isinstance(c, PeCpuOverheadCmd)
)
assert cycles1 == cycles2, (
f"same kernel via same default table must produce identical "
f"overhead cycles; got {cycles1} vs {cycles2}"
)
# Concrete expected for this kernel: 2×load(5) + 1×dot(5) + 1×store(5) = 20.
assert cycles1 == 20, (
f"expected 20 cycles total (2 load + 1 dot + 1 store at 5 ns); "
f"got {cycles1}"
)
# ── T7: back-compat — no table → uniform dispatch_cycles ─────────
def test_back_compat_no_table_uses_dispatch_cycles():
"""ADR-0046 §D6 contract preserved when no issue_cost_table is passed.
Existing call sites doing ``TLContext(dispatch_cycles=0)`` must
continue to emit zero overhead. Existing call sites doing
``TLContext(dispatch_cycles=1)`` must continue to emit uniform 1.
"""
# Zero path (most existing tests use this).
tl_zero = TLContext(pe_id=0, num_programs=1, dispatch_cycles=0)
tl_zero.load(0x1000, shape=(4, 4), dtype="f16")
overheads = [c for c in tl_zero.commands if isinstance(c, PeCpuOverheadCmd)]
assert overheads == [], (
f"dispatch_cycles=0 with no table must emit no overhead; "
f"got {[c.cycles for c in overheads]}"
)
# Uniform path (test_dispatch_overhead_inserted relies on this).
tl_one = TLContext(pe_id=0, num_programs=1, dispatch_cycles=1)
tl_one.load(0x1000, shape=(4, 4), dtype="f16")
overheads = [c for c in tl_one.commands if isinstance(c, PeCpuOverheadCmd)]
assert overheads == [PeCpuOverheadCmd(cycles=1)], (
f"dispatch_cycles=1 with no table must emit cycles=1; "
f"got {[c.cycles for c in overheads]}"
)
# ── T8: live PE_CPU constructs TLContext with the default table ──
def test_live_pe_cpu_uses_default_table(monkeypatch):
"""End-to-end: live PE_CPU greenlet path constructs TLContext with
``issue_cost_table=DEFAULT_CPU_ISSUE_COST``. This is the wiring assertion
for ADR-0064 D4 "active by default" + review item #6 (path parity).
We patch ``TLContext.__init__`` to record its kwargs and run a
single-load kernel through the live PE_CPU. The recorded
``issue_cost_table`` must equal ``DEFAULT_CPU_ISSUE_COST``.
"""
from kernbench.common.cpu_issue_cost import DEFAULT_CPU_ISSUE_COST
from kernbench.triton_emu import tl_context as _tlc
captured: list[dict] = []
real_init = _tlc.TLContext.__init__
def spy_init(self, *args, **kwargs):
captured.append(dict(kwargs))
real_init(self, *args, **kwargs)
monkeypatch.setattr(_tlc.TLContext, "__init__", spy_init)
clear_registry()
hbm_pa = _hbm_pa(sip=0, cube=0, pe_id=0)
def single_load_kernel(tl):
tl.load(hbm_pa, shape=(4, 4), dtype="f16")
register_kernel("test_active_default_table", single_load_kernel)
engine = _engine()
pe_cpu_id = "sip0.cube0.pe0.pe_cpu"
done = engine._env.event()
txn = Transaction(
request=KernelLaunchMsg(
correlation_id="t", request_id="r",
kernel_ref=KernelRef(name="test_active_default_table", kind="builtin"),
args=(),
),
path=[pe_cpu_id], step=0, nbytes=0, done=done,
)
def inject():
yield engine._components[pe_cpu_id]._inbox.put(txn)
yield done
engine._env.process(inject())
engine._env.run()
clear_registry()
# Live PE_CPU constructs TLContext on either the greenlet path
# (kernel_runner.py:run) or the legacy replay path
# (pe_cpu.py:_execute_legacy) — whichever is chosen depends on whether
# MemoryStore is wired. Both must pass DEFAULT_CPU_ISSUE_COST.
pe_ctx_calls = [c for c in captured if c.get("pe_id") == 0
and "issue_cost_table" in c]
assert len(pe_ctx_calls) >= 1, (
f"expected at least one TLContext constructed by a live PE_CPU "
f"path with an issue_cost_table kwarg; got {len(pe_ctx_calls)} "
f"(all captures: {captured})"
)
last = pe_ctx_calls[-1]
assert last["issue_cost_table"] == DEFAULT_CPU_ISSUE_COST, (
f"live PE_CPU must use DEFAULT_CPU_ISSUE_COST; got {last['issue_cost_table']}"
)
+5 -3
View File
@@ -178,9 +178,11 @@ def test_tl_composite_nonblocking():
assert isinstance(h, CompletionHandle)
comp_cmds = [c for c in tl.commands if isinstance(c, CompositeCmd)]
assert len(comp_cmds) == 1
assert comp_cmds[0].op == "gemm"
assert comp_cmds[0].out_addr == 0x3000
assert comp_cmds[0].out_nbytes == 32 * 32 * 2 # M×N×dtype_bytes
# Flat-ops shape (ADR-0065 D1): head op carries kind + write-back handle.
head = comp_cmds[0].ops[0]
assert head.kind == "gemm"
assert head.out.addr == 0x3000
assert head.out.nbytes == 32 * 32 * 2 # M×N×dtype_bytes
# ── 8. tl.wait(handle) → WaitCmd ─────────────────────────────────