# 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`): ```python 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 ```python 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: ```python 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:** ```python 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.