Files
kernbench2/src/kernbench/triton_emu/tl_context.py
T

1199 lines
47 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""TLContext: fake Triton Language module for kernel performance simulation.
Passed as the `tl` parameter to kernel functions. Each API call records a
PeCommand in the internal trace. After the kernel returns, PE_CPU replays
the command list through SimPy.
Kernel code looks like standard Python — no yield, no async:
def my_kernel(a_ptr, b_ptr, out_ptr, tl):
pid = tl.program_id(0)
a = tl.load(a_ptr, shape=(32, 64), dtype="f16")
b = tl.load(b_ptr + pid * stride, shape=(64, 32), dtype="f16")
tl.composite(op="gemm", a=a, b=b, out_ptr=out_ptr)
"""
from __future__ import annotations
import math
from typing import Literal
from kernbench.common.ipcq_types import IpcqRecvCmd, IpcqSendCmd, RecvFuture
from kernbench.common.pe_commands import (
EPILOGUE_OPS,
CompletionHandle,
CompositeCmd,
CopyCmd,
DmaReadCmd,
DmaWriteCmd,
GemmCmd,
MathCmd,
OpSpec,
PeCommand,
PeCpuOverheadCmd,
Scope,
TensorHandle,
WaitCmd,
)
_DTYPE_BYTES: dict[str, int] = {
"f16": 2, "f32": 4, "f64": 8,
"bf16": 2,
"i8": 1, "i16": 2, "i32": 4, "i64": 8,
"u8": 1, "u16": 2, "u32": 4, "u64": 8,
}
class LoadFuture:
"""Lazy ``tl.load`` future (ADR-0062 §D2).
Mirrors ``RecvFuture`` for IPCQ, generalised to HBM loads. Carries
the originating ``DmaReadCmd``, the SimPy completion event (set by
the runner once the DMA has been issued), and a resolved flag.
Consumer ops auto-wait via ``TLContext._await_pending(handle)``,
which yields ``event`` if ``resolved`` is False, then reads the
DMA's bytes into ``handle.data`` and marks the future resolved.
"""
__slots__ = ("cmd", "event", "resolved", "data")
def __init__(self, cmd: DmaReadCmd) -> None:
self.cmd = cmd
self.event: object | None = None # simpy.Event set by runner
self.resolved: bool = False
self.data: object = None
class CompositeFuture:
"""Pending marker on a composite's output handle (ADR-0065 D8).
Carries the composite's ``CompletionHandle`` so a downstream op (via
``_await_pending``) or ``tl.wait`` can await it. Distinct from
``LoadFuture``: composite completion is tracked by completion id in the
runner's pending dict, so awaiting just re-issues a ``WaitCmd`` (a
second wait on an already-drained completion is a harmless no-op).
"""
__slots__ = ("completion", "resolved")
def __init__(self, completion: "CompletionHandle") -> None:
self.completion = completion
self.resolved: bool = False
class _ScratchScope:
"""Context manager that recycles per-tile scratch (ADR-0063 D1).
``__enter__`` snapshots ``_scratch_cursor``; ``__exit__`` restores it,
so every handle allocated inside the ``with``-block has its address
freed for the next iteration. Persistent state (running ``(m, , O)``,
prefetch buffers) lives outside the scope per ADR-0063 D3.
"""
def __init__(self, ctx: "TLContext") -> None:
self._ctx = ctx
self._save: int | None = None
def __enter__(self) -> "_ScratchScope":
self._save = self._ctx._scratch_cursor
return self
def __exit__(self, *exc_info: object) -> bool:
if self._save is not None:
self._ctx._scratch_cursor = self._save
return False
class TLContext:
"""Fake Triton Language context.
Args:
pe_id: program instance index (returned by program_id).
num_programs: total number of program instances.
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__(
self,
pe_id: int = 0,
num_programs: int = 1,
dispatch_cycles: int = 1,
runner: Any = None,
cube_id: int = 0,
num_cubes: int = 1,
scratch_base: int = 0,
scratch_size: int = 1 << 20, # 1 MiB per kernel invocation
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._cost_model = cost_model
self._clock_freq_ghz = clock_freq_ghz
self._commands: list[PeCommand] = []
self._handle_counter = 0
self._completion_counter = 0
self._runner = runner # KernelRunner for greenlet mode (ADR-0020 D3)
# PE-local scratch allocator for math/compute output handles.
# Each binary/unary/reduction op auto-allocates a unique addr from
# this pool so the resulting TensorHandle can be the source of a
# later tl.send / tl.store. Cursor resets on every kernel invocation.
self._scratch_base = scratch_base
self._scratch_size = scratch_size
self._scratch_cursor = 0
def scratch_scope(self) -> _ScratchScope:
"""Per-tile scratch recycling context manager (ADR-0063).
Usage:
with tl.scratch_scope():
s = tl.dot(q, k_t) # per-tile temporaries —
p = tl.softmax(s) # their scratch is rewound
o_j = tl.dot(p, v) # on __exit__
Persistent state (running ``(m, , O)``, lazy-load prefetch
buffers) must be allocated **outside** the scope; only handles
allocated inside are recycled.
"""
return _ScratchScope(self)
def _scratch_alloc(self, nbytes: int) -> int:
"""Allocate a unique scratch address for an output TensorHandle.
Returns 0 if no scratch base was configured (e.g. command-list mode);
in that case the resulting handle has addr=0 and cannot be used as a
send/store source. Greenlet/runner mode always supplies a base.
"""
if self._scratch_base == 0:
return 0
# 16-byte alignment
aligned = (nbytes + 15) & ~15
addr = self._scratch_base + self._scratch_cursor
self._scratch_cursor += aligned
if self._scratch_cursor > self._scratch_size:
raise RuntimeError(
f"TLContext scratch overflow: requested {nbytes}B, "
f"used {self._scratch_cursor}/{self._scratch_size}B"
)
return addr
@property
def commands(self) -> list[PeCommand]:
"""Return the recorded command trace."""
return self._commands
# ── helpers ────────────────────────────────────────────────────
def _next_handle_id(self) -> str:
self._handle_counter += 1
return f"t{self._handle_counter}"
def _next_completion_id(self) -> str:
self._completion_counter += 1
return f"c{self._completion_counter}"
def _dtype_bytes(self, dtype: str) -> int:
return _DTYPE_BYTES.get(dtype, 2)
def _nbytes(self, shape: tuple[int, ...], dtype: str) -> int:
return math.prod(shape) * self._dtype_bytes(dtype)
def _charge_dispatch(self, cmd: PeCommand) -> None:
"""Charge structural PE_CPU dispatch cost (ADR-0064 Rev2 D1).
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._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(
type(cmd), 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,
space: str = "tcm", pinned: bool = False,
) -> TensorHandle:
return TensorHandle(
id=self._next_handle_id(),
addr=addr, shape=shape, dtype=dtype,
nbytes=self._nbytes(shape, dtype),
space=space,
pinned=pinned,
)
def _make_compute_out(
self, shape: tuple[int, ...], dtype: str,
) -> TensorHandle:
"""Allocate an output TensorHandle in PE-local scratch (TCM space).
Used by math/compute ops so the result has a real address that can
be the source of a later send/store. The data field stays None in
Phase 1 — Phase 2 DataExecutor fills the actual ndarray.
"""
nbytes = self._nbytes(shape, dtype)
addr = self._scratch_alloc(nbytes)
return TensorHandle(
id=self._next_handle_id(),
addr=addr, shape=shape, dtype=dtype,
nbytes=nbytes, space="tcm",
# TCM-resident: a downstream composite operand reads it on-chip,
# not via a DMA_READ of its (bit-61) scratch address — same as
# the composite auto-output and recipe scratch handles below.
pinned=True,
)
# ── Reference (no DMA, metadata only) ────────────────────────
def ref(
self, ptr: int, shape: tuple[int, ...], dtype: str = "f16",
) -> TensorHandle:
"""Create a TensorHandle referencing HBM data without issuing DMA.
Used when the scheduler will stream data per-tile (e.g., tensor b
in a composite GEMM) or as a composite's HBM output handle
(ADR-0065 D8). The handle's ``space="hbm"`` — it references HBM
data. No command is generated. (Operand-input DMA stays
``pinned``-based per ADR-0065 D4, so this label does not change
input streaming.)
"""
return self._make_handle(addr=ptr, shape=shape, dtype=dtype,
space="hbm")
# ── Data Movement (blocking, DMA engine) ──────────────────────
def _emit(self, cmd: PeCommand) -> Any:
"""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)
return None
def load(
self, ptr: int, shape: tuple[int, ...], dtype: str = "f16",
) -> TensorHandle:
"""Load tensor from HBM — **lazy** (ADR-0062 §D1/§D2).
Posts the ``DmaReadCmd`` to PE_DMA and returns immediately with a
``TensorHandle`` whose ``pending`` field references a fresh
``LoadFuture``. The actual DMA completion is awaited at the first
consuming op (``tl.dot``, MATH, ``tl.store``, ``tl.send``,
``tl.copy_to``, ``tl.composite``) via ``_await_pending``.
Command-list mode: emits the DmaReadCmd to ``self._commands``,
attaches a LoadFuture for structural compatibility (its event
stays None — no engine, no SimPy event to wait on).
"""
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
# *final* handle. This guarantees ``handle.pending.cmd.handle is handle``.
future = LoadFuture.__new__(LoadFuture)
future.event = None
future.resolved = False
future.data = None
handle = TensorHandle(
id=self._next_handle_id(),
addr=ptr, shape=shape, dtype=dtype, nbytes=nbytes,
data=None, space="hbm", pinned=True, pending=future,
)
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.
self._runner.switch_to_simpy(("load_issue", future))
else:
self._commands.append(cmd)
return handle
def _await_pending(self, *handles: TensorHandle | None) -> None:
"""Auto-wait at first use (ADR-0062 §D2).
For each handle carrying an unresolved ``LoadFuture``, yield on
the DMA completion event (greenlet → runner). Command-list mode
is a no-op.
"""
if self._runner is None:
return
for h in handles:
if h is None:
continue
pending = getattr(h, "pending", None)
if pending is None or pending.resolved:
continue
if isinstance(pending, CompositeFuture):
# Composite output: wait on its completion (ADR-0065 D8).
self._emit(WaitCmd(handle=pending.completion))
pending.resolved = True
else:
self._runner.switch_to_simpy(("load_await", h))
def store(self, ptr: int, handle: TensorHandle) -> None:
"""Store tensor from TCM to HBM."""
self._await_pending(handle)
cmd = DmaWriteCmd(handle=handle, dst_addr=ptr, nbytes=handle.nbytes)
self._emit(cmd)
def copy_to(self, dst: TensorHandle, src: TensorHandle) -> None:
"""TCM-to-TCM byte copy (ADR-0063 §D3.1).
Writes ``src``'s bytes into ``dst``'s address. Both handles must
live in TCM and have matching shape and dtype. Symmetric to
``tl.store`` (which targets HBM) but stays on-chip so it doesn't
emit a DMA entry into op_log.
Used inside ``tl.scratch_scope()`` to persist a scoped result —
typically an updated running ``(m, , O)`` — to an outside-scope
(persistent) handle so its bytes survive the scope's ``__exit__``
cursor rewind (ADR-0063 §D3 two-arena pattern).
"""
if src.shape != dst.shape:
raise ValueError(
f"tl.copy_to: shape mismatch — src.shape={src.shape} "
f"vs dst.shape={dst.shape}"
)
if src.dtype != dst.dtype:
raise ValueError(
f"tl.copy_to: dtype mismatch — src.dtype={src.dtype!r} "
f"vs dst.dtype={dst.dtype!r}"
)
if dst.space != "tcm":
raise ValueError(
f"tl.copy_to: dst must be in TCM (got space={dst.space!r}); "
"writes to HBM go through tl.store"
)
if src.space != "tcm":
raise ValueError(
f"tl.copy_to: src must be in TCM (got space={src.space!r}); "
"reads from HBM go through tl.load"
)
self._await_pending(src)
self._emit(CopyCmd(src=src, dst=dst, nbytes=src.nbytes))
# ── GEMM Engine (blocking) ────────────────────────────────────
def dot(self, a: TensorHandle, b: TensorHandle) -> TensorHandle:
"""Matrix multiply: out = a @ b. Both operands must be in TCM.
a: (M, K), b: (K, N) → out: (M, N)
"""
if len(a.shape) < 2 or len(b.shape) < 2:
raise ValueError("dot requires 2D tensors")
m, k = a.shape[-2], a.shape[-1]
k2, n = b.shape[-2], b.shape[-1]
if k != k2:
raise ValueError(f"dot shape mismatch: a.K={k} != b.K={k2}")
out_shape = (*a.shape[:-2], m, n)
out_dtype = a.dtype
out = self._make_compute_out(shape=out_shape, dtype=out_dtype)
self._await_pending(a, b)
self._emit(GemmCmd(a=a, b=b, out=out, m=m, k=k, n=n))
return out
# ── MATH Engine: unary (blocking) ─────────────────────────────
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(MathCmd(op=op, inputs=(x,), out=out))
return out
def exp(self, x: TensorHandle) -> TensorHandle:
return self._unary_math("exp", x)
def log(self, x: TensorHandle) -> TensorHandle:
return self._unary_math("log", x)
def sqrt(self, x: TensorHandle) -> TensorHandle:
return self._unary_math("sqrt", x)
def abs(self, x: TensorHandle) -> TensorHandle:
return self._unary_math("abs", x)
def sigmoid(self, x: TensorHandle) -> TensorHandle:
return self._unary_math("sigmoid", x)
def cos(self, x: TensorHandle) -> TensorHandle:
return self._unary_math("cos", x)
def sin(self, x: TensorHandle) -> TensorHandle:
return self._unary_math("sin", x)
# ── MATH Engine: reduction (blocking) ─────────────────────────
def _reduction(
self, op: str, x: TensorHandle, axis: int,
) -> TensorHandle:
out_shape = list(x.shape)
out_shape[axis] = 1
out = self._make_compute_out(shape=tuple(out_shape), dtype=x.dtype)
self._await_pending(x)
self._emit(MathCmd(op=op, inputs=(x,), out=out, axis=axis))
return out
def sum(self, x: TensorHandle, axis: int) -> TensorHandle:
return self._reduction("sum", x, axis)
def max(self, x: TensorHandle, axis: int) -> TensorHandle:
return self._reduction("max", x, axis)
def min(self, x: TensorHandle, axis: int) -> TensorHandle:
return self._reduction("min", x, axis)
# ── MATH Engine: binary (blocking) ────────────────────────────
def _binary_math(
self, op: str, a: TensorHandle, b: TensorHandle,
) -> TensorHandle:
out = self._make_compute_out(shape=a.shape, dtype=a.dtype)
self._await_pending(a, b)
self._emit(MathCmd(op=op, inputs=(a, b), out=out))
return out
def where(
self, cond: TensorHandle, a: TensorHandle, b: TensorHandle,
) -> TensorHandle:
out = self._make_compute_out(shape=a.shape, dtype=a.dtype)
self._await_pending(cond, a, b)
self._emit(MathCmd(op="where", inputs=(cond, a, b), out=out))
return out
def maximum(self, a: TensorHandle, b: TensorHandle) -> TensorHandle:
"""Element-wise max of two tensors (real Triton: tl.maximum)."""
return self._binary_math("maximum", a, b)
def minimum(self, a: TensorHandle, b: TensorHandle) -> TensorHandle:
"""Element-wise min of two tensors (real Triton: tl.minimum)."""
return self._binary_math("minimum", a, b)
def fma(
self, a: TensorHandle, b: TensorHandle, c: TensorHandle,
) -> TensorHandle:
"""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(MathCmd(op="fma", inputs=(a, b, c), out=out))
return out
def clamp(
self,
x: TensorHandle,
min: TensorHandle,
max: TensorHandle,
) -> TensorHandle:
"""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(MathCmd(op="clamp", inputs=(x, min, max), out=out))
return out
def softmax(self, x: TensorHandle, axis: int = -1) -> TensorHandle:
"""Numerically-stable softmax along ``axis`` (real Triton: tl.softmax).
Implemented as a single MathCmd (op="softmax") so timing accounts
for one MATH dispatch; Phase 2 DataExecutor expands it to the
canonical (x - max) → exp → sum → div sequence.
"""
out = self._make_compute_out(shape=x.shape, dtype=x.dtype)
self._await_pending(x)
self._emit(MathCmd(op="softmax", inputs=(x,), out=out, axis=axis))
return out
# ── Scalar helpers (real Triton: tl.cdiv etc.) ────────────────
@staticmethod
def cdiv(a: int, b: int) -> int:
"""Ceiling division: (a + b - 1) // b (real Triton: tl.cdiv).
Used by host/kernel grid math; not a tensor op, so no MathCmd
is emitted. Mirrors triton.cdiv.
"""
return -(-int(a) // int(b))
# ── Index / Scalar (PE_CPU, no engine) ────────────────────────
def program_id(self, axis: int = 0) -> int:
"""Return program instance index (ADR-0022).
axis=0: local PE id within cube.
axis=1: cube id.
"""
if axis == 1:
return self._cube_id
return self._pe_id
def num_programs(self, axis: int = 0) -> int:
"""Return total number of program instances (ADR-0022).
axis=0: num PEs per cube.
axis=1: num cubes.
"""
if axis == 1:
return self._num_cubes
return self._num_programs
def arange(self, start: int, end: int, dtype: str = "i32") -> TensorHandle:
"""Create index range tensor in TCM."""
n = end - start
return self._make_handle(addr=0, shape=(n,), dtype=dtype)
def zeros(self, shape: tuple[int, ...], dtype: str = "f16") -> TensorHandle:
"""Create zero-filled tensor in TCM."""
return self._make_handle(addr=0, shape=shape, dtype=dtype)
def full(
self, shape: tuple[int, ...], value: float | int, dtype: str = "f16",
) -> TensorHandle:
"""Create constant-filled tensor in TCM."""
return self._make_handle(addr=0, shape=shape, dtype=dtype)
# ── Metadata (no compute, no DMA) ─────────────────────────────
def trans(self, x: TensorHandle) -> TensorHandle:
"""Transpose — shape change only, no command generated."""
if len(x.shape) < 2:
raise ValueError("trans requires at least 2D tensor")
new_shape = (*x.shape[:-2], x.shape[-1], x.shape[-2])
return TensorHandle(
id=x.id, addr=x.addr, shape=new_shape,
dtype=x.dtype, nbytes=x.nbytes, data=x.data,
)
# ── IPCQ (CCL) collective primitives (ADR-0023 D4) ────────────
def send(
self,
dir: str,
src: TensorHandle | None = None,
*,
src_addr: int | None = None,
nbytes: int | None = None,
shape: tuple[int, ...] | None = None,
dtype: str = "f16",
space: str = "tcm",
) -> None:
"""Send tensor data to the peer in the given direction.
Two calling forms:
tl.send(dir, handle) # use handle's metadata
tl.send(dir, src_addr=..., nbytes=..., shape=..., dtype=..., space=...)
Blocking: returns when PE_IPCQ has accepted the request and
forwarded the IpcqDmaToken to PE_DMA. Backpressure may apply.
"""
if src is not None:
src_addr = src.addr
nbytes = src.nbytes
shape = src.shape
dtype = src.dtype
space = getattr(src, "space", space)
if src_addr is None or nbytes is None or shape is None:
raise ValueError("tl.send: provide either a TensorHandle or src_addr/nbytes/shape")
# ADR-0062: if the source is a lazy-loaded handle, await first so
# the data snapshot below sees the real bytes.
self._await_pending(src)
# Carry the handle's .data snapshot (if available). When the source
# is a recv slot, .data holds the numpy array that was read from
# MemoryStore at recv-time. This prevents a Phase 1 race where a
# 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
cmd = IpcqSendCmd(
direction=dir,
src_addr=src_addr, src_space=space,
nbytes=nbytes, shape=shape, dtype=dtype,
handle_id=self._next_handle_id(),
data=handle_data,
)
self._emit(cmd)
def recv(
self,
dir: str | None = None,
shape: tuple[int, ...] = (),
dtype: str = "f16",
space: str = "tcm",
dst_addr: int | None = None,
dst_space: str | None = None,
) -> TensorHandle:
"""Receive tensor data from a peer.
Args:
dir: specific direction (e.g. "W"), or None for round-robin.
shape, dtype: expected tensor metadata.
dst_addr / dst_space: if both are provided, the slot data is
copied to (dst_space, dst_addr) before the handle is
returned ("copy_to_dst" mode). Otherwise the slot address
is returned directly ("return_slot" mode).
Returns:
TensorHandle pointing to the slot (or dst) where the data has
arrived. In greenlet/runner mode, ``handle.data`` carries the
actual ndarray; in command-list mode the handle is a placeholder.
"""
if dst_addr is not None and dst_space is not None:
cmd = IpcqRecvCmd(
direction=dir,
shape=shape, dtype=dtype,
handle_id=self._next_handle_id(),
recv_mode="copy_to_dst",
dst_addr=dst_addr, dst_space=dst_space,
)
else:
cmd = IpcqRecvCmd(
direction=dir,
shape=shape, dtype=dtype,
handle_id=self._next_handle_id(),
)
result = self._emit(cmd)
if isinstance(result, dict):
slot_addr = int(result.get("src_addr", 0))
slot_space = str(result.get("src_space", "tcm"))
data = result.get("data")
return TensorHandle(
id=self._next_handle_id(),
addr=slot_addr,
shape=shape,
dtype=dtype,
nbytes=self._nbytes(shape, dtype),
data=data,
space=slot_space,
# On-chip (tcm/sram) IPCQ slot: read in place as a composite
# operand, not DMA_READ of its non-HBM slot address. An HBM
# slot is a valid PA, so it stays unpinned (DMA streams it).
pinned=slot_space != "hbm",
)
return self._make_handle(addr=0, shape=shape, dtype=dtype)
def recv_no_consume(
self,
dir: str | None = None,
shape: tuple[int, ...] = (),
dtype: str = "f16",
) -> TensorHandle:
"""DIAGNOSTIC ONLY — recv that blocks for arrival but skips slot read.
Same blocking semantics as ``tl.recv``: the kernel waits until
the payload has landed in the IPCQ slot. Differs from ``tl.recv``
by skipping the slot-read latency charge (slot-IO + PE↔bank
fabric drain) on DST.
This entry point exists solely so the pe2pe overview plot can
draw an apples-to-apples comparison against ``tl.store`` (a
one-sided fabric write that pays no read on DST). Production
kernels MUST use ``tl.recv`` — they need to consume the data
they receive. This API is segregated from ``tl.recv`` so the
diagnostic flag can never accidentally be set in real workloads.
"""
cmd = IpcqRecvCmd(
direction=dir,
shape=shape, dtype=dtype,
handle_id=self._next_handle_id(),
consume=False,
)
result = self._emit(cmd) # type: ignore[arg-type]
if isinstance(result, dict):
slot_addr = int(result.get("src_addr", 0))
slot_space = str(result.get("src_space", "tcm"))
return TensorHandle(
id=self._next_handle_id(),
addr=slot_addr,
shape=shape,
dtype=dtype,
nbytes=self._nbytes(shape, dtype),
data=None,
space=slot_space,
pinned=slot_space != "hbm",
)
return self._make_handle(addr=0, shape=shape, dtype=dtype)
def recv_async(
self,
dir: str,
shape: tuple[int, ...] = (),
dtype: str = "f16",
) -> "RecvFuture":
"""Non-blocking recv. Returns a future to pass into ``tl.wait``."""
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))
return future
# ── Composite + Control ───────────────────────────────────────
def composite(
self,
op: Literal["gemm", "math"],
a: TensorHandle | None = None,
b: TensorHandle | None = None,
math_op: str | None = None,
*,
out: TensorHandle | None = None,
out_ptr: int | None = None,
prologue: list[dict] | None = None,
epilogue: list[dict] | None = None,
acc_dtype: str | None = None,
tile_shape: tuple[int, int, int] | None = None,
) -> TensorHandle:
"""Submit a composite command (non-blocking, tiled pipeline).
``prologue`` (ADR-0065 D5) is an ordered list of recipe dicts — each
with an ``"op"`` naming a ``RECIPE_DESCRIPTORS`` entry plus its
operands. TLContext expands each into flat MATH OpSpecs that run
before the head op and auto-binds the recipe's primary output into
the head GEMM's ``a`` operand (D6.6). ``epilogue`` is an ordered list
of ``EPILOGUE_OPS`` dicts.
``out`` (ADR-0065 D8) is a ``TensorHandle``: an HBM ref
(``tl.ref(addr, shape)`` → ``space="hbm"`` → DMA_WRITE) or an
in-place TCM handle (e.g. a recipe accumulator → STORE only).
``out_ptr`` (int) is shorthand for an HBM output — equivalent to
``out=tl.ref(out_ptr, <output shape>)``. If neither is given, a TCM
scratch is auto-allocated (chainable, like ``tl.dot``).
Returns the **output TensorHandle** — auto-awaited by downstream ops
(and by ``tl.wait(handle)``) via its ``pending`` field.
"""
# 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)
# Prologue recipes → flat MATH OpSpecs + primary-output handle (D5).
prologue_ops: tuple[OpSpec, ...] = ()
primary_out: TensorHandle | None = None
rw_handles: tuple[TensorHandle, ...] = ()
if prologue:
prologue_ops, primary_out, rw_handles = self._expand_prologue(prologue)
# Await any prologue operand that is a pending composite output
# (e.g. the score tile produced by a prior #1 composite) so it is
# computed before this composite reads it (RAW across composites).
for _item in prologue:
self._await_pending(
*[v for v in _item.values() if isinstance(v, TensorHandle)]
)
# Auto-bind (D6.6): the recipe's primary output fills the head GEMM's
# `a` unless the kernel supplied `a` explicitly (then it's ambiguous).
if op == "gemm" and primary_out is not None:
if a is not None:
raise ValueError(
"auto-bind conflict: head GEMM operand 'a' was provided "
"explicitly but the prologue recipe declares a primary_out "
"(ADR-0065 D6.6) — omit 'a' to let the recipe bind it"
)
a = primary_out
if a is None:
raise ValueError(
"composite requires operand 'a' (or a prologue recipe with a "
"primary_out)"
)
# Output geometry.
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
# Write-back handle (ADR-0065 D8). Explicit `out` (HBM ref or an
# in-place TCM accumulator) wins; `out_ptr` is HBM shorthand; else
# auto-allocate a TCM scratch — the chainable default, like tl.dot.
# pinned=True on the auto output: it is TCM-resident, so a downstream
# op reading it needs no DMA_READ.
if out is not None:
out_handle = out
elif out_ptr is not None:
out_handle = TensorHandle(
id=self._next_handle_id(),
addr=out_ptr, shape=out_shape, dtype=out_dtype,
nbytes=out_nbytes, space="hbm",
)
else:
out_handle = TensorHandle(
id=self._next_handle_id(),
addr=self._scratch_alloc(out_nbytes),
shape=out_shape, dtype=out_dtype, nbytes=out_nbytes,
space="tcm", pinned=True,
)
# 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 = (*prologue_ops, head_spec, *epi_specs)
self._validate_composite_ops(ops_tuple, prologue_ops)
completion = CompletionHandle(id=self._next_completion_id())
self._emit(CompositeCmd(
completion=completion, ops=ops_tuple, rw_handles=rw_handles,
))
# Attach the completion to the output handle so downstream ops (and
# tl.wait) auto-await it (ADR-0065 D8 / ADR-0062 lazy pattern), and
# return it so the result chains like tl.dot's (frozen dataclass →
# set pending via object.__setattr__).
object.__setattr__(out_handle, "pending", CompositeFuture(completion))
return out_handle
def _expand_prologue(
self, prologue: list[dict],
) -> tuple[tuple[OpSpec, ...], TensorHandle | None, tuple[TensorHandle, ...]]:
"""Expand prologue recipe items into flat MATH OpSpecs (ADR-0065 D5).
Returns ``(math_ops, primary_out_handle, rw_handles)``. PE_SCHEDULER
never sees the recipe — only the flat ops.
"""
from kernbench.triton_emu.tl_recipes import (
PRIMARY_OUT_SLOT, RECIPE_DESCRIPTORS,
)
math_ops: list[OpSpec] = []
primary_out: TensorHandle | None = None
rw_handles: list[TensorHandle] = []
for item in prologue:
if not isinstance(item, dict) or "op" not in item:
raise ValueError("prologue entry must be a dict with an 'op' key")
name = item["op"]
recipe = RECIPE_DESCRIPTORS.get(name)
if recipe is None:
known = ", ".join(sorted(RECIPE_DESCRIPTORS))
raise ValueError(
f"unknown prologue recipe {name!r} (known: {known})"
)
# Bind recipe operands from the item.
slots: dict[str, TensorHandle] = {}
for opd_name in recipe.operands:
if opd_name not in item:
raise ValueError(
f"prologue recipe {name!r} missing operand {opd_name!r}"
)
slots[opd_name] = item[opd_name]
# Expand engine_seq sequentially — later ops read earlier dsts.
for eop in recipe.engine_seq:
operands: dict[str, Any] = {}
for key in ("src", "src_a", "src_b", "src_c"):
slot_name = getattr(eop, key)
if slot_name is None:
continue
if slot_name not in slots:
raise ValueError(
f"recipe {name!r} op {eop.op_kind!r} reads unbound "
f"slot {slot_name!r}"
)
operands[key] = slots[slot_name]
dst_handle = self._resolve_recipe_dst(eop, recipe, slots, operands)
extra: dict[str, Any] = {}
if eop.reduce_axis is not None:
extra["reduce_axis"] = eop.reduce_axis
if eop.bcast_axis is not None:
extra["bcast_axis"] = eop.bcast_axis
math_ops.append(OpSpec(
kind=eop.op_kind, scope=Scope.KERNEL,
operands=operands, extra=extra, out=dst_handle,
))
slots[eop.dst] = dst_handle
if recipe.primary_out is not None:
primary_out = slots.get(PRIMARY_OUT_SLOT)
for opd_name, rw in recipe.operands.items():
if rw == "RW":
rw_handles.append(slots[opd_name])
return tuple(math_ops), primary_out, tuple(rw_handles)
def _resolve_recipe_dst(
self, eop: Any, recipe: Any, slots: dict, operands: dict,
) -> TensorHandle:
"""Resolve an EngineOp's dst to a handle: reuse an existing slot
(RW operand / already-allocated), else allocate fresh TCM scratch
with an inferred shape."""
from kernbench.triton_emu.tl_recipes import PRIMARY_OUT_SLOT
dst_name = eop.dst
if dst_name in slots:
return slots[dst_name] # write in place
if recipe.primary_out is not None and dst_name == PRIMARY_OUT_SLOT:
ref = slots[recipe.primary_out.from_shape]
shape, dtype = ref.shape, ref.dtype
elif eop.reduce_axis is not None:
# Reductions keep the reduced axis as size 1 (keepdims), so the
# result broadcasts against the running (m, l) — e.g. rmax of
# s=(G,TILE) over axis -1 → (G,1), not (G,) (ADR-0065 N2).
ref = next(iter(operands.values()))
dims = list(ref.shape)
if len(dims) > 1:
ax = eop.reduce_axis % len(dims)
dims[ax] = 1
shape = tuple(dims)
dtype = ref.dtype
else:
ref = next(iter(operands.values()))
shape, dtype = ref.shape, ref.dtype
nbytes = self._nbytes(shape, dtype)
addr = self._scratch_alloc(nbytes)
# pinned=True: recipe scratch / primary-out live in TCM already, so
# the head GEMM's auto-bound `a` (= primary-out P) is consumed in
# place — no DMA_READ (ADR-0065 P3, pinned-based DMA decision).
return TensorHandle(
id=self._next_handle_id(),
addr=addr, shape=shape, dtype=dtype, nbytes=nbytes,
space="tcm", pinned=True,
)
@staticmethod
def _validate_composite_ops(
ops: tuple[OpSpec, ...], prologue_ops: tuple[OpSpec, ...],
) -> None:
"""Emit-time invariants (ADR-0065 D6.1 / D6.7).
D6.1 (GEMM count ≤ 1) applies to the whole composite. D6.7 (MATH
operands TCM-only) applies to **prologue recipe ops only** — the
head op (gemm *or* math) and epilogue ops keep the existing
DMA-staged-from-HBM behavior (a math head's `a` is streamed by
PE_SCHEDULER, like a GEMM operand).
"""
gemm_count = sum(1 for o in ops if o.kind == "gemm")
if gemm_count > 1:
raise ValueError(
f"composite must carry at most 1 GEMM op (ADR-0065 D6.1); "
f"got {gemm_count}"
)
for o in prologue_ops:
for h in o.operands.values():
if getattr(h, "space", "tcm") != "tcm":
raise ValueError(
f"MATH operand must be TCM-resident (ADR-0065 D6.7); "
f"prologue op {o.kind!r} operand has space="
f"{getattr(h, 'space', None)!r}"
)
@staticmethod
def _build_epilogue_spec(entry: dict, idx: int) -> OpSpec:
if not isinstance(entry, dict) or "op" not in entry:
raise ValueError(
f"epilogue[{idx}]: each entry must be a dict with an 'op' key"
)
kind = entry["op"]
if kind not in EPILOGUE_OPS:
known = ", ".join(sorted(EPILOGUE_OPS))
raise ValueError(
f"epilogue[{idx}]: unknown op {kind!r} "
f"(known ops: {known})"
)
required, default_scope = EPILOGUE_OPS[kind]
missing = [f for f in required if f not in entry]
if missing:
raise ValueError(
f"epilogue[{idx}] op {kind!r} missing required field(s): "
f"{', '.join(missing)}"
)
scope = Scope(entry["scope"]) if "scope" in entry else default_scope
operands: dict = {}
scalar: float | None = None
extra: dict = {}
for f in required:
v = entry[f]
if isinstance(v, TensorHandle):
operands[f] = v
elif isinstance(v, (int, float)):
if scalar is None:
scalar = float(v)
else:
extra[f] = v
else:
extra[f] = v
return OpSpec(
kind=kind, scope=scope,
operands=operands, scalar=scalar, extra=extra,
)
def wait(self, handle: "CompletionHandle | RecvFuture | None" = None) -> Any:
"""Wait for a composite, a recv future, or all pending composites.
- ``CompletionHandle`` (or None): wait for composite completion.
- ``RecvFuture``: wait for a non-blocking ``recv_async`` to finish.
Returns the resolved ``TensorHandle``.
"""
if isinstance(handle, RecvFuture):
if handle.resolved:
return handle.result
if self._runner is None:
raise RuntimeError(
"tl.wait(RecvFuture) requires runner mode (greenlet)"
)
result_dict = self._runner.switch_to_simpy(("recv_wait", handle))
slot_addr = int(result_dict.get("src_addr", 0))
slot_space = str(result_dict.get("src_space", "tcm"))
data = result_dict.get("data")
th = TensorHandle(
id=self._next_handle_id(),
addr=slot_addr,
shape=handle.cmd.shape,
dtype=handle.cmd.dtype,
nbytes=self._nbytes(handle.cmd.shape, handle.cmd.dtype),
data=data,
space=slot_space,
pinned=slot_space != "hbm",
)
handle.resolved = True
handle.result = th
return th
# Composite output handle (ADR-0065 D8): wait on its completion.
if isinstance(handle, TensorHandle):
pending = getattr(handle, "pending", None)
if isinstance(pending, CompositeFuture):
if not pending.resolved:
self._emit(WaitCmd(handle=pending.completion))
pending.resolved = True
return None
# Composite path (CompletionHandle or None — existing behaviour).
self._emit(WaitCmd(handle=handle))
return None
def cycles(self, n: int) -> None:
"""Declare PE_CPU scalar execution overhead (cycles)."""
self._emit(PeCpuOverheadCmd(cycles=n))
# ── TensorHandle arithmetic operators ─────────────────────────────
# Enables: a + b, a * b, a - b, a / b in kernel code.
# Each creates a MathCmd via a module-level helper that requires a
# TLContext. We attach the context to handles via a closure approach.
def _enable_tensor_ops() -> None:
"""Patch TensorHandle with arithmetic operators.
Called once at module load. Operators create MathCmd entries via
a thread-local TLContext reference set during kernel execution.
"""
import threading
_local = threading.local()
def set_active_context(ctx: TLContext | None) -> None:
_local.ctx = ctx
def get_active_context() -> TLContext:
ctx = getattr(_local, "ctx", None)
if ctx is None:
raise RuntimeError("TensorHandle ops require an active TLContext")
return ctx
def _binop(op: str):
def method(self: TensorHandle, other: TensorHandle) -> TensorHandle:
ctx = get_active_context()
return ctx._binary_math(op, self, other)
return method
# Patch TensorHandle class with operators
TensorHandle.__add__ = _binop("add") # type: ignore[attr-defined]
TensorHandle.__sub__ = _binop("sub") # type: ignore[attr-defined]
TensorHandle.__mul__ = _binop("mul") # type: ignore[attr-defined]
TensorHandle.__truediv__ = _binop("div") # type: ignore[attr-defined]
# Expose context management
TLContext._set_active = staticmethod(set_active_context) # type: ignore[attr-defined]
TLContext._get_active = staticmethod(get_active_context) # type: ignore[attr-defined]
_enable_tensor_ops()
def run_kernel(
kernel_fn,
tl_ctx: TLContext,
*args,
**kwargs,
) -> list[PeCommand]:
"""Execute a kernel function with the given TLContext and return commands.
Sets tl_ctx as the active context for TensorHandle operators,
calls the kernel, then clears the context.
"""
TLContext._set_active(tl_ctx) # type: ignore[attr-defined]
try:
kernel_fn(*args, tl=tl_ctx, **kwargs)
finally:
TLContext._set_active(None) # type: ignore[attr-defined]
return tl_ctx.commands