f4f55b2c1c
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>
163 lines
6.7 KiB
Markdown
163 lines
6.7 KiB
Markdown
# 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`):
|
||
|
||
```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.
|
||
|
||
### 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.
|