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:
@@ -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)
|
||||
Reference in New Issue
Block a user