Land the new GQA fused-attention kernels (ADR-0060) for prefill/decode
across long and short context, the TL discipline primitives they depend
on (ADR-0062 lazy load, ADR-0063 scratch_scope + copy_to), and the
per-op-type CPU issue cost model (ADR-0064). Remove the pre-ADR-0060
mesh-attention baseline now that the unified kernels supersede it.
ADR-0060 (long context)
- _gqa_decode.py: M-fold + 2-level chain reduce-to-root (Level-2
intra-CUBE row-then-col + Level-1 inter-CUBE) — root-only output.
- _gqa_prefill.py: head-parallel + Ring KV rotation around C CUBEs,
online-softmax merge per ring step, per-CUBE distributed output.
- Each merge stage wraps in scratch_scope() and persists running
(m, l, O) via copy_to() to lift the 1 MiB scratch ceiling.
ADR-0060 §B.split.2 (short context, kv_per_cube in {1,2,4,8})
- _gqa_decode_short.py / _gqa_prefill_short.py: no cube-SP; each CUBE
owns whole KV heads; PE-parallel heads with intra-group chain
reduce. Prefill has no Ring KV (each head fully resident).
ADR-0062 (lazy tl.load): future-bearing TensorHandle, auto-wait at
first consuming op (dot/MATH/store/send/copy_to/composite).
ADR-0063 (tl.scratch_scope + tl.copy_to): scoped per-tile arena with
copy_to writeback primitive for persistent running state.
ADR-0064 (CPU issue cost model)
- common/cpu_issue_cost.py: per-op-type table (composite=40 ns,
primitives=5 ns); ratios are load-bearing per D1.
- TLContext: issue_cost_table param; _emit_dispatch_overhead(kind)
consults table with dispatch_cycles fallback (ADR-0046 §D6
back-compat).
- Live PE_CPU paths (greenlet + legacy) construct TLContext with
DEFAULT_CPU_ISSUE_COST so saturation lever (ADR-0060 §1) is
measurable end-to-end.
P7 headline bench: milestone-gqa-headline writes per-panel
op_log_summary to 1H_milestone_output/gqa_headline/sweep.json. No
figure renderers yet (deferred).
Removals (pre-ADR-0060 baseline now superseded):
- benches: _attention_mesh_kv.py, _attention_mesh_mlo.py,
_attention_mesh_mlo_2d.py, milestone_gqa_llama70b.py
- tests: test_attention_*, test_mesh_*, test_milestone_gqa_llama70b
- topology: llama70b_4sip.yaml (only consumer was the deleted diag)
- artifacts: 1H_milestone_output/gqa/ (sweep.json + 5 PNGs)
- tests/gqa/ plot helper + test (broken on Windows Tcl/Tkinter)
- ADR-0060/0061 references to deleted file paths cleaned up
(EN + KO kept in sync).
Tests: 124/124 focused regression green (attention + Phase E + TL
discipline + triton_emu + pe_components). Full regression: 764 pass,
2 pre-existing test_bench_registry failures (stale EXPECTED_NAMES
across multiple benches, not introduced here).
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
9.3 KiB
ADR-0063: tl.scratch_scope — per-tile scratch recycling for long context
Status
Proposed
Supporting ADR for ADR-0060 (AHBM GQA Fused Attention). Long-context attention sweeps many K/V tiles; each tile's intermediates (
scores,P,exp, partialO, …) currently allocate fresh scratch that is never freed within a kernel invocation. The 1 MiB per-PE scratch budget is exhausted long before a realistic context length.
Context
The bump allocator
TLContext allocates every math/compute output handle from a per-PE
scratch pool with a linear bump cursor
(src/kernbench/triton_emu/tl_context.py; _scratch_alloc):
def _scratch_alloc(self, nbytes):
aligned = (nbytes + 15) & ~15
addr = self._scratch_base + self._scratch_cursor
self._scratch_cursor += aligned
if self._scratch_cursor > self._scratch_size: # default 1 << 20 = 1 MiB
raise RuntimeError("TLContext scratch overflow: ...")
return addr
The docstring states the cursor "resets on every kernel invocation" —
i.e. only at kernel entry, never within the kernel body. Every
tl.dot, tl.softmax, tl.exp, tl.sum, a - b, a * b, … grabs a
fresh slice and nothing is reclaimed until the kernel returns.
Why this bites the GQA kernel
The current milestone bench keeps S_q = S_kv_per_rank = 16 explicitly
because of this limit
(tests/attention/test_milestone_gqa_llama70b.py:123-148):
"S_q_prefill and S_kv_per_rank are deliberately small (16 each) so the simulator's 1 MB per-PE TCM kernel scratch is not exhausted by the bump-allocated handle outputs of softmax/exp/dot/sum chains over n_ranks ring steps."
A FlashAttention sweep over n_tiles tiles allocates O(n_tiles ×
per-tile-intermediates). For any realistic context this overflows. The
math of flash attention needs only O(1) live scratch — the running
(m, l, O) plus the current tile's working set — because each tile's
temporaries are dead once that tile is folded into the running state. The
allocator just doesn't know they're dead.
Decision
Add a scratch scope that lets a kernel mark a region of allocations as reclaimable, so per-tile temporaries are recycled while live running state (and in-flight prefetch buffers) are preserved.
D1. tl surface — context manager
with tl.scratch_scope():
s = tl.dot(q, k_t) # all handles allocated inside the `with`
p = tl.softmax(s) # share a region that is rewound on exit
o_j = tl.dot(p, v)
# ... fold o_j into running (m,l,O) which live OUTSIDE the scope ...
# on __exit__: cursor rewinds to its value at __enter__
Semantics:
__enter__records the current_scratch_cursoras a save-point.__exit__restores the cursor to the save-point, freeing everything allocated inside.- Handles allocated outside the scope (running
m,l,O, prefetch buffers held byLoadFuturefrom ADR-0062) keep their addresses — they were allocated before the save-point or in an enclosing scope.
D2. Safety contract
A handle allocated inside a scope must not be read after the scope exits — its bytes may be overwritten by the next scope's allocations. The flash loop respects this naturally: the only values that survive a tile iteration are the running accumulators, which are allocated outside the per-tile scope and updated by ops inside it writing to outside addresses (the merge writes new running state — see D3).
D3. Interaction with running accumulators
The online-softmax merge reads the old running (m, l, O) and the
current tile's (m_j, l_j, O_j), producing new running values. To keep
the new running values outside the recycled region, the merge writes them
to stable scratch allocated once before the loop (a small fixed
"running-state" arena, distinct from the per-tile scope). Concretely the
kernel keeps two arenas:
- persistent arena (allocated once):
m, l, O(and their double-buffer if needed for the merge). - scoped arena (rewound each tile):
scores, P, exp, O_j, scale_*.
This mirrors how real flash-attention SRAM budgeting works: a small persistent accumulator region + a recycled tile working set.
D3.1 The persistent-arena write mechanism: tl.copy_to(dst, src)
The merge ops (tl.maximum, tl.exp, binary * / +) all call
_make_compute_out(...) which allocates from the bump cursor (D1).
Inside a scratch_scope, their result handles therefore live inside
the scope and vanish on __exit__. To realise D3's two-arena split, the
kernel needs a way to write a scoped result's bytes back to a
persistent address.
The primitive that closes this gap:
def copy_to(self, dst: TensorHandle, src: TensorHandle) -> None:
"""Copy ``src``'s bytes into ``dst``'s address (both TCM).
Shapes and dtypes must match. ``dst`` is typically a handle
allocated outside any active ``scratch_scope`` (the persistent
arena); ``src`` is a scoped handle whose bytes must outlive
scope ``__exit__``.
"""
Symmetric to tl.store (the HBM-side byte copy), kept TCM-only here so
the running-state writeback doesn't pollute op_log with spurious DMA.
Mechanics:
- New
CopyCmd(src, dst, nbytes, data_op=True)command. - op_log:
op_kind="math",op_name="copy"— runs on the vector engine. - Latency:
pe_math._compute_ns(prod(shape))— models on-chip register writeback, not HBM transfer. - Emit-time validation:
dst.shape == src.shape,dst.dtype == src.dtype,dst.space == "tcm",src.space == "tcm". Authoring errors surface in Phase 1, not deep in Phase 2 data execution.
Call-site pattern:
m, l, O = init_running(...) # persistent (outside scope)
for j in range(n_tiles):
with tl.scratch_scope():
... # per-tile work (recycled)
m_new = tl.maximum(m, mj) # scoped scratch
l_new = l * scale_old + l_step * scale_step
O_new = O * scale_old + O_step * scale_step
tl.copy_to(m, m_new) # ← persist new running state
tl.copy_to(l, l_new)
tl.copy_to(O, O_new)
# exit: scoped m_new/l_new/O_new gone; their bytes live in persistent m/l/O
The copy happens before __exit__, so the read of src (scoped) is
valid; after exit only dst (persistent) is read, satisfying D2's
no-read-after-exit safety contract.
Why a dedicated primitive rather than dst= kwargs on every math op
(considered, rejected): adding dst= to tl.maximum, tl.exp,
_binary_math, _unary_math, _reduction is ~25 LOC across 5
op-families and breaks the uniform "call returns a fresh handle"
pattern. tl.copy_to is one primitive, one command, one executor
handler — minimal surface area for the same effect.
D4. Nesting
Scopes nest (stack of save-points). Inner scope exit rewinds to the inner save-point; outer exit rewinds further. This supports an outer "per-query-block" scope around an inner "per-KV-tile" scope for prefill.
Alternatives
A1. Round-trip temporaries through HBM
Store intermediates to HBM and reload to "free" TCM scratch. Rejected: turns a TCM-resident streaming kernel into an HBM-bandwidth-bound one — the opposite of the goal, and it pollutes op_log with spurious DMA.
A2. Tiled tl.composite (scheduler-managed scratch)
tl.composite recycles per-tile scratch inside PE_SCHEDULER
automatically. As in ADR-0062 A1, this is attractive but blocked on a
flash-capable composite kind (two GEMMs + carried (m,l,O)), a much
larger change. scratch_scope gives the same memory behaviour for the
greenlet kernel with a tiny, general primitive. The two can coexist.
A3. Grow the scratch budget
Bump pe_tcm.kernel_scratch_mb. Rejected as a fix: it only pushes the
context-length ceiling out linearly while still leaking O(n_tiles)
scratch, and it misrepresents the hardware (real TCM is small; the point
of flash attention is O(1) working set). Useful only as a coarse knob,
not a substitute for recycling.
Consequences
Positive
- Removes the artificial
S = 16validation-scale ceiling; enables realistic context lengths in both timing and data modes. - Faithfully models flash attention's O(1) working-set property.
- Small, general primitive (any tiled kernel benefits).
Negative
- A use-after-scope bug silently reads stale bytes. Mitigated by the D2 contract, by keeping the scope discipline inside a shared attention helper, and (optionally) by a debug build that poisons rewound regions.
- Must coordinate with ADR-0062: prefetch buffers are live across tile iterations, so they belong to the persistent arena, not the scoped one.
Test Requirements
- Recycling: a loop of
Ntiles insidetl.scratch_scope()keeps peak_scratch_cursorbounded by one tile's footprint, independent ofN(today it grows linearly and overflows). - Correctness: a flash sweep with scopes produces the same
Oas the same sweep without scopes at a smallNthat fits without recycling (Phase 2). - Long context: a sweep at an
Nthat would overflow 1 MiB without scopes completes (the exact failure theS=16cap avoids today). - Persistent-vs-scoped isolation: running
(m,l,O)allocated outside the scope retains correct values across__exit__. - Nesting: nested scopes rewind to the correct save-points.