diff --git a/docs/adr-proposed/ADR-0060-algo-gqa-fused-attention-ahbm.md b/docs/adr-proposed/ADR-0060-algo-gqa-fused-attention-ahbm.md index 27af36a..54076b8 100644 --- a/docs/adr-proposed/ADR-0060-algo-gqa-fused-attention-ahbm.md +++ b/docs/adr-proposed/ADR-0060-algo-gqa-fused-attention-ahbm.md @@ -56,51 +56,105 @@ only the irreducible data moves — `(m,ℓ,O)` (decode) or KV blocks (prefill). Both share §3's composite-hybrid inner tile (GEMMs → `tl.composite`, softmax merge in kernel, lazy `tl.load`). +### Kernel 1 — DECODE + SP (head-replicated, KV static shard, 2-level reduce; NO ring) + +The decode inner tile has a hard chain `Q·Kᵀ → softmax → P·V`: softmax +(`tl.max(Sj)`) waits for the first GEMM, so a naïve loop stalls the CPU on +`Sj` and **leaves the GEMM engine idle during softmax** (a bubble). Three +variants trade CPU/HW complexity against that bubble (ship **opt3** now; +**opt2** needs a new command kind — revisit with the cost model, §8/ADR-0064): + ```python -# ── Kernel 1 — DECODE + SP: head-replicated, KV static shard, 2-level reduce (NO ring) ── -def gqa_decode_sp(q_ptr, k_ptr, v_ptr, o_ptr, S_kv_local, d, C, P, scale, *, tl): +# OPTION 1 — current CompositeCmd (today; HAS the GEMM-engine bubble) +def gqa_decode_v1(q_ptr, k_ptr, v_ptr, o_ptr, S_kv_local, d, C, P, scale, *, tl): cube_id = tl.program_id(axis=1); pe_id = tl.program_id(axis=0) - kv = head_of_group(cube_id) - q_g = tl.load(q_base(kv), (G * 1, d)) # T_q=1; Q replicated: G heads stacked into M (M-fold) - m, l, O = init_running(G, d) # persistent arena - for j in range(ceil(S_kv_local / TILE)): # sweep ONLY my KV shard (resident, no move) + kv = head_of_group(cube_id) # CUBE → its KV head + q_g = tl.load(q_base(kv), (G * 1, d)) # T_q=1; Q replicated: G heads stacked into M (M-fold) + m, l, O = init_running(G, d) # persistent arena (-inf, 0, 0) + for j in range(ceil(S_kv_local / TILE)): # sweep ONLY my KV shard (resident, no move) with tl.scratch_scope(): - Sj = tl.composite("gemm", a=q_g, b=tl.ref(k_tile(j), (d, TILE))) * scale + Sj = tl.composite("gemm", a=q_g, b=tl.ref(k_tile(j), (d, TILE)), epi=[scale]) + # ↓ CPU auto-waits on Sj → GEMM engine IDLE during softmax (bubble) m2 = tl.maximum(m, tl.max(Sj, -1)); P = tl.exp(Sj - m2); corr = tl.exp(m - m2) l = l * corr + tl.sum(P, -1) O = O * corr + tl.composite("gemm", a=P, b=tl.ref(v_tile(j), (TILE, d))); m = m2 - hierarchical_reduce_and_store(m, l, O, cube_id, pe_id, C, P, o_base(kv)) # ← reduce, NO ring + hierarchical_reduce_and_store(m, l, O, cube_id, pe_id, C, P, o_base(kv)) # §4 2-level reduce +``` -# ── Kernel 2 — PREFILL + SP: 1 Q head per CUBE (head-parallel), Ring KV (NO reduce) ── +```python +# OPTION 3 — software pipelining (current primitives; bubble removed) ← SHIP THIS +def gqa_decode_v3(q_ptr, k_ptr, v_ptr, o_ptr, S_kv_local, d, C, P, scale, *, tl): + cube_id = tl.program_id(axis=1); pe_id = tl.program_id(axis=0) + kv = head_of_group(cube_id); q_g = tl.load(q_base(kv), (G * 1, d)) + m, l, O = init_running(G, d); n = ceil(S_kv_local / TILE) + Sb = double_buffer() # 2 persistent Sj buffers (outside scratch_scope) + h = tl.composite("gemm", a=q_g, b=tl.ref(k_tile(0), (d, TILE)), out=Sb[0], epi=[scale]) # prime + for j in range(n): + Sj = Sb[j % 2] + if j + 1 < n: # ← issue NEXT Q·Kᵀ before softmax → fills GEMM engine + h = tl.composite("gemm", a=q_g, b=tl.ref(k_tile(j+1), (d, TILE)), out=Sb[(j+1)%2], epi=[scale]) + with tl.scratch_scope(): + m2 = tl.maximum(m, tl.max(Sj, -1)); P = tl.exp(Sj - m2); corr = tl.exp(m - m2) + l = l * corr + tl.sum(P, -1) + O = O * corr + tl.composite("gemm", a=P, b=tl.ref(v_tile(j), (TILE, d))); m = m2 + hierarchical_reduce_and_store(m, l, O, cube_id, pe_id, C, P, o_base(kv)) +``` + +```python +# OPTION 2 — ex_composite, 2-split (NEW flash-epilogue cmd; gives K-before-V DMA priority) +def gqa_decode_v2(q_ptr, k_ptr, v_ptr, o_ptr, S_kv_local, d, C, P, scale, *, tl): + cube_id = tl.program_id(axis=1); pe_id = tl.program_id(axis=0) + kv = head_of_group(cube_id); q_g = tl.load(q_base(kv), (G * 1, d)) + acc = init_running(G, d) # (m,l,O): scheduler-updated flash accumulator + for j in range(ceil(S_kv_local / TILE)): + Sj = tl.composite("gemm", a=q_g, b=tl.ref(k_tile(j), (d, TILE)), epi=[scale]) # #1: reads K (priority) + tl.ex_composite("softmax_pv", s=Sj, v=tl.ref(v_tile(j), (TILE, d)), acc=acc, scale=scale) # #2: reads V + # ↑ both non-blocking; CPU never waits intra-tile → max run-ahead, fewest issues + hierarchical_reduce_and_store(*acc, cube_id, pe_id, C, P, o_base(kv)) +``` + +| | new HW cmd | GEMM bubble | CPU intra-tile wait | issues | now | +|---|---|---|---|---|---| +| **opt1 current** | no | **yes** | yes | `O(tiles·ops)` | ✓ | +| **opt3 sw-pipe** | no | no | yes (reordered) | `O(tiles·ops)` | ✓ | +| **opt2 ex_composite** | **#2 only** | no | **no** | `O(tiles)` | ✗ (build #2) | + +(`#1` = the *existing* GEMM composite + `scale`; only `#2` = softmax + P·V + +the stateful online-softmax accumulator is new. MATH engine already has +max/sum/exp — the new part is the flash accumulator, not the ops.) + +### Kernel 2 — PREFILL + SP (1 Q head per CUBE, head-parallel; Ring KV, NO reduce) + +```python def gqa_prefill_sp(q_ptr, k_ptr, v_ptr, o_ptr, T_q, S_kv_local, d, C, scale, q_block, cube_start, *, tl): - i = tl.program_id(axis=1) - cube_start # this CUBE's query head = its start KV slice - q = tl.load(q_ptr, (T_q, d)) # MY head's query rows (resident) + i = tl.program_id(axis=1) - cube_start # this CUBE's Q head = its start KV slice + q = tl.load(q_ptr, (T_q, d)) # MY Q head's query rows (resident) Kc = tl.load(k_ptr, (S_kv_local, d)); Vc = tl.load(v_ptr, (S_kv_local, d)); src = i m, l, O = init_running(T_q, d) - for step in range(C): # ── Ring KV: rotate blocks around C CUBEs ── + for step in range(C): # ── Ring KV: rotate blocks around C CUBEs ── f = None if step < C - 1: tl.send("ring+", Kc); tl.send("ring+", Vc) # rotate out f = (tl.recv_async("ring-", (S_kv_local, d)), tl.recv_async("ring-", (S_kv_local, d))) if not block_all_future(q_block, slice_pos(src)): # causal skip whole-future blocks - S = tl.composite("gemm", a=q, b=tl.ref(Kc, (d, TILE_S))) * scale + S = tl.composite("gemm", a=q, b=tl.ref(Kc, (d, TILE_S)), epi=[scale]) if block_partial(q_block, slice_pos(src)): S = S + causal_mask(q_block, slice_pos(src)) m2 = tl.maximum(m, tl.max(S, -1)); P = tl.exp(S - m2); corr = tl.exp(m - m2) l = l * corr + tl.sum(P, -1); O = O * corr + tl.composite("gemm", a=P, b=tl.ref(Vc)); m = m2 if f: Kc, Vc = tl.wait(f[0]), tl.wait(f[1]); src = (src - 1) % C - tl.store(o_ptr, O / l) # MY head's rows — NO reduce + tl.store(o_ptr, O / l) # MY Q head's rows — NO reduce ``` > **Why two shapes.** Decode keeps KV put and moves the tiny `(m,ℓ,O)` > (§4 2-level reduce); prefill moves KV (Ring, §5.5) and keeps each head's -> big output local. Both use §3's composite-hybrid tile (GEMMs → -> `tl.composite`, softmax merge in kernel, lazy `tl.load`). `K` is pre-stored +> big output local. Both use §3's composite-hybrid tile. `K` is pre-stored > transposed to sidestep the reshape-not-transpose caveat (§3, §B); no -> bespoke "flash-composite" kind (§8 item 4). **Output head distribution -> differs** — decode lands all `G` heads at the CUBE-Group root; prefill -> leaves one Q head per CUBE (§0.5.4). The decode kernel above is the -> reference shape; **§5.6 gives 3 CPU-pipelining variants** of it -> (current / software-pipelined / `ex_composite`). +> bespoke "flash-composite" kind on the decode critical path (§8 item 4). +> **Output head distribution differs** — decode lands all `G` heads at the +> CUBE-Group root; prefill leaves one Q head per CUBE (§0.5.4). +> The opt2/opt3 variants are a **decode** concern: in prefill the causal `if` +> is kernel control flow that **cannot enter a composite**, and the ring +> already overlaps via `recv_async`. --- @@ -722,74 +776,25 @@ The baseline `_attention_mesh_kv` already implements the ring fold; this ADR adds GQA reuse, the head-parallel placement, causal step-skip, and the composite-hybrid inner tile (§3). -### 5.6 Decode CPU-pipelining variants (3 kernels) +### 5.6 Decode CPU-pipelining variants (opt1 / opt3 / opt2) -The decode inner loop has a hard intra-tile chain `Q·Kᵀ → softmax → -P·V`: the softmax `tl.max(Sj)` waits for the first GEMM, so a naïve loop -stalls the CPU on `Sj` and **leaves the GEMM engine idle during the -softmax** (a bubble). Three variants trade CPU/HW complexity against that -bubble (assume a realistic non-zero per-op CPU issue cost — §9/ADR-0064): +The three decode variants are shown **in full in the TL;DR** (Kernel 1): +**opt1** current `CompositeCmd` (has a GEMM-engine bubble while the CPU +auto-waits on `Sj`), **opt3** software pipelining (issue the next tile's +`Q·Kᵀ` before this tile's softmax, `Sj` in a persistent double buffer — +removes the bubble, no new command kind), **opt2** `ex_composite` (split +into `#1` Q·Kᵀ = existing composite reading K first, `#2` softmax+P·V+ +accumulator = the only new flash-epilogue machinery, reading V later). -**Option 1 — current `CompositeCmd` (today; has the bubble):** +**Recommend:** ship **opt3** now (no new machinery, removes the bubble); +revisit **opt2** once the cost model (ADR-0064) makes the fewer-CPU-issues +win measurable (§8 item 4). The MATH engine already has max/sum/exp — the +only genuinely new part of opt2 is the composite's **stateful `(m,ℓ,O)` +accumulator**, not the ops. -```python -for j in range(n_tiles): - with tl.scratch_scope(): - Sj = tl.composite("gemm", a=q_g, b=tl.ref(k_tile(j),(d,TILE)), epi=[scale]) - # ↓ CPU auto-waits on Sj → GEMM engine IDLE while softmax runs (bubble) - m2 = tl.maximum(m, tl.max(Sj,-1)); P = tl.exp(Sj-m2); corr = tl.exp(m-m2) - l = l*corr + tl.sum(P,-1) - O = O*corr + tl.composite("gemm", a=P, b=tl.ref(v_tile(j),(TILE,d))); m = m2 -``` - -**Option 3 — software pipelining (current primitives; bubble removed):** -issue the *next* tile's `Q·Kᵀ` **before** this tile's softmax, so the GEMM -engine runs `Q·Kᵀ_{j+1}` during `softmax_j`. `Sj` lives in a persistent -**double buffer** (outside `scratch_scope`, so the next composite does not -clobber it). - -```python -Sb = double_buffer() # 2 persistent Sj buffers -h = tl.composite("gemm", a=q_g, b=tl.ref(k_tile(0),(d,TILE)), out=Sb[0], epi=[scale]) -for j in range(n_tiles): - Sj = Sb[j % 2] - if j+1 < n_tiles: # ← next Q·Kᵀ before softmax: fills GEMM engine - h = tl.composite("gemm", a=q_g, b=tl.ref(k_tile(j+1),(d,TILE)), out=Sb[(j+1)%2], epi=[scale]) - with tl.scratch_scope(): - m2 = tl.maximum(m, tl.max(Sj,-1)); P = tl.exp(Sj-m2); corr = tl.exp(m-m2) - l = l*corr + tl.sum(P,-1) - O = O*corr + tl.composite("gemm", a=P, b=tl.ref(v_tile(j),(TILE,d))); m = m2 -``` - -**Option 2 — extended composite (`ex_composite`; needs a new command kind):** -split into **two** composites so DMA can prioritise **K first, V later** -(V is only needed after softmax). `#1` is the *existing* composite (GEMM + -`scale`); only `#2` (softmax + P·V + the online-softmax accumulator merge) -needs the new flash-epilogue machinery (reduction epilogues + a stateful -`(m,ℓ,O)` accumulator, §8 item 4). The CPU issues both non-blocking and -**never waits intra-tile** → maximal run-ahead, fewest issues. - -```python -for j in range(n_tiles): - Sj = tl.composite("gemm", a=q_g, b=tl.ref(k_tile(j),(d,TILE)), epi=[scale]) # #1: reads K (priority) - tl.ex_composite("softmax_pv", s=Sj, v=tl.ref(v_tile(j),(TILE,d)), acc=(m,l,O), scale=scale) # #2: reads V -``` - -| | new HW cmd | GEMM bubble | CPU intra-tile wait | issues | available now | -|---|---|---|---|---|---| -| **opt1 current** | no | **yes** | yes | `O(tiles·ops)` | ✓ | -| **opt3 sw-pipe** | no | no | yes (reordered) | `O(tiles·ops)` | ✓ | -| **opt2 ex_composite** | **#2 only** | no | **no** | `O(tiles)` | ✗ (build #2) | - -All three end with the §4 2-level reduce. **Recommend:** ship **opt3** now -(no new machinery, removes the bubble); revisit **opt2** once the cost -model (ADR-0064) makes the fewer-issues win measurable. - -> **Prefill note:** these variants are a **decode** concern. In prefill -> (§5.5) the causal `if` (skip-future / partial-mask) is data-dependent -> kernel control flow that **cannot enter a composite** and makes -> pre-issuing speculative; prefill's overlap is the `recv_async` KV -> prefetch, already present. So opt2/opt3 give little for prefill. +These variants are a **decode** concern: in prefill (§5.5) the causal `if` +is kernel control flow that cannot enter a composite, and the ring already +overlaps via `recv_async`. ---