gqa(adr/ddd): iteration-2 corrections — GQA reuse via M-fold, transpose caveat

Verified against sim_engine data path (memory_store, data_executor):

- GQA reuse does NOT need a broadcast op. The baseline's h_q==h_kv limit
  is its head-packing reshape hack, not a missing primitive. Correct fix:
  per-KV-head loop with G folded into matmul M dim (byte-conserving
  reshape) — runs today, timing correct (m=G*T_q), data mode runs.
- ADR-0061 broadcast demoted from 'the blocker' to optional convenience.
- Surfaced tl.trans = reshape-not-transpose (memory_store reshapes;
  data_executor np.matmul on reshaped operands) -> numeric parity is
  bounded; verification is structural/timing/determinism-first (matches
  SPEC perf-model purpose). Optional tl.transpose deferred.
- Reordered DDD phase plan (P1 GQA needs no new feature; P3 scratch_scope
  is the key scale feature); added open decisions 10.10 (transpose) and
  10.11 (GQA-via-M-fold finding).

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
2026-06-03 17:51:03 -07:00
parent fb600b3cdd
commit 6d24b9306f
3 changed files with 177 additions and 69 deletions
@@ -14,10 +14,12 @@ KV-load-bound decode; sequence-parallel (Ring KV) for long-context prefill.
kernbench work* — that code is the baseline this ADR upgrades to a real kernbench work* — that code is the baseline this ADR upgrades to a real
GQA, causal, long-context kernel. GQA, causal, long-context kernel.
**Supporting ADRs (prerequisites):** **Supporting ADRs** (efficiency / scale enablers — *not* GQA blockers;
- **ADR-0061** `tl.broadcast` — data-faithful GQA head reuse. see §8 correction): **ADR-0063** `tl.scratch_scope` (per-tile scratch
- **ADR-0062** `tl.load_async` — non-blocking HBM tile load (KV prefetch). recycling — required for realistic context length), **ADR-0062**
- **ADR-0063** `tl.scratch_scope` — per-tile scratch recycling. `tl.load_async` (KV prefetch overlap — efficiency), **ADR-0061**
`tl.broadcast` (optional mask/general convenience). Real GQA itself needs
only kernel restructuring (§5.2).
**Algorithm lineage.** This kernel is **FlashAttention** (tiling + **Algorithm lineage.** This kernel is **FlashAttention** (tiling +
online/streaming softmax with fused P·V — no full score matrix online/streaming softmax with fused P·V — no full score matrix
@@ -55,9 +57,15 @@ loop. This **matters for this ADR's design** (see §1).
*efficient GQA* kernel must lift: *efficient GQA* kernel must lift:
1. **No GQA reuse.** `h_q == h_kv == 1` 1. **No GQA reuse.** `h_q == h_kv == 1`
(`test_milestone_gqa_llama70b.py:137-142`). Real GQA (`h_q = G·h_kv`) (`test_milestone_gqa_llama70b.py:137-142`). The test attributes this to
is blocked by the MemoryStore byte-conservation check on the symbolic a MemoryStore byte-conservation failure on a *broadcast view*, but that
broadcast view → fixed by **ADR-0061**. failure is a property of the baseline's **head-packing hack**
(`_view(K, (h_q·d, S_kv))`, which conflates all heads into one matmul
dim and only conserves bytes when `h_q == h_kv`). The correct fix is
**kernel restructuring**, not a broadcast op: process **one KV head at
a time** and fold the `G` group rows into the matmul **M** dimension
(§5.2). That uses only byte-conserving reshapes, so real GQA
(`h_q = G·h_kv`) runs with **no new primitive** — see §8.
2. **O(N) reduction.** The baseline does an all-to-all bidirectional 2. **O(N) reduction.** The baseline does an all-to-all bidirectional
fan-out so *every* rank ends with the full answer (`n_ranks 1` fan-out so *every* rank ends with the full answer (`n_ranks 1`
steps). Attention only needs `O` at the query owner → a **tree steps). Attention only needs `O` at the query owner → a **tree
@@ -302,10 +310,16 @@ with tl.scratch_scope(): # per-tile temporaries rec
``` ```
Notes: Notes:
- `q_g` is the GQA-batched query: `[G, d]` (decode) or `[G, T_q, d]` - `q_g` is the GQA-batched query reshaped to `[G·T_q, d]` (the `G` group
(prefill). One K/V tile load serves all `G` rows — the GQA reuse lever. rows folded into the matmul M dim; byte-conserving). One K/V tile load
Where the matmul needs the KV head expanded to `G`, use serves all `G·T_q` rows — the GQA reuse lever — with no broadcast.
`tl.broadcast` (ADR-0061), never the symbolic `_view`. - `tl.trans(Kj)` is **metadata-only** in kernbench (`tl_context.py:390`),
and `MemoryStore.read` *reshapes* rather than transposes
(`memory_store.py:73`). For zero/structural runs this is harmless; for
non-trivial numeric data it yields a reshape-not-transpose. Data-mode
*numeric* parity therefore needs care (§11) — store K pre-transposed,
or add a real `tl.transpose` (a candidate further primitive, likely
unnecessary given the simulator's performance-modeling purpose).
- The V load is issued (prefetched) *before* it is needed so its DMA - The V load is issued (prefetched) *before* it is needed so its DMA
overlaps the Q·Kᵀ + softmax of the same/earlier tile. overlaps the Q·Kᵀ + softmax of the same/earlier tile.
- Masking: the kernel builds the boundary-tile mask from query/KV global - Masking: the kernel builds the boundary-tile mask from query/KV global
@@ -398,9 +412,15 @@ def attn_kernel(q_ptr, k_ptr, v_ptr, o_ptr, counter, start_pe, N,
- `T_q = 1`, attends to all past KV ⇒ no future tiles, mask only on the - `T_q = 1`, attends to all past KV ⇒ no future tiles, mask only on the
final ragged tile. final ragged tile.
- **GQA reuse is the whole game** (decode is KV-load-bound): the `G=8` - **GQA reuse is the whole game** (decode is KV-load-bound): the `G=8`
query rows of the KV head are batched (`q_g = [G, d]`), so each K/V query rows of the KV head are folded into the matmul **M** dimension
tile is loaded once and reused by 8 rows (Q·Kᵀ is an `8×TILE` batched (`q_g` reshaped `[G, T_q, d] → [G·T_q, d]`, a byte-conserving reshape).
GEMV, P·V an `8×d`). `tl.broadcast` expands the KV head where needed. Then `Q·Kᵀ` is `tl.dot([G·T_q, d], Kᵀ[d, TILE]) → [G·T_q, TILE]` and
`P·V` is `tl.dot([G·T_q, TILE], V[TILE, d]) → [G·T_q, d]`. The KV tile
(`[TILE, d]`) is the shared `K`/`V` operand — **loaded once, reused by
all `G·T_q` rows automatically** because they are the M rows of the
GEMM. No broadcast of K/V is needed; `m = G·T_q` in the emitted
`GemmCmd` also makes the Phase-1 timing count all `G` rows' work
correctly (a leading batch axis would *not* be counted — see §8).
- If `S_pe` fits in scratch (small/medium context) this degenerates to a - If `S_pe` fits in scratch (small/medium context) this degenerates to a
**one-shot** partial attention (one `tl.dot` for Q·Kᵀ, one softmax, one **one-shot** partial attention (one `tl.dot` for Q·Kᵀ, one softmax, one
`tl.dot` for P·V) — exactly the baseline `_attention_mesh_mlo` `tl.dot` for P·V) — exactly the baseline `_attention_mesh_mlo`
@@ -474,20 +494,16 @@ abstraction.
## 8. Required kernbench changes ## 8. Required kernbench changes
**Genuinely new (each has a supporting ADR):** **Correction from design iteration:** real GQA (`h_q > h_kv`) needs **no
new primitive** — only the kernel restructuring in §5.2 (per KV head,
1. **GQA head broadcast in data mode** — **ADR-0061** (`tl.broadcast`). `G` folded into M, byte-conserving reshapes). The supporting ADRs are
*The* blocker for `h_q > h_kv`. Without it the GQA reuse lever cannot *efficiency / scale* enablers, not GQA blockers.
run under `enable_data=True`.
2. **Async HBM tile load / KV prefetch** — **ADR-0062** (`tl.load_async`
+ `tl.wait`). The KV-load-bound overlap lever for decode/long-context.
3. **Per-tile scratch recycling** — **ADR-0063** (`tl.scratch_scope`).
Removes the `S=16` ceiling so realistic context lengths run.
**Algorithm work in the kernel (no new primitive; existing `tl` API):** **Algorithm work in the kernel (no new primitive; existing `tl` API):**
- GQA Q-axis batching (`G` rows share each K/V load) — `tl.broadcast` + - **GQA Q-axis batching** (the reuse lever) — fold `G·T_q` into the matmul
2-D `tl.dot`. M dim per KV head (§5.2); `_view`-style byte-conserving reshape + 2-D
`tl.dot`. Runs today in both timing and data mode.
- Tree reduction to root (§4) replacing the baseline all-to-all fan-out — - Tree reduction to root (§4) replacing the baseline all-to-all fan-out —
pure kernel control flow over `tl.send`/`tl.recv`. pure kernel control flow over `tl.send`/`tl.recv`.
- Causal tile skip + additive boundary mask (§3/§5.4) — kernel `if` + - Causal tile skip + additive boundary mask (§3/§5.4) — kernel `if` +
@@ -495,6 +511,21 @@ abstraction.
- Round-robin KV placement / valid-length arithmetic (§2) — launch-arg - Round-robin KV placement / valid-length arithmetic (§2) — launch-arg
arithmetic + `DPPolicy(pe="row_wise")`. arithmetic + `DPPolicy(pe="row_wise")`.
**New primitives for efficiency / scale (each has a supporting ADR):**
1. **Per-tile scratch recycling** — **ADR-0063** (`tl.scratch_scope`).
*Required for scale*: removes the `S=16` ceiling (1 MiB bump
allocator) so realistic context lengths run. Highest-value of the
three.
2. **Async HBM tile load / KV prefetch** — **ADR-0062** (`tl.load_async`
+ `tl.wait`). *Efficiency*: the KV-load-bound overlap lever for
decode/long-context. Without it the kernel is correct but serial.
3. **GQA head / mask broadcast** — **ADR-0061** (`tl.broadcast`).
*Optional convenience*, not a GQA blocker (see correction above).
Useful for additive-mask construction across the `G·T_q` rows and for
general kernels; `np.matmul` already broadcasts in data mode, so it is
not needed for correctness. Lowest priority.
**Explicitly REJECTED (efficient alternative chosen):** **Explicitly REJECTED (efficient alternative chosen):**
4. ~~Composite that chains DMA→MM→VEC→DMA→MM→VEC with carried `(m,,O)` 4. ~~Composite that chains DMA→MM→VEC→DMA→MM→VEC with carried `(m,,O)`
@@ -557,10 +588,21 @@ projection **downstream**; score `S` and probs `P` are never materialised.
SPEC/ADR coverage: R5 (PE↔PE IPCQ, PE↔HBM), R2 (latency by traversal), SPEC/ADR coverage: R5 (PE↔PE IPCQ, PE↔HBM), R2 (latency by traversal),
ADR-0023/0025 (IPCQ), ADR-0046 (`tl` contract), ADR-0054 (eval bench). ADR-0023/0025 (IPCQ), ADR-0046 (`tl` contract), ADR-0054 (eval bench).
1. **Correctness vs reference (Phase 2, `enable_data=True`):** for a The simulator's contract is **latency by traversal + determinism +
small `(G, T_q, S, N)`, kernel `O` matches a numpy FlashAttention structural correctness** (SPEC §0, §0.1), not bit-exact numerics — Phase 2
reference within fp tolerance — for all four cases. This is the test data exists mainly to exercise the data path, `tl.trans` is
that *requires* ADR-0061 (GQA) to even run with `h_q>h_kv`. reshape-not-transpose, and `bf16` is modelled as `f16`. Verification is
therefore **structural/timing-first**, with numeric parity as a bounded
secondary check.
1. **Runs in data mode (`enable_data=True`):** the GQA kernel
(`h_q = G·h_kv`) completes without the byte-conservation error that the
baseline head-packing hits — for all four cases. (This needs the §5.2
restructuring, *not* a new primitive.)
**Numeric parity (secondary):** for symmetric/identity inputs where
reshape-as-transpose is exact, kernel `O` matches a numpy
FlashAttention reference within fp tolerance. Full asymmetric parity is
gated on a real `tl.transpose` (out of scope; flagged).
2. **GQA reuse:** with `h_q = G·h_kv`, the K/V `dma_read_count` is 2. **GQA reuse:** with `h_q = G·h_kv`, the K/V `dma_read_count` is
independent of `G` (one load per tile, reused across the group), while independent of `G` (one load per tile, reused across the group), while
GEMM work scales with `G`. Asserts the lever actually fires. GEMM work scales with `G`. Asserts the lever actually fires.
@@ -4,11 +4,21 @@
Proposed Proposed
> Supporting ADR for **ADR-0060** (AHBM GQA Fused Attention). This ADR > Supporting ADR for **ADR-0060** (AHBM GQA Fused Attention).
> removes the single hard blocker that keeps the existing attention >
> kernels at `h_q == h_kv == 1`: there is no way to expand a KV head to > **⚠ Re-evaluated during design iteration — this is now OPTIONAL, not a
> the `G` query heads that share it without violating the MemoryStore > GQA blocker.** The original framing (below) held that GQA `h_q > h_kv`
> byte-conservation check. > is blocked by the MemoryStore byte-conservation check on a broadcast
> view. That is true *only of the baseline's head-packing hack*
> (`_view(K, (h_q·d, S_kv))`). The correct GQA design (ADR-0060 §5.2)
> processes **one KV head at a time** and folds the `G` group rows into
> the matmul **M** dimension — using only byte-conserving reshapes — so it
> needs **no broadcast op at all**, and `np.matmul` already broadcasts a
> shared operand in Phase-2 data mode. This ADR therefore drops to a
> *convenience* primitive: it cleans up additive-mask construction across
> the `G·T_q` rows and serves general kernels. **Lowest priority of the
> three supporting ADRs.** The rest of this document is retained as the
> design for that convenience op.
## Context ## Context
@@ -140,24 +140,44 @@ ctx.launch(
`H_q > H_kv` is the change from the baseline's `H_q=H_kv=1`; it requires `H_q > H_kv` is the change from the baseline's `H_q=H_kv=1`; it requires
ADR-0061 to survive Phase 2. ADR-0061 to survive Phase 2.
### 4.3 GQA-batched Q load ### 4.3 GQA-batched Q load — fold `G` into the matmul M dim (no broadcast)
For a KV head `kv`, its group is query heads `[kv·G, (kv+1)·G)`. The For a KV head `kv`, its group is query heads `[kv·G, (kv+1)·G)`. Load
kernel loads those `G` rows and, where the matmul needs the KV head those `G` rows and **reshape the group into the matmul M dimension**
expanded to `G`, broadcasts the KV head — **not** the symbolic `_view`: this is the key to GQA reuse and it needs only a byte-conserving reshape,
**not** a broadcast:
```python ```python
def load_q_group(q_ptr, kv, G, T_q, d, tl): def load_q_group(q_ptr, kv, G, T_q, d, tl):
# Q laid out [H_q, T_q, d]; the group is contiguous in head axis. # Q laid out [H_q, T_q, d]; the group is contiguous in the head axis.
return tl.load(q_ptr + kv*G*T_q*d, shape=(G, T_q, d), dtype="f16") qg = tl.load(q_ptr + kv*G*T_q*d, shape=(G, T_q, d), dtype="f16")
return _view(qg, (G*T_q, d)) # byte-conserving reshape → M = G*T_q
def kv_for_group(k_tile, G, tl): # ADR-0061
# k_tile: [TILE, d] (one KV head). Expand to [G, TILE, d] for batched dot.
return tl.broadcast(k_tile[None, :, :], (G, k_tile.shape[0], k_tile.shape[1]))
``` ```
(Whether the batched `tl.dot` consumes `[G,T_q,d]·[G,d,TILE]` directly or Then per KV tile (`Kj`, `Vj` are `[TILE, d]`, one KV head, loaded once):
the kernel loops `g` is an Open Decision — §10.4.)
```python
Sj = tl.dot(q_g, tl.trans(Kj)) * scale # [G*T_q, d]·[d, TILE] -> [G*T_q, TILE]
Oj = tl.dot(P, Vj) # [G*T_q, TILE]·[TILE, d] -> [G*T_q, d]
```
K/V are the **shared operands** (the matmul's `b`), loaded once and reused
across all `G·T_q` M-rows automatically. The emitted `GemmCmd` carries
`m = G·T_q`, so Phase-1 timing counts all `G` rows' work correctly.
> **Why not a leading batch axis** (`[G,T_q,d]·[G,d,TILE]`)? `tl.dot`
> derives `m,k,n` from the **last two dims only** (`tl_context.py:220`)
> and `GemmCmd` carries no batch dim — so a leading `G` axis is computed
> correctly by `np.matmul` in Phase 2 (it broadcasts) but **under-counted
> by `G×` in Phase-1 timing**. Folding into M fixes both. This supersedes
> the earlier batched-vs-loop question.
> **Transpose caveat.** `tl.trans` is metadata-only and `MemoryStore.read`
> *reshapes* (not transposes) (`memory_store.py:73`,
> `data_executor.py:153-168` does `np.matmul` on the reshaped operands).
> Fine for timing/structure and for the baseline's zero data; for
> non-trivial numeric parity, store K pre-transposed as `[d, S_pe]` or add
> a real `tl.transpose` (§10.10).
--- ---
@@ -273,16 +293,17 @@ inner update — running `(m,l,O)` carried across steps as Python handles
## 6. The three new primitives — integration points ## 6. The three new primitives — integration points
### 6.1 ADR-0061 `tl.broadcast` ### 6.1 ADR-0061 `tl.broadcast` (OPTIONAL — not on the GQA critical path)
GQA reuse is achieved by folding `G` into M (§4.3), so broadcast is **not**
required to unblock `h_q > h_kv`. It remains a clean primitive for additive
causal-mask construction across the `G·T_q` rows and for general kernels.
- `tl_context.py`: new `broadcast(x, new_shape)` → emits `BroadcastCmd`, - `tl_context.py`: new `broadcast(x, new_shape)` → emits `BroadcastCmd`,
output handle from `_make_compute_out` (so `nbytes` is correct). output handle from `_make_compute_out` (so `nbytes` is correct).
- `data_executor.py`: `_execute_broadcast` = `np.broadcast_to(src, - `data_executor.py`: `_execute_broadcast` = `np.broadcast_to(src,
shape).copy()` → `store.write`. shape).copy()` → `store.write`.
- Latency: `pe_math._compute_ns(prod(new_shape))`; logs `math/broadcast`, - Latency: `pe_math._compute_ns(prod(new_shape))`; logs `math/broadcast`,
**no** `dma_read`. **no** `dma_read`.
- Unblocks: the `tl.dot(q_g[G,…], kv_for_group(...)[G,…])` consuming a - Lowest priority; deliver after the efficiency/scale features land.
broadcast KV head passes the MemoryStore nbytes check
(`memory_store.py:67-73`) that fails today.
### 6.2 ADR-0062 `tl.load_async` ### 6.2 ADR-0062 `tl.load_async`
- `tl_context.py`: `load_async(ptr, shape, dtype) → LoadFuture`; extend - `tl_context.py`: `load_async(ptr, shape, dtype) → LoadFuture`; extend
@@ -306,16 +327,18 @@ on a green predecessor; every phase keeps the existing baseline green.
| Phase | Deliverable | Gate | | Phase | Deliverable | Gate |
|---|---|---| |---|---|---|
| **P0** | ADR-0061 `tl.broadcast` + unit test | broadcast Phase 2 parity; downstream dot passes nbytes check | | **P1** | Real GQA by restructuring (per KV head, `G`→M fold); one-shot, no tiling. **No new primitive.** | data-mode run completes for `h_q=G·h_kv`; K/V `dma_read_count` independent of `G`; `gemm` m=`G·T_q` |
| **P1** | Real GQA in baseline decode kernel (`h_q=G·h_kv`), one-shot, no tiling, using broadcast | numpy parity at small `S`; K/V `dma_read_count` independent of `G` | | **P2** | Tree reduction replacing fan-out (decode SP) | `⌈log₂ N⌉` reduction rounds; structure preserved |
| **P2** | Tree reduction replacing fan-out (decode SP) | `⌈log₂ N⌉` reduction rounds; parity preserved | | **P3** | ADR-0063 `tl.scratch_scope` + tiled sweep | long-`S` decode completes (today caps at S=16) |
| **P3** | ADR-0063 `tl.scratch_scope` + tiled sweep | long-`S` decode completes; parity | | **P4** | ADR-0062 `tl.load_async` + prefetch in sweep | latency < serial Σ(load+compute) |
| **P4** | ADR-0062 `tl.load_async` + prefetch in sweep | latency < serial Σ(load+compute); parity | | **P5** | Causal masking + tile/step skip (prefill, ring) | GEMM count = triangular tile count |
| **P5** | Causal masking + tile/step skip (prefill, ring) | GEMM count = triangular tile count; parity |
| **P6** | Headline-scale milestone panels (real `G`, tiling, prefetch) + figures | sweep.json headline rows; figures render | | **P6** | Headline-scale milestone panels (real `G`, tiling, prefetch) + figures | sweep.json headline rows; figures render |
| **P7** *(opt)* | ADR-0061 `tl.broadcast` (mask convenience); optional `tl.transpose` for numeric parity | broadcast unit test; asymmetric numeric parity |
P0P1 alone deliver "GQA actually runs in data mode" — the single most **P1 alone delivers "GQA actually runs" — and it needs no new feature**,
important unblock. P2P5 deliver efficiency. P6 is the eval payoff. only the §4.3/§5 restructuring (the iteration-2 finding). P3 is the most
important *new* feature (removes the S=16 scale ceiling). P4 is the
efficiency lever. P2/P5 are algorithmic. P7 is optional polish.
--- ---
@@ -324,10 +347,14 @@ important unblock. P2P5 deliver efficiency. P6 is the eval payoff.
Grounded in SPEC R2/R5, ADR-0023/0025 (IPCQ), ADR-0046 (`tl`), Grounded in SPEC R2/R5, ADR-0023/0025 (IPCQ), ADR-0046 (`tl`),
ADR-0054 (eval bench). Mirrors ADR-0060 §11. ADR-0054 (eval bench). Mirrors ADR-0060 §11.
**Correctness (Phase 2, `enable_data=True`):** **Runs + numeric (Phase 2, `enable_data=True`):**
- `test_gqa_correctness.py`: for `(G∈{1,8}, T_q∈{1,16}, S, N∈{1,8})`, - `test_gqa_correctness.py`: for `(G∈{1,8}, T_q∈{1,16}, S, N∈{1,8})` the
kernel `O` ≈ numpy FlashAttention reference. The `G=8` cases are the GQA kernel **completes** (no byte-conservation error) — the `G=8` cases
ones that cannot run today (ADR-0061 gate). that the baseline head-packing cannot express. This is gated on the
§4.3 restructuring, **not** on a new primitive.
- *Numeric parity (secondary):* `O ≈ numpy FlashAttention` for
symmetric/identity inputs where reshape-as-transpose is exact; full
asymmetric parity is gated on an optional `tl.transpose` (§10.10).
**Levers (op_log assertions):** **Levers (op_log assertions):**
- GQA amortisation: K/V `dma_read_count` constant as `G` grows 1→8; - GQA amortisation: K/V `dma_read_count` constant as `G` grows 1→8;
@@ -397,12 +424,14 @@ on root only. Attention needs root-only, so tree wins. If any downstream
consumer needs O replicated (none found), we'd keep fan-out. *Assuming consumer needs O replicated (none found), we'd keep fan-out. *Assuming
root-only is fine.* root-only is fine.*
### 10.4 Batched 2-D dot vs per-`g` loop for GQA **[recommend: try batched first]** ### 10.4 GQA matmul shape — RESOLVED in iteration 2 **[fold G into M]**
`tl.dot` semantics for a leading group axis (`[G,T_q,d]·[G,d,TILE]`) need Earlier this asked "batched leading axis vs per-`g` loop." Resolved: do
confirming in `tl_context.dot` (it computes `m,k,n` from the last two **neither** — fold `G·T_q` into the matmul M dimension (§4.3). A leading
dims; `tl_context.py:213-229`). If batched dot over a leading axis isn't `G` axis would be computed by `np.matmul` in data mode but **under-counted
supported, fall back to a `for g in range(G)` loop (ADR-0061 A1) — correct `G×` in Phase-1 timing** (`GemmCmd` carries no batch dim,
but `G×` GEMM ops. *I assumed batched works or is a small extension.* `tl_context.py:220`). Folding into M gives correct timing *and* correct
data with a single byte-conserving reshape and one shared K/V operand. No
new primitive, no per-`g` loop.
### 10.5 Persistent-vs-scoped arena split **[recommend: two explicit arenas]** ### 10.5 Persistent-vs-scoped arena split **[recommend: two explicit arenas]**
ADR-0063 requires `m,l,O` to live outside the recycled scope. The cleanest ADR-0063 requires `m,l,O` to live outside the recycled scope. The cleanest
@@ -438,13 +467,40 @@ parity tests must use tolerances that accommodate this, and we should not
claim bf16 *accuracy* results — only timing/structure. Pre-existing claim bf16 *accuracy* results — only timing/structure. Pre-existing
simulator limitation; flagged so results aren't over-interpreted. simulator limitation; flagged so results aren't over-interpreted.
### 10.10 `tl.trans` is reshape-not-transpose **[decision: structural-first verify]**
Discovered in iteration 2: `tl.trans` only swaps shape metadata
(`tl_context.py:390`) and `MemoryStore.read` *reshapes* the bytes
(`memory_store.py:73`); `data_executor._execute_gemm` then `np.matmul`s
the reshaped operands (`data_executor.py:153-168`). So for **non-trivial
numeric data**, `Q·Kᵀ` via `tl.trans(K)` is mathematically a reshape, not
a transpose — wrong numbers (the baseline never notices: it runs on
all-zeros and only asserts op_log structure). **Decision:** treat the
simulator as a *performance* model (SPEC §0) — verify structure + timing +
determinism first; treat numeric parity as a bounded secondary check
(symmetric/identity inputs, or store K pre-transposed). A real
data-materialising `tl.transpose` is a *candidate* further primitive if
true numeric validation is ever required — **recommend deferring it**; it
is not needed for the performance-modeling goal. *Confirm you're content
with structural-first verification.*
### 10.11 GQA reuse = fold `G` into M, NOT a broadcast op **[finding, iteration 2]**
The headline mechanism finding. The deferred-GQA limitation
(`test_milestone_gqa_llama70b.py:137-142`) was blamed on a missing
broadcast / the MemoryStore nbytes check. Iteration 2 shows the real
cause is the baseline's head-packing reshape, and the fix is kernel
restructuring (§4.3), which needs **no new primitive**. This demoted
ADR-0061 from "the blocker" to "optional convenience" and reordered the
phase plan (§7). Recorded explicitly because it changes the project's
critical path: *GQA itself is not gated on any new feature; scale
(ADR-0063) and efficiency (ADR-0062) are the features that matter.*
--- ---
## 11. Risks ## 11. Risks
| Risk | Likelihood | Mitigation | | Risk | Likelihood | Mitigation |
|---|---|---| |---|---|---|
| Batched dot over leading axis unsupported | med | per-`g` fallback (§10.4); still correct | | Numeric parity blocked by trans-as-reshape | med | structural-first verify (§10.10); pre-transpose K or defer `tl.transpose` |
| Tree pairs not 1-hop on installed mesh | med | verify SFR neighbour table (§10.2); relabel PEs | | Tree pairs not 1-hop on installed mesh | med | verify SFR neighbour table (§10.2); relabel PEs |
| scratch_scope use-after-scope bug | med | two-arena discipline in shared helper; debug poison | | scratch_scope use-after-scope bug | med | two-arena discipline in shared helper; debug poison |
| prefetch buffer recycled while in flight | med | live `LoadFuture` buffers in persistent arena (§10.5) | | prefetch buffer recycled while in flight | med | live `LoadFuture` buffers in persistent arena (§10.5) |