gqa: ADR-0060/0062/0063/0064 unified GQA kernels + CPU cost model
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>
This commit is contained in:
@@ -102,6 +102,68 @@ kernel keeps two arenas:
|
||||
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
|
||||
|
||||
Reference in New Issue
Block a user