Files
kernbench2/docs/adr-proposed/ADR-0063-prog-tl-scratch-scope.md
T
ywkang f4f55b2c1c gqa(adr): add supporting feature ADRs 0061-0063 for GQA fused attention
Proposed prerequisites surfaced while evaluating the GQA fused-attention
ADR against the actual kernbench tl/sim_engine implementation:

- ADR-0061 tl.broadcast: data-faithful GQA head reuse (fixes the
  MemoryStore nbytes check that forces h_q==h_kv==1 today).
- ADR-0062 tl.load_async: non-blocking HBM tile load for KV prefetch
  (KV-load-bound decode/long-context overlap).
- ADR-0063 tl.scratch_scope: per-tile scratch recycling (removes the
  1 MiB bump-allocator ceiling that caps context at S=16).

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
2026-06-03 17:38:55 -07:00

6.7 KiB
Raw Blame History

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, partial O, …) 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_cursor as 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 by LoadFuture from 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.

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 = 16 validation-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

  1. Recycling: a loop of N tiles inside tl.scratch_scope() keeps peak _scratch_cursor bounded by one tile's footprint, independent of N (today it grows linearly and overflows).
  2. Correctness: a flash sweep with scopes produces the same O as the same sweep without scopes at a small N that fits without recycling (Phase 2).
  3. Long context: a sweep at an N that would overflow 1 MiB without scopes completes (the exact failure the S=16 cap avoids today).
  4. Persistent-vs-scoped isolation: running (m,l,O) allocated outside the scope retains correct values across __exit__.
  5. Nesting: nested scopes rewind to the correct save-points.