Files
kernbench2/docs/adr/ADR-0063-prog-tl-scratch-scope.md
T
mukesh 7fad0371c5 gqa: tile-granular Ring KV (P3c) + rename to gqa_attention_* + ADR-0060/62/63/64 → Accepted
Three logically distinct changes, bundled for atomic test green:

1. **P3c — prefill_long tile-granular Ring KV** (ADR-0060 §5.5.1 amendment).
   Convert the ring from slice-granular (one full ``(d_head, S_local)``
   KV slice per step) to tile-granular (``n_tiles`` tiles of
   ``TILE_S_KV`` per step). Nested loop with outer tile, inner ring step:
   each tile propagates through all C ring positions before the next
   tile starts, so IPCQ in-flight depth stays at 1 per direction.
   Bootstrap at ``(t=0, k=0)`` outside the scratch_scope establishes the
   persistent ``(m, ℓ, O)``; every other iteration scope-wraps + persists
   via ``copy_to``. Per-rank persistent scratch shrinks to ~1 KB; per-tile
   scope bounded by TILE_S_KV regardless of S_local. Headline:
   prefill_long now completes at S_kv=128K (previously overflowed).
   New: ``tests/attention/test_gqa_prefill_long_tile_ring.py``
   (3 tests — ceiling-lift + tile-granular ipcq_copy count +
   per-CUBE distributed output regression guard).

2. **Rename ``gqa_*`` → ``gqa_attention_*``** across kernel files,
   function names, and importers. The "attention" name makes the role
   explicit (GQA is grouped-query attention) and matches upstream Triton
   FlashAttention naming conventions. Renames:
     _gqa_decode_long.py        -> _gqa_attention_decode_long.py
     _gqa_decode_short.py       -> _gqa_attention_decode_short.py
     _gqa_prefill_long.py       -> _gqa_attention_prefill_long.py
     _gqa_prefill_short.py      -> _gqa_attention_prefill_short.py
   And function names ``gqa_<phase>_<context>_kernel`` →
   ``gqa_attention_<phase>_<context>_kernel``. Updated 1 bench file
   (milestone_gqa_headline.py) and 10 test files.

3. **ADR-0060 / 0062 / 0063 / 0064: Proposed → Accepted**.
   All four are reflected in production code and covered by tests:
   - ADR-0060 (GQA fused attention): 4 kernels deployed; §5.5.1
     amendment added for the tile-granular Ring KV introduced by P3c
     (EN + KO mirror).
   - ADR-0062 (lazy tl.load): LoadFuture + _await_pending live in
     tl_context.py.
   - ADR-0063 (tl.scratch_scope + tl.copy_to): used in every chain
     reduce + tile sweep + ring step. EN-only previously; KO
     translation authored as part of this commit (CLAUDE.md
     bidirectional rule).
   - ADR-0064 (per-op-type CPU issue cost): cpu_issue_cost.py +
     issue_cost_table wiring in tl_context.py (Phase E).
   Files git mv'd from docs/adr-proposed/ to docs/adr/ (EN) and
   docs/adr-ko/ (KO). ADR-0061 (tl.broadcast) stays Proposed — no
   implementation; documented as optional convenience primitive in
   the ADR itself.

Tests: 88/88 focused regression green
(tests/attention/ + Phase E + TL discipline).
ADR pair verification: ``python tools/verify_adr_lang_pairs.py`` OK.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-06-10 16:20:00 -07:00

9.3 KiB
Raw Blame History

ADR-0063: tl.scratch_scope — per-tile scratch recycling for long context

Status

Accepted

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.

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 = 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.