ADR-0064 Revision 2: replace per-op-type calibration table with a structural formula `FIXED_PER_CMD + cmd.logical_bytes × R`, defaults anchored at typical composite ≈ 45 ns. Topology yaml override under `pe_cost_model:` block; logical_bytes property per PE command. ADR-0065: implement ADR-0060 §5.6 / §8 item 4 carve-out as a flat-ops CompositeCmd (no head/epilogue structural fields — position + scope drives placement) + first stateful recipe `softmax_merge` (MATH-only 8-step). RECIPE_DESCRIPTORS lives in TLContext-adjacent module only; PE_SCHEDULER stays recipe-free and auto-inserts DMAs from operand `space`. Strict-FIFO RW hazard tracker; ≤1 GEMM per composite invariant. User-facing `tl.composite(prologue=[...], op=, epilogue=[...])` API preserved; existing benches unchanged (meaning-preserving refactor). DDD-0065: implementation-ready phased plan (P0=ADR-0064 Rev2 first, P1-P6 for ADR-0065), file plan, recipe engine sequence, scheduler plan-gen algorithm, RW tracker design, test matrix. Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
17 KiB
ADR-0065: Flat-ops CompositeCmd + first stateful recipe (softmax_merge)
Status
Proposed (verification gated on ADR-0064 Revision 2 land)
Supporting ADR for ADR-0060 (AHBM GQA Fused Attention). Implements the §5.6 / §8 item 4 carve-out exactly: opt2 of decode = (#1 existing GEMM composite for Q·Kᵀ) + (#2 a single composite carrying softmax + P·V + online-softmax merge). The "ex_composite" / "flash_pv_merge" name referenced in ADR-0060 §5.6 is retired — this ADR fixes the canonical shape using the existing
tl.compositeentry with two structural additions: (a) a flatopstuple, and (b) a first stateful recipesoftmax_mergethat lowers to a sequence of MATH micro-ops.
Context
What ADR-0060 left to this ADR
ADR-0060 §5.6 shipped opt3 (software pipelining) as the recommended now, and explicitly carved out opt2 — fewer per-tile dispatches, K-before-V DMA priority, measurable only after ADR-0064's cost model — as the follow-up.
§8 item 4 sized the carve-out: "if revisited, split it into two
composites — #1 = Q·Kᵀ (existing composite + scale), #2 = softmax +
P·V + online-softmax accumulator merge". This ADR implements #2.
Why composite needs a structural change
Today's CompositeCmd has the shape (op, a, b, out_addr, ops=(head, *epi)) with implicit conventions:
opselects the head engine (gemm | math)ops[0]is the head OpSpec,ops[1:]are epilogue OpSpecs that run after the head per-output-tile or per-K-tile (scope-determined)
For decode opt2 we need MATH ops to run before the head GEMM (the
online-softmax m, l, O update + emit P for the GEMM input). The
current shape has no slot for that.
Two failed paths considered earlier:
- Embedding softmax_merge as an epilogue of #1 (Q·Kᵀ) — wrong: it
needs to read
Sj(= #1's output) and writes a freshPconsumed by the next GEMM (#2's P·V); circular. - A single mixed-engine recipe
flash_pv_mergeabsorbing MATH + GEMM + MATH — violates the engine boundary, complicates PE_SCHEDULER.
The clean shape is: drop the head/epilogue distinction from the
command format. Make ops a flat ordered tuple; let each op's
scope + position determine its tile-loop placement. The "prologue"
concept stays in the user-facing tl.composite(...) API as
ergonomics — but the HW command never sees it.
Why a recipe (not a generic op list)
Decode opt2's #2 MATH chain is 8 steps with a fixed structure (online
softmax merge). Letting the kernel author wire those 8 steps would be
both verbose and a hazard surface (correctness depends on step order).
A recipe — a single named op (softmax_merge) that TLContext expands
to the 8 steps with addresses pre-filled — keeps the kernel concise and
the structure under one source of truth. The recipe table lives in
TLContext (the compiler analog); PE_SCHEDULER never reads it. See
boundary in D5.
Decision
D1. CompositeCmd is a flat ordered op list
@dataclass(frozen=True)
class CompositeCmd:
completion: CompletionHandle
ops: tuple[OpSpec, ...] # ordered MATH/GEMM ops
rw_handles: tuple[TensorHandle, ...] = () # cross-composite hazard
data_op: bool = True
@property
def logical_bytes(self) -> int: ... # ADR-0064 D2
The previous (op, a, b, out_addr, out_nbytes, math_op) fields are
absorbed into ops. The first OpSpec in ops (if kind == "gemm") is
the head GEMM; preceding OpSpecs are pre-loop ops, following OpSpecs
are post-loop / per-tile epilogue ops depending on scope.
D2. OpSpec carries named operands
@dataclass(frozen=True)
class OpSpec:
kind: str # "gemm" | "rmax" | ...
scope: Scope # KERNEL | K_TILE | OUTPUT_TILE
operands: dict[str, TensorHandle] # named — see D5
extra: dict[str, Any] # scalars, axes, m/k/n, etc.
out: TensorHandle | None = None # explicit write-back handle
@property
def logical_bytes(self) -> int: ...
Named operands let PE_SCHEDULER route inputs to engine ports (e.g., GEMM
takes operands["a"], operands["b"]) without positional convention.
D3. Position + scope determines tile-loop placement
PE_SCHEDULER scans cmd.ops for the GEMM op (≤1 per CompositeCmd, see
D6). Calling its index g:
| OpSpec position | scope | Placement in plan |
|---|---|---|
0 .. g-1 |
KERNEL | pre-loop stages (run once before tile loop) |
g |
KERNEL | head GEMM (drives tile loop with extra["m"], extra["k"], extra["n"]) |
g+1 .. |
K_TILE | per K-tile epilogue (inside accumulation loop) |
g+1 .. |
OUTPUT_TILE | per (m, n) tile epilogue (after K accumulation) |
g+1 .. |
KERNEL | post-loop stages (run once after tile loop) |
A composite with no GEMM op (e.g., MATH-only composite) treats all ops as KERNEL-scope sequential.
D4. PE_SCHEDULER auto-inserts DMAs from operand space
PE_SCHEDULER scans the GEMM's operands and out. For each:
space == "hbm"→ emit DMA_READ Stage (or DMA_WRITE forout) before the FETCH/GEMM Stagesspace == "tcm"→ no DMA Stage, operand consumed in place
This is already partially modelled today via a_pinned/b_pinned; D4
makes the rule uniform and visible: the kernel never emits explicit
DMA cmds for composite operands; PE_SCHEDULER infers them from
handles. tl.ref(addr, shape) returns a handle with space="hbm";
tl.zeros / tl.full / tl.load outputs return space="tcm".
Non-GEMM ops (MATH chain in prologue / epilogue) have TCM-resident
operands (kernel allocates via tl.zeros / tl.full); no DMA inserted
for them.
D5. RECIPE_DESCRIPTORS — TLContext-internal, NOT in pe_commands.py
Recipes live in a new TLContext-adjacent module
(src/kernbench/triton_emu/tl_recipes.py). PE_SCHEDULER does not
import it. The first recipe:
@dataclass(frozen=True)
class RecipeDescriptor:
# Operand contract: name → "R" (read-only) | "RW" (read-write)
operands: dict[str, str]
# Type / shape rule for the implicit primary output (auto-bound to
# head's first matrix operand). None means recipe emits no value.
primary_out: PrimaryOutSpec | None
# Single-shot prologue (full reduction) vs tile-aligned execution.
tile_alignment: Literal["single_shot", "tile_aligned"]
# Internal scratch budget (bytes), function of operand dims.
internal_scratch_bytes_fn: Callable[..., int]
# Engine micro-op sequence: TLContext expands this into flat OpSpecs.
engine_seq: tuple[EngineOp, ...]
RECIPE_DESCRIPTORS["softmax_merge"] = RecipeDescriptor(
operands={"s": "R", "m": "RW", "l": "RW", "O": "RW"},
primary_out=PrimaryOutSpec(
from_shape="s", from_dtype="s", transform="identity",
),
tile_alignment="single_shot",
internal_scratch_bytes_fn=lambda G, TILE, d, bpe: bpe * (
G + G + G + G * TILE + G # m_loc + m_new + corr + P + l_loc
),
engine_seq=(
EngineOp("MATH", "rmax", src="s", dst="m_loc", reduce_axis=-1),
EngineOp("MATH", "max_elem", src_a="m", src_b="m_loc", dst="m_new"),
EngineOp("MATH", "exp_diff", src_a="m", src_b="m_new", dst="corr"),
EngineOp("MATH", "exp_diff", src_a="s", src_b="m_new", dst="P", bcast_axis=0),
EngineOp("MATH", "rsum", src="P", dst="l_loc", reduce_axis=-1),
EngineOp("MATH", "fma", src_a="l", src_b="corr", src_c="l_loc", dst="l"),
EngineOp("MATH", "mul_bcast", src_a="O", src_b="corr", dst="O", bcast_axis=1),
EngineOp("MATH", "copy", src="m_new", dst="m"),
),
)
TLContext, on tl.composite(prologue=[{"op": "softmax_merge", ...}], op="gemm", ...):
- Validates operand R/RW against
RECIPE_DESCRIPTORS["softmax_merge"]. - Allocates scratch (m_loc, m_new, corr, P, l_loc) via the scratch_scope helper (ADR-0063).
- Derives the primary output P's shape from
s.shape(identity). - Expands
engine_seqinto 8 flat MATH OpSpecs with all addresses / sizes filled. Each OpSpec getsscope=KERNEL. - Auto-binds head GEMM's
operands["a"] = P_handle. - Emits
CompositeCmd(ops=(8 MATH + 1 GEMM + epilogue), rw_handles= (m, l, O)).
RECIPE_DESCRIPTORS appears nowhere in the HW path. PE_SCHEDULER sees only the flat ops list.
D6. Invariants
- GEMM count.
0 ≤ count(op.kind == "gemm" in cmd.ops) ≤ 1. A composite with 0 GEMMs is MATH-only. - softmax_merge has no V operand.
RECIPE_DESCRIPTORS["softmax_merge" ].operandsdoes not include any HBM-resident handle → no DMA inserted during the prologue → V DMA can only happen during the GEMM head → natural K-before-V DMA priority across #1 / #2 of decode opt2. - Strict FIFO across composites. PE_SCHEDULER tracks in-flight
rw_handles. A new composite whoserw_handles(or any of its op inputs) intersect with an in-flight composite'srw_handleswaits for all earlier composites to complete — no out-of-order reorder. - No backwards-incompatibility for legacy callers. The user-facing
API
tl.composite(op="gemm", a, b, epilogue=[...])is preserved. TLContext internally lowers it to a flat-opsCompositeCmd. - PE_SCHEDULER recipe-free. PE_SCHEDULER does not import RECIPE_DESCRIPTORS, does not know that "softmax_merge" exists, and has no recipe-specific code branches.
D7. Boundary summary (compiler vs scheduler vs engine)
HOST DEVICE
TLContext (compiler analog) PE_SCHEDULER (scheduler/dispatcher)
- validate against RECIPE_DESC. - scan flat ops list
- derive primary_out shape - identify GEMM (≤1)
- allocate scratch - position-based stage placement
- expand recipe → flat OpSpecs - auto-insert DMA stages from space
- compute rw_handles - strict-FIFO RW hazard tracking
- emit CompositeCmd - emit Stage list
imports: RECIPE_DESCRIPTORS imports: nothing recipe-related
ENGINES (PE_MATH / PE_GEMM / PE_DMA)
- read Stage.params (op_kind, addresses, n_elements)
- existing latency models unchanged
imports: nothing recipe-related
Alternatives
A1. Three composites per tile (no prologue concept)
Split #2 further into softmax_merge + gemm(P·V) + add(O) as three
separate CompositeCmds. Pros: no flat-ops restructuring (add reuses
existing epilogue path). Cons: 3 dispatches per tile instead of 2 —
diverges from ADR-0060 §8 item 4 sizing ("split into two"); loses
~32 ns × N_tiles of fixed-cost savings. Rejected.
A2. Mixed-engine recipe (flash_pv_merge)
One recipe absorbing MATH + GEMM + MATH (ADR-0060 §5.6 original phrasing). Pros: tightest dispatch (1 composite). Cons: PE_SCHEDULER must know engine-boundary recipes, violates the "engines see only Stage.op_kind" boundary, recipe table becomes ISA-like and pollutes HW interface. Rejected.
A3. Generic op-list DAG (tl.ex_composite([...]))
User-defined sequence of any ops with named slots and dataflow analysis. Pros: maximally flexible. Cons: a single user (softmax_merge) does not justify a generic DAG language; PE_SCHEDULER would need dataflow analysis; reintroduces ADR-0060 §8 item 4's "absorb the inner loop" trap. Rejected (Simplicity First).
A4. RW-aware scheduler reorder (not strict FIFO)
PE_SCHEDULER could allow a later CompositeCmd to overtake an in-flight
one when their rw_handles don't intersect (e.g., next-tile #1 could
overlap with this-tile #2 since #1 doesn't touch m/l/O). Pros: more
overlap. Cons: more scheduler state, gains are bounded (next #1 waits
for K DMA anyway). Deferred — strict FIFO ships first; RW-aware
reorder may be a future follow-up ADR.
A5. Embedding softmax_merge as an epilogue of #1 (Q·Kᵀ)
Wrong: softmax_merge's output P is the input of #2's P·V GEMM, which
runs after #1's epilogue. Embedding it would force #2 to read its own
input from #1's epilogue scratch — circular tile-loop dependency.
Rejected (incorrect).
Consequences
Positive
- ADR-0060 §5.6 / §8 item 4 carve-out implemented exactly: opt2 = 2 composites per tile, K-before-V DMA priority natural.
- Composite contract becomes uniform: flat ops list with scope-based placement. No structural distinction between "head" and "epilogue" in the HW cmd.
- DMAs auto-inferred from operand
space; kernel surface simpler. - HW interface (
CompositeCmdshape) stays small even as recipes grow — recipe expansion happens host-side. - New fused patterns (linear+gelu+norm, rmsnorm+linear, …) can be added
by inserting one
RECIPE_DESCRIPTORSentry; no PE_SCHEDULER changes.
Negative
- CompositeCmd's struct shape changes — a refactor for all current callers. Semantics preserved (D6.4); existing goldens unchanged once ADR-0064 Revision 2 lands.
OpSpec.operandschanges from positionaltuple[Any, ...]todict[str, TensorHandle]— touches the existing epilogue lowering.- New TLContext code: recipe expansion + scratch-slot allocation + primary_out auto-bind. ~80–120 LOC.
- PE_SCHEDULER's
_generate_plangets a position-based scan + DMA auto-insertion + strict-FIFO RW tracker. ~60–100 LOC.
Open review items
- Recipe shape derivation rule beyond identity. Future recipes may
need non-identity transforms (e.g., transposed primary_out). The
PrimaryOutSpec.transformfield is forward-compat; add new transforms as needed. - Recipe scratch allocator integration with ADR-0063's
scratch_scope. Initial wiring: TLContext allocates inside the currentscratch_scopecontext if active, else in a kernel-scope persistent slot. tl_recipes.pylocation. New module undertriton_emu/. Keeps recipe knowledge co-located with the compiler analog (TLContext).- GEMM
out=TensorHandle. New form alongside legacyout_ptr: int. Both accepted; handle form preferred for new code (enables strict-FIFO RW tracking by handle identity).
Test Requirements
- CompositeCmd flat-ops refactor — meaning preservation. Every
existing bench's op_log is byte-equal pre/post refactor (legacy API
tl.composite(op="gemm", a, b, epilogue=[...])lowers to the same Stage sequence). - softmax_merge recipe lowering. TLContext call
tl.composite( prologue=[{"op": "softmax_merge", "s": Sj, "m": m, "l": l, "O": O}], op="gemm", b=V_ref, out=O, epilogue=[{"op": "add", "other": O}])produces aCompositeCmdwith exactly 10 ops in the expected order (8 MATH + 1 GEMM + 1 MATH(add)) andrw_handles == (m, l, O). - K-before-V DMA priority invariant. op_log of decode opt2 shows no V-related DMA during #2's MATH prologue (#2's V DMA begins only when the GEMM head starts).
- Strict-FIFO RW serialization. Two consecutive composites sharing
any
rw_handlecomplete in dispatch order; their stages do not interleave. - DMA auto-insertion from
space. A GEMM composite withb=tl.ref(...)(space=hbm) emits a DMA_READ Stage; withbfromtl.zeros(...)(space=tcm) it does not. Output handle withspace=hbmemits DMA_WRITE;space=tcmdoes not. - opt2 numeric parity with opt3. In data mode, opt2's final
(m, l, O)matches opt3's within fp tolerance. - opt2 dispatch ratio (after ADR-0064 Rev2). opt3 vs opt2 per-tile PE_CPU dispatch cycles ratio ≈ 2.4× at default calibration.
- GEMM-count invariant. A composite with two GEMM OpSpecs raises a validation error at TLContext emit.
Dependencies
- ADR-0060 §5.6, §8 item 4 — the carve-out this ADR implements.
- ADR-0064 Revision 2 — dispatch cost model; lands first, makes opt2's fewer-issues win measurable.
- ADR-0063
scratch_scope— recipe intermediates allocated inside;m, l, Oallocated outside as persistent arena (per DDD-0060 §6.2). - ADR-0046
tl_contextcontract — extended byprologue=[...]kwarg,out=TensorHandle, RECIPE_DESCRIPTORS lookup. - ADR-0042 tile plan generators — extended (not redesigned) to process the flat ops list and auto-insert DMAs.
- ADR-0014 PE pipeline — boundary preserved: scheduler stays recipe-free, engines stay op_kind-opaque.
Migration
Land order:
- ADR-0064 Revision 2 (separate PR): structural dispatch cost +
logical_bytesproperty + topology config override + one-time goldens regeneration. - ADR-0065 Phase 1 (this ADR, PR a):
CompositeCmdflat-ops restructure + legacy-API lowering. Refactor only; goldens unchanged. - ADR-0065 Phase 2 (this ADR, PR b):
tl_recipes.py+softmax_mergerecipe + PE_SCHEDULER position-scan + DMA auto-insertion + strict-FIFO RW tracker. - ADR-0065 Phase 3 (this ADR, PR c):
_gqa_decode_long.pyopt2 variant + numeric parity test + dispatch-ratio measurement against opt3 baseline.
Detailed implementation plan: see DDD-0065.