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
-52
View File
@@ -1,52 +0,0 @@
"""Per-op-type CPU issue cost table (ADR-0064 D1).
Replaces the single uniform ``dispatch_cycles`` scalar with a cost table
keyed by command kind. Charged on PE_CPU at issue time (before the command
is dispatched to PE_SCHEDULER) so the hybrid's CPU-saturation win
(ADR-0060 §1) becomes measurable.
The table is consulted by ``TLContext._emit_dispatch_overhead(kind)``;
live PE_CPU paths (greenlet via ``kernel_runner.py``, legacy replay via
``pe_cpu.py:_execute_legacy``) construct TLContext with
``issue_cost_table=DEFAULT_CPU_ISSUE_COST`` so all benches see the cost.
Absolute ns values are provisional (ADR-0064 review item #1). The
defensible claim is the **ratio** — composite ≫ primitive.
"""
from __future__ import annotations
from typing import Literal
OpKind = Literal[
"composite",
"load",
"store",
"dot",
"math",
"ipcq_send",
"ipcq_recv",
"copy_to",
]
DEFAULT_CPU_ISSUE_COST: dict[str, int] = {
"composite": 40,
"load": 5,
"store": 5,
"dot": 5,
"math": 5,
"ipcq_send": 5,
"ipcq_recv": 5,
"copy_to": 5,
}
def get_issue_cost(kind: str, table: dict[str, int] | None = None) -> int:
"""Return per-op-type CPU issue cost in ns.
Unknown kinds return 0 (no charge) so adding a new ``tl.*`` op kind
doesn't accidentally over-charge before the table is updated.
"""
if table is None:
table = DEFAULT_CPU_ISSUE_COST
return table.get(kind, 0)
+12
View File
@@ -123,6 +123,12 @@ class IpcqSendCmd:
data: Any = None
data_op: bool = True # ADR-0020 op_log recording flag
@property
def logical_bytes(self) -> int:
# framing + direction enum + src_addr + src_space enum + nbytes
# + shape (len marker + 4·rank) + dtype tag (ADR-0064 D2).
return 4 + 1 + 4 + 1 + 4 + (1 + 4 * len(self.shape)) + 1
# ── D12: IpcqRecvCmd (PE_CPU → PE_IPCQ) ──────────────────────────────
@@ -155,6 +161,12 @@ class IpcqRecvCmd:
data_op: bool = True
consume: bool = True # DIAGNOSTIC: see docstring
@property
def logical_bytes(self) -> int:
# framing + direction enum + shape (len marker + 4·rank) + dtype tag
# + dst_addr + dst_space enum (ADR-0064 D2).
return 4 + 1 + (1 + 4 * len(self.shape)) + 1 + 4 + 1
# ── D12: IpcqDmaToken (PE_IPCQ → PE_DMA, vc_comm) ───────────────────
+72 -12
View File
@@ -10,7 +10,7 @@ from __future__ import annotations
from dataclasses import dataclass, field
from enum import Enum
from typing import TYPE_CHECKING, Any, Literal
from typing import TYPE_CHECKING, Any
if TYPE_CHECKING:
import simpy
@@ -22,6 +22,24 @@ class Scope(Enum):
KERNEL = "kernel"
def _extra_bytes(v: Any) -> int:
"""Type-aware HW-logical byte count for an OpSpec.extra value (ADR-0064 D2).
bool→1, int/float→4 (scalar), tuple/list→1 + 4·len (length marker + 4
per element, e.g. shape/axes), str→1 (opcode-like tag). Default→4.
``bool`` is checked first because it subclasses ``int``.
"""
if isinstance(v, bool):
return 1
if isinstance(v, (int, float)):
return 4
if isinstance(v, (tuple, list)):
return 1 + 4 * len(v)
if isinstance(v, str):
return 1
return 4
@dataclass(frozen=True)
class OpSpec:
"""One operation in a multi-op composite (head + epilogue, ADR-0014 D3.3).
@@ -33,9 +51,22 @@ class OpSpec:
kind: str # "gemm" | "bias" | "relu" | ...
scope: "Scope" = Scope.OUTPUT_TILE
operands: tuple[Any, ...] = () # tuple[TensorHandle, ...]
operands: dict[str, Any] = field(default_factory=dict) # name → TensorHandle
scalar: float | None = None
extra: dict[str, Any] = field(default_factory=dict)
out: "TensorHandle | None" = None # explicit write-back handle
@property
def logical_bytes(self) -> int:
"""HW-logical byte size (ADR-0064 D2). ``scalar`` is a transitional
field not part of the flat-ops model (ADR-0065 D2) — excluded,
pending its removal."""
return (
1 + 1 # opcode + scope enum
+ 1 + 8 * len(self.operands) # len marker + handles
+ (8 if self.out is not None else 0) # out handle
+ 1 + sum(_extra_bytes(v) for v in self.extra.values())
)
# Epilogue op contracts: kind → (required field names, default scope).
@@ -103,6 +134,10 @@ class DmaReadCmd:
nbytes: int
data_op: bool = True
@property
def logical_bytes(self) -> int:
return 4 + 8 + 4 + 4 # framing + handle + src_addr + nbytes
@dataclass(frozen=True)
class DmaWriteCmd:
@@ -113,6 +148,10 @@ class DmaWriteCmd:
nbytes: int
data_op: bool = True
@property
def logical_bytes(self) -> int:
return 4 + 8 + 4 + 4 # framing + handle + dst_addr + nbytes
@dataclass(frozen=True)
class GemmCmd:
@@ -129,6 +168,10 @@ class GemmCmd:
n: int
data_op: bool = True
@property
def logical_bytes(self) -> int:
return 4 + 8 * 3 + 4 * 3 # framing + 3 handles + m/k/n scalars
@dataclass(frozen=True)
class MathCmd:
@@ -145,6 +188,13 @@ class MathCmd:
axis: int | None = None # for reductions
data_op: bool = True
@property
def logical_bytes(self) -> int:
return (
4 + 1 + 1 + 8 * len(self.inputs) + 8 # framing+opcode+len+inputs+out
+ (4 if self.axis is not None else 0)
)
@dataclass(frozen=True)
class CopyCmd:
@@ -161,6 +211,10 @@ class CopyCmd:
nbytes: int
data_op: bool = True
@property
def logical_bytes(self) -> int:
return 4 + 8 * 2 + 4 # framing + src/dst handles + nbytes
@dataclass(frozen=True)
class CompositeCmd:
@@ -168,20 +222,26 @@ class CompositeCmd:
Non-blocking — submitted to PE_SCHEDULER which manages tile splitting
and pipeline overlaps (ADR-0014 D3.2).
Flat-ops shape (ADR-0065 D1): ``ops`` is an ordered list of OpSpecs.
The GEMM op (if any, ≤1) drives the tile loop; preceding/following
OpSpecs are placed by position + scope. ``rw_handles`` carries
cross-composite hazard metadata (ADR-0065 D6.3); unused in P1.
"""
completion: CompletionHandle
op: Literal["gemm", "math"]
a: TensorHandle
b: TensorHandle | None
out_addr: int
out_nbytes: int
math_op: str | None = None # for op="math": which math operation
data_op: bool = True
# Multi-op composite (ADR-0014 D3.3): when non-empty, ops[0] is the
# head and ops[1:] are epilogue stages with explicit scope. When empty,
# the legacy single-op semantics (op/a/b/math_op) apply.
ops: tuple[OpSpec, ...] = ()
rw_handles: tuple["TensorHandle", ...] = ()
data_op: bool = True
@property
def logical_bytes(self) -> int:
"""HW-logical byte size (ADR-0064 D2). Per-op summation, no dedup."""
return (
4 # framing
+ 1 + sum(op.logical_bytes for op in self.ops)
+ 1 + 8 * len(self.rw_handles)
)
@dataclass(frozen=True)
+57
View File
@@ -0,0 +1,57 @@
"""Structural PE_CPU dispatch cost model (ADR-0064 Revision 2).
Replaces the Rev1 per-op-type calibration table (``cpu_issue_cost.py``,
removed) with a structural formula derived from each command's
``logical_bytes`` (ADR-0064 D2):
dispatch_cycles(cmd) = FIXED_PER_CMD + cmd.logical_bytes * R
The per-command ``FIXED_PER_CMD`` term models the queue-tail update /
MMIO-class RTT / completion-event registration; the byte term ``R`` models
queue-write bandwidth. The primary signal the model exposes is
**command-count reduction** (FIXED-dominated); the byte term is a secondary
refinement (ADR-0064 Context).
Knobs are cycle-domain only; cycle→ns uses the PE node's ``clock_freq_ghz``
(ADR-0064 D3). Defaults anchor a typical 1-OpSpec GEMM composite (≈54 bytes)
at ≈43 ns on a 16 B/cycle on-die descriptor queue.
A composite whose ``logical_bytes`` exceeds ``max_composite_logical_bytes``
is rejected at emit time with a ``ValueError`` (ADR-0064 D7, revised: hard
cap, no auto-segmentation).
"""
from __future__ import annotations
from dataclasses import dataclass
@dataclass(frozen=True)
class PeCostModel:
"""Cycle-domain dispatch cost knobs (ADR-0064 D3/D4)."""
fixed_per_cmd_cycles: float = 40.0
byte_cycles_recip: float = 0.0625 # = 16 bytes/cycle
max_composite_logical_bytes: int = 1024 # D7 hard cap
def dispatch_cycles(self, logical_bytes: int) -> float:
"""FIXED + logical_bytes × R (ADR-0064 D1)."""
return self.fixed_per_cmd_cycles + logical_bytes * self.byte_cycles_recip
DEFAULT_PE_COST_MODEL = PeCostModel()
def from_node_attrs(attrs: dict) -> PeCostModel:
"""Build a PeCostModel from a PE node's ``pe_cost_model:`` attrs block
(ADR-0064 D4). Missing keys fall back to defaults."""
block = attrs.get("pe_cost_model", {}) or {}
d = DEFAULT_PE_COST_MODEL
return PeCostModel(
fixed_per_cmd_cycles=float(
block.get("fixed_per_cmd_cycles", d.fixed_per_cmd_cycles)),
byte_cycles_recip=float(
block.get("byte_cycles_recip", d.byte_cycles_recip)),
max_composite_logical_bytes=int(
block.get("max_composite_logical_bytes",
d.max_composite_logical_bytes)),
)
+12 -2
View File
@@ -66,6 +66,14 @@ class PeCpuComponent(ComponentBase):
| (self._cube_idx << 32)
| (self._pe_idx << 24)
)
# ADR-0064 Rev2: structural dispatch cost model. Reads an optional
# `pe_cost_model:` block from this PE_CPU node's attrs (D4); missing
# keys fall back to defaults. Clock for cycle→ns is this node's
# `clock_freq_ghz` (D3), same attr PE_GEMM/PE_MATH use.
from kernbench.common.pe_cost_model import from_node_attrs
self._cost_model = from_node_attrs(node.attrs)
self._clock_freq_ghz = float(node.attrs.get("clock_freq_ghz", 1.0))
def _find_shard(self, shards: tuple) -> Any:
"""Find shard matching this PE's (sip, cube, pe). Fallback to positional index."""
@@ -176,6 +184,8 @@ class PeCpuComponent(ComponentBase):
store=store,
scratch_base=self._tl_scratch_base,
scratch_size=self._tl_scratch_size,
cost_model=self._cost_model,
clock_freq_ghz=self._clock_freq_ghz,
)
yield from runner.run(env, kernel_fn, kernel_args, num_programs)
return getattr(runner, "_composite_results", [])
@@ -184,7 +194,6 @@ class PeCpuComponent(ComponentBase):
self, env, kernel_fn, kernel_args, num_programs, scheduler_id,
) -> Generator:
"""Legacy Phase 0 + replay: generate command list, then dispatch."""
from kernbench.common.cpu_issue_cost import DEFAULT_CPU_ISSUE_COST
from kernbench.common.pe_commands import (
CompositeCmd, PeCpuOverheadCmd, PeInternalTxn, WaitCmd,
)
@@ -194,7 +203,8 @@ class PeCpuComponent(ComponentBase):
pe_id=self._pe_idx, num_programs=num_programs,
cube_id=self._cube_idx, num_cubes=self._num_cubes,
dispatch_cycles=0,
issue_cost_table=DEFAULT_CPU_ISSUE_COST,
cost_model=self._cost_model,
clock_freq_ghz=self._clock_freq_ghz,
)
run_kernel(kernel_fn, tl, *kernel_args)
commands = tl.commands
@@ -156,19 +156,21 @@ class PeSchedulerComponent(ComponentBase):
pp = self._pe_prefix
bpe = 2 # default bytes per element (f16)
if cmd.op == "gemm" and cmd.b is not None:
a = cmd.a
b = cmd.b
# Flat-ops (ADR-0065 D1): ops[0] is the head, ops[1:] are epilogue
# specs placed by scope. The head's kind selects the engine path.
head = cmd.ops[0]
epi_specs = tuple(cmd.ops[1:])
if head.kind == "gemm" and "b" in head.operands:
a = head.operands["a"]
b = head.operands["b"]
M, K = a.shape[-2], a.shape[-1]
N = b.shape[-1]
# When CompositeCmd.ops is populated, ops[0] is the head and
# ops[1:] is the epilogue spec list. Empty ops → legacy path.
epi_specs = tuple(cmd.ops[1:]) if cmd.ops else ()
return generate_gemm_plan(
M=M, K=K, N=N,
tile_m=self.TILE_M, tile_k=self.TILE_K, tile_n=self.TILE_N,
bytes_per_element=bpe,
A_addr=a.addr, B_addr=b.addr, C_addr=cmd.out_addr,
A_addr=a.addr, B_addr=b.addr, C_addr=head.out.addr,
pe_prefix=pp,
a_pinned=getattr(a, "pinned", False),
b_pinned=getattr(b, "pinned", False),
@@ -176,14 +178,14 @@ class PeSchedulerComponent(ComponentBase):
)
else:
# Math composite
a = cmd.a
a = head.operands["a"]
M = a.shape[-2] if len(a.shape) >= 2 else a.shape[0]
N = a.shape[-1] if len(a.shape) >= 2 else 1
return generate_math_plan(
M=M, N=N,
tile_m=self.TILE_M, tile_n=self.TILE_N,
bytes_per_element=bpe,
math_op=cmd.math_op or "identity",
src_addr=a.addr, dst_addr=cmd.out_addr,
math_op=head.extra.get("math_op") or "identity",
src_addr=a.addr, dst_addr=head.out.addr,
pe_prefix=pp,
)
+20 -14
View File
@@ -248,27 +248,33 @@ def _extract_op_info(msg: Any) -> tuple[str, str, dict[str, Any]]:
"nbytes": msg.nbytes,
}
if isinstance(msg, CompositeCmd):
# Flat-ops (ADR-0065 D1): the head op carries op kind, operands, and
# the write-back handle that previously lived on the cmd directly.
head = msg.ops[0]
op = head.kind
params: dict[str, Any] = {
"op": msg.op,
"out_addr": msg.out_addr,
"out_nbytes": msg.out_nbytes,
"op": op,
"out_addr": head.out.addr,
"out_nbytes": head.out.nbytes,
}
# ADR-0027: preserve operand info so Phase 2 DataExecutor can replay
# the composite's numerical effect (treat it like a GemmCmd).
if msg.op == "gemm" and msg.a is not None and msg.b is not None:
a = head.operands.get("a")
b = head.operands.get("b")
if op == "gemm" and a is not None and b is not None:
params.update({
"src_a_addr": msg.a.addr,
"src_b_addr": msg.b.addr,
"shape_a": msg.a.shape,
"shape_b": msg.b.shape,
"dtype_in": msg.a.dtype,
"dtype_out": msg.a.dtype,
"src_a_space": getattr(msg.a, "space", "hbm"),
"src_b_space": getattr(msg.b, "space", "hbm"),
"src_a_addr": a.addr,
"src_b_addr": b.addr,
"shape_a": a.shape,
"shape_b": b.shape,
"dtype_in": a.dtype,
"dtype_out": a.dtype,
"src_a_space": getattr(a, "space", "hbm"),
"src_b_space": getattr(b, "space", "hbm"),
"dst_space": "hbm",
# dst_addr alias so DataExecutor._execute_gemm picks it up.
"dst_addr": msg.out_addr,
"dst_addr": head.out.addr,
})
return "gemm" if msg.op == "gemm" else "math", f"composite_{msg.op}", params
return "gemm" if op == "gemm" else "math", f"composite_{op}", params
# Fallback for unknown data_op messages
return "unknown", type(msg).__name__, {}
+7 -2
View File
@@ -55,6 +55,8 @@ class KernelRunner:
ipcq_id: str | None = None,
scratch_base: int = 0,
scratch_size: int = 1 << 20,
cost_model: Any = None,
clock_freq_ghz: float = 1.0,
) -> None:
self._pe_prefix = pe_prefix
self._pe_idx = pe_idx
@@ -72,6 +74,9 @@ class KernelRunner:
# ops produce a result that may later be used as a send/store source.
self._scratch_base = scratch_base
self._scratch_size = scratch_size
# ADR-0064 Rev2 structural dispatch cost model (+ clock for cycle→ns).
self._cost_model = cost_model
self._clock_freq_ghz = clock_freq_ghz
def run(
self,
@@ -89,7 +94,6 @@ class KernelRunner:
4. Dispatches each command through SimPy components
5. Returns results to the kernel
"""
from kernbench.common.cpu_issue_cost import DEFAULT_CPU_ISSUE_COST
from kernbench.triton_emu.tl_context import TLContext
self._parent = greenlet.getcurrent()
@@ -103,7 +107,8 @@ class KernelRunner:
runner=self,
scratch_base=self._scratch_base,
scratch_size=self._scratch_size,
issue_cost_table=DEFAULT_CPU_ISSUE_COST,
cost_model=self._cost_model,
clock_freq_ghz=self._clock_freq_ghz,
)
self._tl = tl # exposed so switch_to_simpy can re-set on restore
+86 -69
View File
@@ -93,17 +93,16 @@ class TLContext:
Args:
pe_id: program instance index (returned by program_id).
num_programs: total number of program instances.
dispatch_cycles: uniform PE_CPU overhead per tl API call. Used as
a fallback when ``issue_cost_table`` is None (ADR-0046 §D6
back-compat). When ``issue_cost_table`` is provided, the
per-kind table value is used instead.
issue_cost_table: optional per-op-type CPU issue cost table
(ADR-0064 D1). When provided, each ``tl.*`` call charges the
table value keyed by op kind ("composite", "load", "store",
"dot", "math", "ipcq_send", "ipcq_recv", "copy_to"). Unknown
kinds fall back to ``dispatch_cycles``. Live PE_CPU paths
construct TLContext with ``DEFAULT_CPU_ISSUE_COST`` so the
dispatch_cycles: uniform PE_CPU overhead per dispatched command.
Back-compat fallback used only when ``cost_model`` is None
(ADR-0046 §D6).
cost_model: structural dispatch cost model (ADR-0064 Rev2). When
provided, every dispatched PeCommand is charged
``FIXED + cmd.logical_bytes × R`` cycles (÷ ``clock_freq_ghz``
for ns) as a ``PeCpuOverheadCmd`` emitted just before it. Live
PE_CPU paths construct TLContext with the per-PE model so the
hybrid's CPU-saturation lever is measurable.
clock_freq_ghz: cycle→ns conversion for the cost model (ADR-0064 D3).
"""
def __init__(
@@ -116,14 +115,16 @@ class TLContext:
num_cubes: int = 1,
scratch_base: int = 0,
scratch_size: int = 1 << 20, # 1 MiB per kernel invocation
issue_cost_table: dict[str, int] | None = None,
cost_model: "PeCostModel | None" = None,
clock_freq_ghz: float = 1.0,
) -> None:
self._pe_id = pe_id
self._num_programs = num_programs
self._cube_id = cube_id
self._num_cubes = num_cubes
self._dispatch_cycles = dispatch_cycles
self._issue_cost_table = issue_cost_table
self._cost_model = cost_model
self._clock_freq_ghz = clock_freq_ghz
self._commands: list[PeCommand] = []
self._handle_counter = 0
self._completion_counter = 0
@@ -193,23 +194,30 @@ class TLContext:
def _nbytes(self, shape: tuple[int, ...], dtype: str) -> int:
return math.prod(shape) * self._dtype_bytes(dtype)
def _emit_dispatch_overhead(self, kind: str | None = None) -> None:
"""Charge per-op-type CPU issue cost (ADR-0064 D1).
def _charge_dispatch(self, cmd: PeCommand) -> None:
"""Charge structural PE_CPU dispatch cost (ADR-0064 Rev2 D1).
When ``issue_cost_table`` was provided, look up the per-kind cost
and emit ``PeCpuOverheadCmd(cycles=N)`` if N > 0. Unknown kinds
fall back to the uniform ``dispatch_cycles`` for forward-compat
when a new ``tl.*`` op is added before the table is updated.
When ``issue_cost_table`` is None, preserve the ADR-0046 §D6
contract: emit ``PeCpuOverheadCmd(dispatch_cycles)`` if positive.
With a ``cost_model``: emit ``PeCpuOverheadCmd`` carrying
``round((FIXED + cmd.logical_bytes × R) / clock_freq_ghz)`` ns just
before ``cmd``. A CompositeCmd over the descriptor-size cap is
rejected (D7). Without a cost_model, fall back to the uniform
``dispatch_cycles`` contract (ADR-0046 §D6).
"""
if self._issue_cost_table is not None and kind is not None:
cycles = self._issue_cost_table.get(kind, self._dispatch_cycles)
else:
cycles = self._dispatch_cycles
if cycles > 0:
self._emit(PeCpuOverheadCmd(cycles=cycles))
if self._cost_model is not None:
if isinstance(cmd, CompositeCmd):
lb = cmd.logical_bytes
cap = self._cost_model.max_composite_logical_bytes
if lb > cap:
raise ValueError(
f"CompositeCmd logical_bytes {lb} exceeds "
f"max_composite_logical_bytes {cap} (ADR-0064 D7)"
)
cycles = self._cost_model.dispatch_cycles(cmd.logical_bytes)
ns = round(cycles / self._clock_freq_ghz)
if ns > 0:
self._emit(PeCpuOverheadCmd(cycles=ns))
elif self._dispatch_cycles > 0:
self._emit(PeCpuOverheadCmd(cycles=self._dispatch_cycles))
def _make_handle(
self, addr: int, shape: tuple[int, ...], dtype: str,
@@ -255,7 +263,15 @@ class TLContext:
# ── Data Movement (blocking, DMA engine) ──────────────────────
def _emit(self, cmd: PeCommand) -> Any:
"""Emit command: greenlet switch if runner available, else append to list."""
"""Emit command: greenlet switch if runner available, else append to list.
Each dispatched PeCommand is preceded by a ``PeCpuOverheadCmd``
carrying its structural dispatch cost (ADR-0064 Rev2 D1).
``PeCpuOverheadCmd`` (incl. manual ``tl.cycles``) and ``WaitCmd``
bypass the charge (D5).
"""
if not isinstance(cmd, (PeCpuOverheadCmd, WaitCmd)):
self._charge_dispatch(cmd)
if self._runner is not None:
return self._runner.switch_to_simpy(cmd)
self._commands.append(cmd)
@@ -276,7 +292,6 @@ class TLContext:
attaches a LoadFuture for structural compatibility (its event
stays None — no engine, no SimPy event to wait on).
"""
self._emit_dispatch_overhead("load")
nbytes = self._nbytes(shape, dtype)
# LoadFuture is mutable; create it first, attach to the handle,
# then point its ``cmd`` at the DmaReadCmd that references the
@@ -292,6 +307,9 @@ class TLContext:
)
cmd = DmaReadCmd(handle=handle, src_addr=ptr, nbytes=nbytes)
future.cmd = cmd
# Lazy load bypasses _emit (it posts via a load_issue future), so it
# charges its own dispatch cost here (ADR-0064 D1).
self._charge_dispatch(cmd)
if self._runner is not None:
# Lazy: runner posts the DmaReadCmd, sets future.event, then
# switches back immediately. No yield on completion here.
@@ -320,7 +338,6 @@ class TLContext:
def store(self, ptr: int, handle: TensorHandle) -> None:
"""Store tensor from TCM to HBM."""
self._await_pending(handle)
self._emit_dispatch_overhead("store")
cmd = DmaWriteCmd(handle=handle, dst_addr=ptr, nbytes=handle.nbytes)
self._emit(cmd)
@@ -358,7 +375,6 @@ class TLContext:
"reads from HBM go through tl.load"
)
self._await_pending(src)
self._emit_dispatch_overhead("copy_to")
self._emit(CopyCmd(src=src, dst=dst, nbytes=src.nbytes))
# ── GEMM Engine (blocking) ────────────────────────────────────
@@ -378,7 +394,6 @@ class TLContext:
out_dtype = a.dtype
out = self._make_compute_out(shape=out_shape, dtype=out_dtype)
self._await_pending(a, b)
self._emit_dispatch_overhead("dot")
self._emit(GemmCmd(a=a, b=b, out=out, m=m, k=k, n=n))
return out
@@ -387,7 +402,6 @@ class TLContext:
def _unary_math(self, op: str, x: TensorHandle) -> TensorHandle:
out = self._make_compute_out(shape=x.shape, dtype=x.dtype)
self._await_pending(x)
self._emit_dispatch_overhead("math")
self._emit(MathCmd(op=op, inputs=(x,), out=out))
return out
@@ -421,7 +435,6 @@ class TLContext:
out_shape[axis] = 1
out = self._make_compute_out(shape=tuple(out_shape), dtype=x.dtype)
self._await_pending(x)
self._emit_dispatch_overhead("math")
self._emit(MathCmd(op=op, inputs=(x,), out=out, axis=axis))
return out
@@ -441,7 +454,6 @@ class TLContext:
) -> TensorHandle:
out = self._make_compute_out(shape=a.shape, dtype=a.dtype)
self._await_pending(a, b)
self._emit_dispatch_overhead("math")
self._emit(MathCmd(op=op, inputs=(a, b), out=out))
return out
@@ -450,7 +462,6 @@ class TLContext:
) -> TensorHandle:
out = self._make_compute_out(shape=a.shape, dtype=a.dtype)
self._await_pending(cond, a, b)
self._emit_dispatch_overhead("math")
self._emit(MathCmd(op="where", inputs=(cond, a, b), out=out))
return out
@@ -468,7 +479,6 @@ class TLContext:
"""Fused multiply-add: a * b + c (real Triton: tl.fma)."""
out = self._make_compute_out(shape=a.shape, dtype=a.dtype)
self._await_pending(a, b, c)
self._emit_dispatch_overhead("math")
self._emit(MathCmd(op="fma", inputs=(a, b, c), out=out))
return out
@@ -481,7 +491,6 @@ class TLContext:
"""Clamp x to [min, max] (real Triton: tl.clamp)."""
out = self._make_compute_out(shape=x.shape, dtype=x.dtype)
self._await_pending(x, min, max)
self._emit_dispatch_overhead("math")
self._emit(MathCmd(op="clamp", inputs=(x, min, max), out=out))
return out
@@ -494,7 +503,6 @@ class TLContext:
"""
out = self._make_compute_out(shape=x.shape, dtype=x.dtype)
self._await_pending(x)
self._emit_dispatch_overhead("math")
self._emit(MathCmd(op="softmax", inputs=(x,), out=out, axis=axis))
return out
@@ -597,7 +605,6 @@ class TLContext:
# later IPCQ inbound overwrites the slot before the outbound
# PE_DMA reads it.
handle_data = getattr(src, "data", None) if src is not None else None
self._emit_dispatch_overhead("ipcq_send")
cmd = IpcqSendCmd(
direction=dir,
src_addr=src_addr, src_space=space,
@@ -631,7 +638,6 @@ class TLContext:
arrived. In greenlet/runner mode, ``handle.data`` carries the
actual ndarray; in command-list mode the handle is a placeholder.
"""
self._emit_dispatch_overhead("ipcq_recv")
if dst_addr is not None and dst_space is not None:
cmd = IpcqRecvCmd(
direction=dir,
@@ -682,7 +688,6 @@ class TLContext:
they receive. This API is segregated from ``tl.recv`` so the
diagnostic flag can never accidentally be set in real workloads.
"""
self._emit_dispatch_overhead("ipcq_recv")
cmd = IpcqRecvCmd(
direction=dir,
shape=shape, dtype=dtype,
@@ -711,13 +716,15 @@ class TLContext:
dtype: str = "f16",
) -> "RecvFuture":
"""Non-blocking recv. Returns a future to pass into ``tl.wait``."""
self._emit_dispatch_overhead("ipcq_recv")
cmd = IpcqRecvCmd(
direction=dir,
shape=shape, dtype=dtype,
handle_id=self._next_handle_id(),
blocking=False,
)
# recv_async bypasses _emit (posts via a recv_async future), so it
# charges its own dispatch cost here (ADR-0064 D1).
self._charge_dispatch(cmd)
future = RecvFuture(cmd=cmd)
if self._runner is not None:
self._runner.switch_to_simpy(("recv_async", future))
@@ -749,39 +756,49 @@ class TLContext:
# ADR-0062: composite operand DMA paths still need their inputs
# to be resolved before the composite reads them via PE_SCHEDULER.
self._await_pending(a, b)
# Compute output size based on op
# Compute output geometry based on op.
if op == "gemm" and b is not None:
m, k = a.shape[-2], a.shape[-1]
n = b.shape[-1]
out_dtype = a.dtype
out_shape: tuple[int, ...] = (m, n)
out_nbytes = m * n * self._dtype_bytes(out_dtype)
else:
out_dtype = a.dtype
out_shape = a.shape
out_nbytes = a.nbytes
ops_tuple: tuple[OpSpec, ...] = ()
if epilogue is not None:
head_operands = (a, b) if (op == "gemm" and b is not None) else (a,)
head_spec = OpSpec(
kind=op, scope=Scope.OUTPUT_TILE, operands=head_operands,
extra={
k: v for k, v in (
("acc_dtype", acc_dtype),
("tile_shape", tile_shape),
("math_op", math_op),
) if v is not None
},
)
epi_specs = tuple(self._build_epilogue_spec(e, i)
for i, e in enumerate(epilogue))
ops_tuple = (head_spec, *epi_specs)
# Head op's write-back handle carries the output address (ADR-0065 D1).
out_handle = TensorHandle(
id=self._next_handle_id(),
addr=out_ptr, shape=out_shape, dtype=out_dtype,
nbytes=out_nbytes, space="tcm",
)
# Head op (flat-ops, ADR-0065 D2): GEMM takes named a/b and carries
# m/k/n in extra; MATH takes named a and carries math_op in extra.
if op == "gemm" and b is not None:
head_operands: dict[str, Any] = {"a": a, "b": b}
head_extra: dict[str, Any] = {"m": m, "k": k, "n": n}
else:
head_operands = {"a": a}
head_extra = {}
for _k, _v in (("acc_dtype", acc_dtype),
("tile_shape", tile_shape),
("math_op", math_op)):
if _v is not None:
head_extra[_k] = _v
head_spec = OpSpec(
kind=op, scope=Scope.OUTPUT_TILE,
operands=head_operands, extra=head_extra, out=out_handle,
)
epi_specs = tuple(self._build_epilogue_spec(e, i)
for i, e in enumerate(epilogue or []))
ops_tuple = (head_spec, *epi_specs)
completion = CompletionHandle(id=self._next_completion_id())
self._emit_dispatch_overhead("composite")
self._emit(CompositeCmd(
completion=completion, op=op,
a=a, b=b, out_addr=out_ptr, out_nbytes=out_nbytes,
math_op=math_op, ops=ops_tuple,
))
self._emit(CompositeCmd(completion=completion, ops=ops_tuple))
return completion
@staticmethod
@@ -805,13 +822,13 @@ class TLContext:
f"{', '.join(missing)}"
)
scope = Scope(entry["scope"]) if "scope" in entry else default_scope
operands: list = []
operands: dict = {}
scalar: float | None = None
extra: dict = {}
for f in required:
v = entry[f]
if isinstance(v, TensorHandle):
operands.append(v)
operands[f] = v
elif isinstance(v, (int, float)):
if scalar is None:
scalar = float(v)
@@ -821,7 +838,7 @@ class TLContext:
extra[f] = v
return OpSpec(
kind=kind, scope=scope,
operands=tuple(operands), scalar=scalar, extra=extra,
operands=operands, scalar=scalar, extra=extra,
)
def wait(self, handle: "CompletionHandle | RecvFuture | None" = None) -> Any: