gqa(adr): pivot ADR-0060 to composite hybrid + lazy tl.load; add ADR-0064 cost model

- ADR-0060: GEMMs (Q.Kt, P.V) via existing tl.composite (scheduler-managed
  tiling + K/V DMA streaming); softmax merge + IPCQ tree reduction stay in
  kernel. Front TL;DR pseudocode of the final composite kernel; new section
  B lists open design items (DDD sync, K pre-transpose, dma_read lever,
  kernel-vs-scheduler tiling, ring path).
- ADR-0062: redefined from a new load_async op to global lazy tl.load
  (non-blocking + auto-wait on first use; API unchanged; goldens regenerate).
- ADR-0064 (new): per-op-type CPU issue cost model (composite ~40ns >>
  primitive) so the hybrid's CPU-saturation win becomes measurable
  (currently dispatch_cycles=0 hides it). Cost-model impl deferred.
- KO mirrors for ADR-0060/0062/0064 (-ko suffix, adr-proposed).

Rationale: non-blocking CompositeCmd offloads tiling to PE_SCHEDULER,
decoupling CPU issue-rate from execution so the CPU can saturate the
engines; the prior 'composite = no latency benefit' claim was an artifact
of dispatch_cycles=0. Docs only; no production code changed.

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
2026-06-04 10:19:37 -07:00
parent 6d24b9306f
commit 7a9d4ec47b
6 changed files with 1675 additions and 215 deletions
@@ -16,10 +16,12 @@ GQA, causal, long-context kernel.
**Supporting ADRs** (efficiency / scale enablers — *not* GQA blockers;
see §8 correction): **ADR-0063** `tl.scratch_scope` (per-tile scratch
recycling — required for realistic context length), **ADR-0062**
`tl.load_async` (KV prefetch overlap — efficiency), **ADR-0061**
`tl.broadcast` (optional mask/general convenience). Real GQA itself needs
only kernel restructuring (§5.2).
recycling — required for realistic context length), **ADR-0062** lazy
`tl.load` (non-blocking load with auto-wait on first use → load/compute
overlap), **ADR-0061** `tl.broadcast` (optional mask/general
convenience). The two GEMMs (Q·Kᵀ, P·V) are issued as scheduler-managed
`tl.composite` commands (the existing `CompositeCmd`; no new command
kind); real GQA itself needs only kernel restructuring (§5.2).
**Algorithm lineage.** This kernel is **FlashAttention** (tiling +
online/streaming softmax with fused P·V — no full score matrix
@@ -32,6 +34,69 @@ known algorithms onto the kernbench **greenlet `tl` programming model**
---
## TL;DR — final GQA kernel (composite hybrid) pseudocode
The decision (§1) in one place: **GEMMs → `tl.composite`; softmax merge +
tree reduction → kernel; `tl.load` is lazy.** Per-KV-head, `G` folded into
the matmul M dim (real GQA, no broadcast). This is the reference shape; the
prose sections elaborate each piece.
```python
def gqa_attention(q_ptr, k_ptr, v_ptr, o_ptr,
counter, start_pe, N, q_block, scale, *, tl):
# ---- geometry (kernel arithmetic, §2) ----
pe_id = tl.program_id(axis=rank_axis)
G = H_q // H_kv # query heads per KV head (=8)
causal = q_block.is_prefill
for kv in range(H_kv): # one KV head per iteration (§5.2)
# Q group: G rows folded into M → [G·T_q, d]. Lazy load (ADR-0062):
# issue now, auto-wait at first use inside the first composite.
q_g = tl.load(q_base(kv), (G * q_block.T_q, d)) # [G·T_q, d]
my_len = valid_len(counter, start_pe, pe_id, N) if not causal else S_kv
n_tiles = ceil(my_len / TILE)
# persistent arena (outside scratch_scope, ADR-0063): -inf, 0, zeros
m, l, O = init_running(G * q_block.T_q, d)
for j in range(n_tiles):
if causal and tile_all_future(j, q_block):
continue # causal skip (kernel if)
with tl.scratch_scope(): # per-tile temporaries (ADR-0063)
# --- GEMM #1 on the scheduler: Q·Kⱼᵀ; K streamed by composite ---
Sj = tl.composite("gemm", a=q_g,
b=tl.ref(k_tile(kv, j), (d, TILE))) * scale # Kᵀ pre-stored [d,TILE]
if causal and tile_partial(j, q_block):
Sj = Sj + causal_mask(j, q_block) # additive boundary mask
# --- online softmax merge in the kernel (MATH ops) ---
m_new = tl.maximum(m, tl.max(Sj, axis=-1))
P = tl.exp(Sj - m_new)
corr = tl.exp(m - m_new)
l = l * corr + tl.sum(P, axis=-1)
# --- GEMM #2 on the scheduler: P·Vⱼ; V streamed by composite ---
Oj = tl.composite("gemm", a=P,
b=tl.ref(v_tile(kv, j), (TILE, d)))
O = O * corr + Oj # running merge
m = m_new
# ---- cross-PE combine (§4): log-sum-exp tree to root, or store ----
if N == 1:
tl.store(o_base(kv), O / l)
else:
tree_reduce_and_store(m, l, O, pe_id, N, o_base(kv)) # tl.send/tl.recv
```
> **Why this shape:** the two `tl.composite` calls offload tiling + K/V DMA
> streaming + cross-tile pipelining to PE_SCHEDULER (CPU issues coarse
> descriptors and runs ahead → engines stay saturated, §1); the softmax
> merge and the IPCQ tree reduction stay in the kernel because the existing
> `CompositeCmd` cannot carry cross-tile state. `K` is pre-stored
> transposed `[d, TILE]` to sidestep the reshape-not-transpose caveat
> (§3, §B). A bespoke "flash-composite" command kind is **not** introduced
> (§8 item 4).
---
## A. Relationship to existing kernbench work (read first)
kernbench **already runs FlashAttention with an online-softmax `(m, , O)`
@@ -47,11 +112,14 @@ Both are driven by `milestone-gqa-llama70b` (4 panels:
`src/kernbench/benches/milestone_gqa_llama70b.py`) and tested in
`tests/attention/test_milestone_gqa_llama70b.py`.
**They are written in the greenlet `tl` API — not composites:** `tl.load`,
`tl.dot`, `tl.softmax`/`tl.max`/`tl.sum`/`tl.exp`, `tl.send`/`tl.recv`,
and Python `-`/`*`/`/` on `TensorHandle` (each emits a `MathCmd`). The
running `(m, , O)` is just Python `TensorHandle`s threaded through the
loop. This **matters for this ADR's design** (see §1).
**They are written in the greenlet `tl` API:** `tl.load`, `tl.dot`,
`tl.softmax`/`tl.max`/`tl.sum`/`tl.exp`, `tl.send`/`tl.recv`, and Python
`-`/`*`/`/` on `TensorHandle` (each emits a `MathCmd`) — the GEMMs as
**blocking `tl.dot`**, not composites. The running `(m, , O)` is just
Python `TensorHandle`s threaded through the loop. This ADR keeps the
running state and the softmax merge in the kernel but **moves the two
GEMMs onto the scheduler-managed `tl.composite` path** (see §1) — this
**matters for the design**.
**Three deliberate limitations of the baseline** — exactly what an
*efficient GQA* kernel must lift:
@@ -74,7 +142,8 @@ loop. This **matters for this ADR's design** (see §1).
allocator leaks per-tile temporaries
(`test_milestone_gqa_llama70b.py:123-148`) and there is no causal
masking / tiling → fixed by **ADR-0063** (recycling) + §5 (tiling,
causal skip) + **ADR-0062** (prefetch).
causal skip) + composite K/V streaming (§3) + **ADR-0062** (lazy load
overlap).
**Documentation debt (out of scope but recorded):** the baseline cites
ADR-0055/0056/0057/0058/0059, **none of which exist as files** — they are
@@ -103,10 +172,14 @@ Hardware recap:
(`tl_context.py:402-499`).
- **Composite command** (`CompositeCmd`, `pe_commands.py:144-162`): a
single GEMM (or MATH) *head* plus element-wise *epilogue* stages
(`bias/relu/scale/add/...`). It is **not** a general multi-op DAG: it
cannot chain two GEMMs, cannot carry register state across instances,
and cannot pop/wait on IPCQ. This ADR therefore does **not** require
composites for the attention inner loop (see §1, §8).
(`bias/relu/scale/add/...`), issued **non-blocking** to PE_SCHEDULER,
which generates a tile plan and streams DMA→GEMM→write per tile
(ADR-0014 D6; `pe_scheduler.py:104-143`). It is **not** a general
multi-op DAG: it cannot chain two GEMMs, cannot carry register state
across instances, and cannot pop/wait on IPCQ. This ADR therefore issues
**each** of the two attention GEMMs (Q·Kᵀ, P·V) as its own composite and
keeps the cross-GEMM softmax merge + the IPCQ reduction in the kernel —
it does **not** need a new "flash-composite" command kind (see §1, §8).
- **Allocation policy (SP):** one query head per CUBE in the multi-user
panels; the `G` query heads of a KV group map within a CUBE's PEs.
@@ -194,51 +267,82 @@ projection (downstream), score `S` and probs `P` (never materialised).
## 1. Decision (mechanism)
**Implement the kernel in the greenlet `tl` programming model, not as
composites.** Rationale grounded in kernbench's execution + latency model:
**Implement the kernel as a *hybrid*: issue the two GEMMs (Q·Kᵀ and P·V)
as scheduler-managed `tl.composite(op="gemm")` commands, and keep the
online-softmax merge and the cross-PE reduction as kernel-level `tl`
ops.** `tl.load` is **lazy** (non-blocking; the wait is auto-inserted at
first use of the loaded data — ADR-0062), so explicit HBM loads overlap
the compute that follows. Rationale grounded in kernbench's execution +
latency model:
- The simulator charges latency **per op on its modelled component**
(GEMM on PE_GEMM, vector ops on PE_MATH, DMA on PE_DMA — `pe_gemm.py`,
`pe_math.py`, `pe_dma.py`). "Fusing" several ops into one
`CompositeCmd` does **not** reduce modelled latency; the stages still
run on the same engines. The win a real fused kernel gets — *overlap*
(prefetch) and *small working set* (scratch recycling) — is obtained
here by **ADR-0062** and **ADR-0063**, which are smaller and reusable.
- The running `(m, , O)` flash state is naturally a set of Python
`TensorHandle`s threaded through the loop (the baseline already does
this). No composite-carried register state is needed.
- Cross-PE combination uses kernel-level `tl.send`/`tl.recv` (the
baseline already does this and it is tested). No composite-driven IPCQ
push is needed.
- **GEMM tiling is offloaded to PE_SCHEDULER.** A `CompositeCmd` is
non-blocking (`kernel_runner.py:182-191`, `pe_scheduler.py:104-121`):
the kernel pushes **one coarse descriptor** (M = `G·T_q`, the whole
per-PE tile sweep) and the scheduler generates the tile plan and streams
DMA→GEMM→write per tile (ADR-0014 D6). K/V are `tl.ref` operands the
scheduler streams from HBM, so per-tile **K/V prefetch is the
scheduler's job** — no explicit prefetch op. The CPU (greenlet) is freed
to issue the **next** composite while the current one runs, so the
scheduler keeps the GEMM engine saturated across tiles.
- **This reflects the hardware** and decouples CPU issue-rate from
execution-rate. The blocking per-op `tl.dot` path, by contrast, stalls
the CPU on every GEMM and leaves GEMM-engine **bubbles** during the
interleaved softmax MATH ops; it is realistic only if the CPU can keep
up with fine-grained per-tile issue.
- **The running `(m, , O)` flash state stays Python `TensorHandle`s**
threaded through the loop (the baseline already does this); the softmax
merge (max/exp/sum/rescale) is kernel-level `tl` MATH **between** the two
GEMM composites. The existing `CompositeCmd` cannot chain two GEMMs or
carry cross-tile register state (§0), so the merge necessarily lives in
the kernel — this is the hybrid split, not a limitation worked around.
- **Cross-PE combination** is a log-sum-exp **tree** over `(m, , O)`
after P·V, via kernel-level `tl.send`/`tl.recv` (§4) — unchanged.
So the per-tile inner pipeline is the op sequence
So the per-tile inner pipeline is:
```
K_j load → Q·Kⱼᵀ → (+ causal mask) → online-softmax update → V_j load → P·Vⱼ → running (m,,O) update
q_g = tl.load(Q group) # lazy; auto-wait at first use
per tile j:
Sⱼ = tl.composite("gemm", a=q_g, b=tl.ref(Kⱼ)) → Sⱼ # scheduler streams Kⱼ DMA + GEMM
Sⱼ += maskⱼ # kernel MATH, boundary tile only
online-softmax: mⱼ, m_new, P, corr, # kernel MATH
Oⱼ = tl.composite("gemm", a=P, b=tl.ref(Vⱼ)) → Oⱼ # scheduler streams Vⱼ DMA + GEMM
O = O*corr + Oⱼ; m = m_new # kernel MATH (running merge)
```
issued as ordinary `tl.*` ops, with the **next** tile's K/V issued via
`tl.load_async` (ADR-0062) so its DMA overlaps the current tile's
compute, and each tile's temporaries wrapped in `tl.scratch_scope`
(ADR-0063) so scratch stays O(1).
Cross-PE combination (KV-parallel / SP) is a **log-sum-exp tree
reduction** over `(m, , O)` after P·V, flowed through IPCQ with
kernel-level `tl.send`/`tl.recv` (§4).
with each tile's MATH temporaries wrapped in `tl.scratch_scope`
(ADR-0063) so scratch stays O(1), and the next tile's composites issued
before the current tile's results are waited on (non-blocking handles) so
the scheduler pipelines across tiles.
Control flow (tile skip, mask generation, reduction scheduling, address
arithmetic) lives in the **kernel** (plain Python `if`/arithmetic in the
greenlet body). This is exactly what kernbench's greenlet model already
permits (`kernel_runner.py`, ADR-0020 D3).
> **Why this is the efficient choice.** It reuses the proven baseline
> kernels and the proven IPCQ collective; the only genuinely new
> machinery is three small, general primitives (ADR-0061/0062/0063). The
> originally-proposed "everything-in-one-composite + composite IPCQ push
> + composite-carried state" would require a bespoke flash-composite
> command type, a register-lifetime model across composites, and an
> IPCQ-push epilogue — large, special-purpose, and with no latency
> benefit over the greenlet path. Those are **rejected**; see §8.
> **What this supersedes.** An earlier iteration proposed a pure greenlet
> primitive path (all `tl.dot`, no composite) on the argument that
> "composite yields no latency benefit." That holds **only because** the
> simulator currently charges **zero** per-op CPU issue cost
> (`dispatch_cycles=0`, `pe_cpu.py`) — it models away exactly the CPU
> issue-rate / DMA-program cost that descriptor offload exists to hide.
> The hybrid is the faithful representation of an efficient kernel. The
> **measurable** size of the win (can the CPU saturate the engines for
> many tiles?) is gated on modelling an op-type-differentiated issue cost,
> tracked as future work (cost model; §9). Even at `dispatch_cycles=0` the
> non-blocking composite path fills the GEMM-engine bubbles the blocking
> `tl.dot` path leaves.
>
> **Why this is the efficient choice.** GEMM tiling + DMA streaming +
> cross-tile pipelining are offloaded to the proven `CompositeCmd`
> scheduler path; the softmax merge and the proven IPCQ collective stay in
> the kernel. The only genuinely new machinery is two small, general
> primitives (ADR-0062 lazy `tl.load`, ADR-0063 `tl.scratch_scope`); the
> reduction reuses `tl.send`/`tl.recv`. A bespoke "flash-composite"
> command (one kind internalising the softmax merge + carried register
> state + an IPCQ-push epilogue) is **not** built — large, special-purpose,
> and its only delta over this hybrid (full softmax offload) is not
> justified at the current modelling fidelity; see §8.
---
@@ -283,17 +387,20 @@ Driver = bases + counter + rotation. The single formula
## 3. Per-tile op sequence (greenlet `tl`)
One iteration = one KV tile on one PE. Real `tl` names (`tl_context.py`),
with `tl.load_async` (ADR-0062), `tl.broadcast` (ADR-0061),
`tl.scratch_scope` (ADR-0063):
One iteration = one KV tile on one PE. The two GEMMs are
`tl.composite(op="gemm")` (scheduler-managed tiling + K/V DMA streaming);
the softmax merge is kernel `tl` MATH between them. Real `tl` names
(`tl_context.py`), with lazy `tl.load` (ADR-0062), `tl.scratch_scope`
(ADR-0063):
```python
# running state (persistent arena — allocated once, outside the scope)
# m: [G, T_q] l: [G, T_q] O: [G, T_q, d]
q_g = tl.load(Q_group_ptr, (G*T_q, d)) # lazy; auto-wait at first use (ADR-0062)
with tl.scratch_scope(): # per-tile temporaries recycled
Kj = tl.wait(f_k[j]) # prefetched (ADR-0062)
Sj = tl.dot(q_g, tl.trans(Kj)) * softmax_scale # [G, TILE]; q_g is GQA-batched
with tl.scratch_scope(): # per-tile MATH temporaries recycled
Sj = tl.composite("gemm", a=q_g, # [G·T_q, TILE]; scheduler streams Kⱼ DMA
b=tl.ref(K_base + j*TILE*d, (TILE, d))) * softmax_scale
if mask_j is not None:
Sj = Sj + mask_j # additive causal mask (boundary tile)
m_j = tl.max(Sj, axis=-1)
@@ -301,32 +408,35 @@ with tl.scratch_scope(): # per-tile temporaries rec
P = tl.exp(Sj - m_new) # no full-matrix softmax; streaming
corr = tl.exp(m - m_new) # rescale factor for old accumulators
l = l * corr + tl.sum(P, axis=-1)
Vj = tl.wait(f_v[j])
O = O * corr + tl.dot(P, Vj) # P·V folded into running O
Oj = tl.composite("gemm", a=P, # [G·T_q, d]; scheduler streams Vⱼ DMA
b=tl.ref(V_base + j*TILE*d, (TILE, d)))
O = O * corr + Oj # running merge (kernel MATH)
m = m_new
if j + PREFETCH < n_tiles: # keep the pipeline full
f_k[j + PREFETCH] = tl.load_async(K_base + (j+PREFETCH)*TILE*d, (TILE, d))
f_v[j + PREFETCH] = tl.load_async(V_base + (j+PREFETCH)*TILE*d, (TILE, d))
```
Notes:
- `q_g` is the GQA-batched query reshaped to `[G·T_q, d]` (the `G` group
rows folded into the matmul M dim; byte-conserving). One K/V tile load
serves all `G·T_q` rows — the GQA reuse lever — with no broadcast.
- `tl.trans(Kj)` is **metadata-only** in kernbench (`tl_context.py:390`),
and `MemoryStore.read` *reshapes* rather than transposes
rows folded into the matmul M dim; byte-conserving). One K/V tile serves
all `G·T_q` rows — the GQA reuse lever — with no broadcast.
- **K/V are `tl.ref` operands** the composite scheduler streams from HBM
per tile (`pe_scheduler.py:104-143`): that *is* the prefetch/pipeline,
so there is no explicit prefetch op. Issuing tile `j+1`'s composites
before waiting on tile `j` (non-blocking handles) keeps the scheduler
pipelined across tiles and the GEMM engine saturated.
- `tl.trans` 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
non-trivial numeric data it yields a reshape-not-transpose, so Q·Kᵀ via
a transposed K needs care (§11) — store K pre-transposed `[d, TILE]`, 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
overlaps the Q·Kᵀ + softmax of the same/earlier tile.
- Masking: the kernel builds the boundary-tile mask from query/KV global
offsets and adds it (`Sj + mask_j`); full-past tiles pass `None`;
full-future tiles are **skipped** (the `if` never enqueues them).
- No `CompositeCmd` is used. (A future `tl.composite` flash kind is an
*optional* optimisation — §8 item 4 — not a requirement.)
- A **new** "flash-composite" command kind (one that internalises the
softmax merge + carried `(m,,O)`) is **not** used; the existing
`CompositeCmd` covers each GEMM and the merge stays in the kernel
(§1, §8 item 4).
---
@@ -395,13 +505,12 @@ def attn_kernel(q_ptr, k_ptr, v_ptr, o_ptr, counter, start_pe, N,
pe_id = tl.program_id(axis=rank_axis)
my_len = valid_len(counter, start_pe, pe_id, N)
n_tiles = ceil(my_len / TILE)
q_g = load_Q_group(q_ptr) # [G,d] decode / [G,T_q,d] prefill; tl.broadcast for GQA
q_g = load_Q_group(q_ptr) # [G·T_q, d]; lazy tl.load, G folded into M (no broadcast)
m, l, O = init_running() # persistent arena: -inf, 0, zeros
prime_prefetch(k_ptr, v_ptr, n_tiles) # tl.load_async first PREFETCH tiles
for j in range(n_tiles):
for j in range(n_tiles): # K/V streamed per tile by the composite scheduler (§3)
if tile_all_future(j, q_block): # causal skip (kernel if)
continue
run_tile(j, mask_or_null(j, q_block)) # §3 op sequence, in tl.scratch_scope
run_tile(j, mask_or_null(j, q_block)) # §3: 2 composites + softmax MATH, in tl.scratch_scope
if N == 1:
tl.store(o_ptr, O / l) # no reduction
else:
@@ -414,18 +523,18 @@ def attn_kernel(q_ptr, k_ptr, v_ptr, o_ptr, counter, start_pe, N,
- **GQA reuse is the whole game** (decode is KV-load-bound): the `G=8`
query rows of the KV head are folded into the matmul **M** dimension
(`q_g` reshaped `[G, T_q, d] → [G·T_q, d]`, a byte-conserving reshape).
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
Then `Q·Kᵀ` is `composite([G·T_q, d], Kᵀ[d, TILE]) → [G·T_q, TILE]` and
`P·V` is `composite([G·T_q, TILE], V[TILE, d]) → [G·T_q, d]`. The KV tile
(`[TILE, d]`) is the shared `K`/`V` operand — **streamed 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).
GEMM. No broadcast of K/V is needed; `m = G·T_q` in the composite's tile
plan also makes the 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
**one-shot** partial attention (one `tl.dot` for Q·Kᵀ, one softmax, one
`tl.dot` for P·V) — exactly the baseline `_attention_mesh_mlo`
`_partial_attention`, just GQA-batched. Tiling (§3) only kicks in when
`S_pe` exceeds the scratch scope's tile budget.
**one-shot** partial attention (one composite for Q·Kᵀ, one softmax, one
composite for P·V) — exactly the baseline `_attention_mesh_mlo`
`_partial_attention`, just GQA-batched and on the composite path. Tiling
(§3) only kicks in when `S_pe` exceeds the scratch scope's tile budget.
### 5.3 DECODE, with SP / KV-parallel (`N=8`)
- One request's KV is round-robin across 8 PEs; each owns ≈`my_len`
@@ -484,11 +593,13 @@ ADR adds GQA reuse, causal step-skip, and `recv_async` overlap.
| tile skip (future), mask generation, causal bounds | **kernel** (Python `if` + arithmetic in greenlet body) |
| address / offset / valid-length arithmetic | **kernel** (from counter) |
| reduction scheduling, IPCQ send/recv ordering | **kernel** (static tree) |
| K/V load, Q·Kᵀ, mask add, softmax math, P·V | **`tl` ops** on PE engines |
| Q·Kᵀ, P·V (incl. per-tile K/V DMA streaming + tiling) | **`tl.composite`** PE_SCHEDULER |
| Q load, mask add, softmax math, running `(m,,O)` merge | **`tl` ops** on PE engines (kernel-issued) |
The kernel decides; the `tl` ops execute already-decided work. This is
exactly the greenlet model kernbench already supports — no new control
abstraction.
The kernel decides; the GEMMs are offloaded to the scheduler as
composites; the remaining `tl` ops execute already-decided work. This is
exactly the greenlet + composite model kernbench already supports — no new
control abstraction.
---
@@ -502,8 +613,13 @@ new primitive** — only the kernel restructuring in §5.2 (per KV head,
**Algorithm work in the kernel (no new primitive; existing `tl` API):**
- **GQA Q-axis batching** (the reuse lever) — fold `G·T_q` into the matmul
M dim per KV head (§5.2); `_view`-style byte-conserving reshape + 2-D
`tl.dot`. Runs today in both timing and data mode.
M dim per KV head (§5.2); `_view`-style byte-conserving reshape; the
GEMM is a `tl.composite(op="gemm")` with M = `G·T_q`. Runs today in both
timing and data mode.
- **GEMMs via composite** (§1/§3) — Q·Kᵀ and P·V each issued as a
non-blocking `tl.composite(op="gemm")`; PE_SCHEDULER tiles them and
streams the `tl.ref` K/V operands' DMA (existing `CompositeCmd`; no new
command kind).
- Tree reduction to root (§4) replacing the baseline all-to-all fan-out —
pure kernel control flow over `tl.send`/`tl.recv`.
- Causal tile skip + additive boundary mask (§3/§5.4) — kernel `if` +
@@ -517,9 +633,12 @@ new primitive** — only the kernel restructuring in §5.2 (per KV head,
*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.
2. **Lazy `tl.load`** — **ADR-0062** (non-blocking load + auto-wait on
first use; API surface unchanged). *Efficiency*: overlaps explicit
loads (the Q group, non-composite kernels) with following compute. The
per-tile **K/V** prefetch is handled by the composite scheduler (§1),
so this covers the remaining explicit loads. Global semantics change →
existing goldens regenerate (ADR-0062 D3).
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
@@ -528,16 +647,16 @@ new primitive** — only the kernel restructuring in §5.2 (per KV head,
**Explicitly REJECTED (efficient alternative chosen):**
4. ~~Composite that chains DMA→MM→VEC→DMA→MM→VEC with carried `(m,,O)`
register state and a tail IPCQ push.~~ Replaced by the greenlet path:
per-op latency is identical; overlap comes from ADR-0062; small
working set from ADR-0063; running state is Python handles; the IPCQ
push is `tl.send`. Building a bespoke flash-composite command type +
cross-composite register lifetime + an IPCQ-push epilogue is large,
special-purpose, and yields no latency benefit. **A tiled
`tl.composite` "flash" kind remains a possible *future* optimisation**
(it would fold prefetch+recycling into the scheduler), but it is not
required for an efficient kernel and is out of scope here.
4. ~~A bespoke "flash-composite" command kind that internalises the whole
inner loop — DMA→MM→VEC→DMA→MM→VEC with carried `(m,,O)` register
state and a tail IPCQ push.~~ The two GEMMs **do** use the existing
`CompositeCmd` (§1/§3) — that gives scheduler-managed tiling, K/V DMA
streaming, and cross-tile pipelining. What is rejected is a **new**
command kind that also absorbs the softmax merge + cross-tile register
lifetime + an IPCQ-push epilogue: it is large and special-purpose, and
its only delta over the hybrid (full softmax offload) is not justified
at the current modelling fidelity (per-op CPU issue cost = 0; see §1).
Revisit if the cost model (§9) makes full offload measurably worthwhile.
5. ~~Hardware `pop`-as-dependency.~~ Out of scope (§6).
6. ~~RoPE / QKV projection / KV-cache write inside this kernel.~~ Upstream
`qkv_rope` (P1P5). Folding RoPE in would force re-rotating past tiles
@@ -547,16 +666,24 @@ new primitive** — only the kernel restructuring in §5.2 (per KV head,
## 9. Open tuning items (measured in kernbench, not blocking)
1. **KV tile prefetch depth** (ADR-0062) — dominant for KV-load-bound;
sweep 2→4.
1. **Composite tile-pipeline depth** — dominant for KV-load-bound; how far
ahead the kernel issues non-blocking composites before waiting, and the
scheduler's per-tile streaming depth.
2. **PE↔mesh-neighbour mapping for the reduction tree** — ensure each
depth-`⌈log₂ N⌉` pair is a physical `N/S/E/W` neighbour; bad mapping
adds hops. Verify against the SFR install.
3. **TILE size** — balance scratch residency (`K_tile + S/P + V_tile +
O_acc + G-way GQA`) against DMA efficiency; interacts with ADR-0063.
3. **TILE size** — balance scratch residency (`S/P tiles + O_acc + G-way
GQA`) against DMA efficiency; interacts with ADR-0063 and the
scheduler's `TILE_M/K/N` (`pe_scheduler.py`).
4. **Ring buffer ping-pong vs `recv_async` depth** in §5.5.
5. **One-shot vs tiled crossover for decode** (§5.2) — the `S_pe`
threshold where tiling beats a single `tl.dot`.
threshold where tiling beats a single composite.
6. **Per-op CPU issue cost (cost model)** — currently `dispatch_cycles=0`
(`pe_cpu.py`), so composite-vs-primitive issue overhead is invisible.
An op-type-differentiated issue cost (a `tl.composite` descriptor push
≫ a primitive op) is what makes the hybrid's CPU-saturation win
**measurable** (§1). Specified in **ADR-0064**; tracked as separate
future work.
---
@@ -612,7 +739,79 @@ secondary check.
count equals the lower-triangular tile count, not the full grid.
5. **Long context (ADR-0063):** a sweep at `S` that overflows 1 MiB
without scopes completes and matches the reference.
6. **Prefetch overlap (ADR-0062):** end-to-end latency of the tiled sweep
is below the serial `Σ(load+compute)` (overlap is real).
7. **Determinism:** identical inputs → identical op_log + latency
6. **Load/compute overlap:** end-to-end latency of the tiled sweep is
below the serial `Σ(load+compute)` — from the composite scheduler
streaming K/V per tile (§1/§3) and lazy `tl.load` (ADR-0062) overlapping
the Q load. (Overlap is real modelled concurrency, not a subtraction.)
7. **Composite GEMM offload (structural):** each tile's Q·Kᵀ and P·V emit a
`CompositeCmd` (non-blocking) to PE_SCHEDULER, not a blocking `tl.dot`;
op_log shows the composite tile plan and the kernel issues the next
tile's composites before waiting (cross-tile pipelining).
8. **Determinism:** identical inputs → identical op_log + latency
(SPEC §0.1).
---
## B. Open design items from the hybrid pivot (review later)
These arose when the decision moved from a pure greenlet primitive path to
the **composite hybrid + lazy `tl.load`** (this revision). None blocks the
design; each needs a verification pass during implementation. Recorded here
(rather than asked) per the working agreement — the recommendation is my
predicted default; revise on review.
1. **DDD-0060 is not yet synced.** The Detailed Design Document still
describes the old `tl.load_async` double-buffer path and primitive
`tl.dot` inner loop (its §4.3/§5/§10). It must be updated to the hybrid
(composite GEMMs, lazy load, K pre-transposed). *Left for review*
because the DDD is a derived how-to and a large rewrite; the ADR is now
the authoritative record. **Recommend:** sync DDD as a follow-up before
implementation starts.
2. **K operand orientation for the composite GEMM.** Q·Kᵀ needs `b =
[d, TILE]`, but the KV cache stores K as `[S_pe, d]`. `tl.trans` is
metadata-only and `MemoryStore.read` reshapes, not transposes
(`memory_store.py:73`) — so a runtime transpose is wrong for non-trivial
data. **Recommend:** store K **pre-transposed** `[d, S_pe]` in the cache
(the pseudocode and §3 assume this), making `tl.ref(k_tile, (d, TILE))`
a contiguous slice. Verify the upstream `qkv_rope` write layout supports
this, or add a real `tl.transpose` (heavier; deferred).
3. **Composite output buffer vs `tl.scratch_scope`.** Each Q·Kᵀ composite
writes `Sj` to an `out_addr`; the kernel then reads it for the softmax
MATH. That output buffer, and the in-flight composites' targets, must
live where the per-tile `scratch_scope` (ADR-0063) will **not** recycle
them before they are consumed — same discipline as in-flight lazy loads
(ADR-0062 D-Negative). **Recommend:** composite outputs for the *current*
tile live in the scoped arena (consumed same iteration); the persistent
`(m,,O)` stays outside. Verify no use-after-recycle when the next
tile's composites are issued early (cross-tile pipelining).
4. **GQA `dma_read_count` lever under composite streaming.** The lever
(§11.2: K/V `dma_read_count` independent of `G`) assumes the composite
emits **one** K/V tile DMA reused across all `G·T_q` M-rows. The
scheduler's `generate_gemm_plan` tiles by `TILE_M/K/N`
(`pe_scheduler.py:35-37`, 32/64/32) — confirm the M-tiling over `G·T_q`
does **not** re-issue the shared K/V tile DMA per M-tile (i.e. operand
DMA is shared across M-tiles, or the lever weakens). **Recommend:**
assert it in the levers test; if violated, the GQA win is in compute
only, not DMA — still correct, but the headline changes.
5. **Kernel TILE vs scheduler `TILE_M/K/N`.** The kernel reasons about a
logical KV `TILE`; the scheduler re-tiles internally at fixed
`TILE_M/K/N`. Two tiling layers interact (scratch residency, pipeline
depth). **Recommend:** treat the kernel TILE as the K/V streaming
granularity and let the scheduler sub-tile the GEMM; document the
relationship in the DDD and sweep both (§9 items 1, 3).
6. **Cost model is a separate ADR.** The hybrid's CPU-saturation benefit is
invisible while `dispatch_cycles=0`. The per-op-type issue-cost model is
specified in **ADR-0064**; this ADR's §1/§9 depend on it for the
*measurable* (not just structural) win. **Recommend:** land ADR-0064's
model before claiming hybrid latency wins in the eval.
7. **Ring path (§5.5) GEMMs.** §5.5 still describes the ring fold with
primitive ops + `recv_async`. For consistency the ring's per-step Q·Kᵀ /
P·V should also be composites; the IPCQ `recv_async` overlap is
orthogonal and stays. **Recommend:** apply the same hybrid shape to the
ring step during implementation; low risk, mirrors §3.