From 8176cdf2871e0309bd7494513714905bc24041f1 Mon Sep 17 00:00:00 2001 From: Yangwook Date: Thu, 4 Jun 2026 21:32:08 -0700 Subject: [PATCH] =?UTF-8?q?gqa(adr):=20split=20ADR-0060=20SP=20into=20two?= =?UTF-8?q?=20kernels=20=E2=80=94=20decode=3Dreduce,=20prefill=3Dring?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Decode-SP and prefill-SP are structurally different and cannot share one kernel (principle: move the smaller thing): - Decode: O=[G,d] tiny, KV cache big -> keep KV statically sharded/resident, G heads replicated (M-fold), move only (m,l,O) via the 2-level reduce (S4). - Prefill: O=[S,d] big -> shard heads (1 query head per CUBE, C=G), rotate KV (Ring KV, S5.5), no (m,l,O) reduce; each CUBE writes its own head. Rewrites TL;DR (two kernels), S0 (head map differs by case), S0.5.4 (output head distribution differs -> downstream out-proj impact), S4 (scoped to decode; S4.1 = intra-CUBE KV-split + PE reduce, the only way decode uses P PEs), S5.1 (decode skeleton), S5.5 (head-parallel Ring KV), S9/S10/S11, and adds SB items (two head mappings, output asymmetry, prefill within-CUBE PE, C=G coupling, reconcile with _attention_mesh_mlo_2d). KO mirror deferred until the design stabilizes (adr-proposed is mirror-exempt). Docs only. Co-Authored-By: Claude Opus 4.8 (1M context) --- .../ADR-0060-algo-gqa-fused-attention-ahbm.md | 332 +++++++++++------- 1 file changed, 203 insertions(+), 129 deletions(-) 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 6675f62..a79b8c8 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 @@ -34,69 +34,70 @@ known algorithms onto the kernbench **greenlet `tl` programming model** --- -## TL;DR — final GQA kernel (composite hybrid) pseudocode +## TL;DR — two SP kernels (decode = reduce, prefill = ring) -The decision (§1) in one place: **GEMMs → `tl.composite`; softmax merge + -2-level reduction → kernel; `tl.load` is lazy.** The KV head is selected by -**CUBE coordinate** (one CUBE Group per head, §0) — **no `for kv` loop**. -`G` is folded into the matmul M dim (real GQA, no broadcast). The KV -sequence is sharded `C × P` ranks (Level-1 inter-CUBE × Level-2 intra-CUBE -PE); the combine is a 2-level reduce-to-root (§4). Reference shape: +Decode-SP and prefill-SP are **structurally different** and are **two +kernels** — one kernel cannot do both. The principle is *move the smaller +thing*: + +- **Decode** (T_q=1): the output `O = [G, d]` is tiny, the KV cache is big. + → keep KV **statically sharded & resident**, do the local sweep, and move + only the small `(m,ℓ,O)` via a **2-level reduce** (§4). `G` heads are + **replicated** (folded into the matmul M dim). No KV movement. +- **Prefill** (T_q=S): the output `O = [S, d]` per head is big — reducing it + would move `[S,d]` per rank. → instead make **each CUBE own one query head** + (head-parallel) and **rotate the KV** so each head sees all KV + (**Ring KV**, §5.5). No `(m,ℓ,O)` reduce (each CUBE outputs a different + head). Only KV blocks move. + +Memory has room, so **replicate Q/weights to avoid communicating them**; +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`). ```python -def gqa_attention(q_ptr, k_ptr, v_ptr, o_ptr, - counter, start_pe, q_block, scale, *, tl): - # ---- geometry: this rank within its CUBE Group (§0, §2) ---- - cube_id = tl.program_id(axis=1) # which CUBE in the SIP mesh - pe_id = tl.program_id(axis=0) # which PE in the CUBE - kv = head_of_cube(cube_id) # CUBE coord → KV head (no for-kv loop) - C, P = cube_group_size, pes_per_cube # Level-1 × Level-2 SP extents - G = H_q // H_kv # query heads per KV head (=8) - causal = q_block.is_prefill +# ── 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): + 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; all G heads replicated (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) + with tl.scratch_scope(): + Sj = tl.composite("gemm", a=q_g, b=tl.ref(k_tile(j), (d, TILE))) * scale + 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 - # Q group: G rows folded into M → [G·T_q, d]. Lazy load (ADR-0062). - q_g = tl.load(q_base(kv), (G * q_block.T_q, d)) # [G·T_q, d] - my_len = valid_len_2level(counter, start_pe, cube_id, pe_id, C, P) \ - if not causal else S_kv_local - 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 - - # ---- 2-level combine to the CUBE Group root (§4); data-driven, no barrier ---- - hierarchical_reduce_and_store(m, l, O, cube_id, pe_id, C, P, o_base(kv)) +# ── Kernel 2 — PREFILL + SP: 1 Q head per CUBE (head-parallel), Ring KV (NO reduce) ── +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) + 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 ── + 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 + 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 ``` -> **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 reduction stay in the kernel because the existing -> `CompositeCmd` cannot carry cross-tile state. The reduction is **2-level** -> (Level-2 PE tree within a CUBE → Level-1 CUBE-mesh reduce within the CUBE -> Group), both reduce-to-root, log-sum-exp, sent as each rank finishes its -> local sweep (§4). `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). +> **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 +> 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 head per CUBE (§0.5.4). --- @@ -195,12 +196,19 @@ two levels** of sequence-parallelism (SP): - **Level-2 (intra-CUBE, across PEs):** each CUBE's slice is split again across its `P` PEs. -So the KV-parallel reduction group for one (KV head, request) is **`C × P` -ranks**, all within one SIP. The `G` query heads are folded into the -matmul M dim (replicated across all ranks). `C` is a **tuning knob** -(e.g. 8 or 4): large `C` for long-context prefill, small `C` (→ 1 = a -single CUBE, intra-CUBE SP only) for short decode where reduction -dominates. +So one KV head maps to **`C × P` ranks**, all within one SIP. **How the +`G` query heads map onto those ranks differs by case** (the two kernels, +TL;DR / §5): + +- **Decode** (§4): `G` heads are **replicated** (folded into the matmul M + dim) on every rank; KV is sequence-sharded `C × P` ways; outputs reduce. +- **Prefill** (§5.5): `G` heads are **head-parallel** — with `C = G`, CUBE + `i` owns query head `i`; KV rotates (Ring); no reduce. + +`C` is a **tuning knob** (e.g. 8 or 4): for prefill set `C = G` (one head +per CUBE); for decode `C` trades inter-CUBE reduction against KV-parallel +breadth (small `C`, even `C = 1` single-CUBE, for short context where +reduction dominates). **Topology grounding** (`topology.yaml`): a SIP is a `4×4` CUBE mesh (16 CUBEs, ADR-0017 NOC); a CUBE has `P = 8` PEs @@ -282,14 +290,19 @@ V_block` via IPCQ into ping-pong buffers (post-RoPE), plus ### 0.5.4 OUTPUTS -| Output | Shape | Location | Notes | -|---|---|---|---| -| `O` (final) | `[G, T_q, d]` | `O_base`, **CUBE-Group root only** | normalised `O_acc / ℓ_acc`, cast to storage dtype. | +**The output head distribution differs by kernel** — downstream +out-projection must consume them accordingly: -**Intermediate (non-root rank, on IPCQ — not a kernel-visible output):** -`(m_i, ℓ_i, O_i)` (`m,ℓ`: `[G, T_q]`; `O_i`: `[G, T_q, d]` unnormalised) -pushed up the 2-level tree to the parent (Level-2 PE parent, then Level-1 -CUBE parent) via `tl.send` (§4). `O_i` is the heavy part. +| Kernel | Output | Location | Notes | +|---|---|---|---| +| **Decode** (§4 reduce) | `O = [G, 1, d]` (all heads) | `O_base` at the **CUBE-Group root** | head-replicated → all `G` heads end on one rank after the 2-level reduce. | +| **Prefill** (§5.5 ring) | `O_i = [1, T_q, d]` (one head each) | per-CUBE `O_base[i]`, **distributed** | head-parallel → CUBE `i` writes head `i`'s rows in place; no reduce. | + +**Decode intermediate** (non-root rank, on IPCQ — not kernel-visible): +`(m_i, ℓ_i, O_i)` pushed up the 2-level tree to the parent (Level-2 PE +parent, then Level-1 CUBE parent) via `tl.send` (§4). `O_i` is the heavy +part. **Prefill** has no such partials (no reduce); instead KV blocks +circulate (§5.5). **No-SP (`C=P=1`):** no IPCQ partials; the single rank's running `(m,ℓ,O)` is normalised in place and written to `O_base`. @@ -487,6 +500,11 @@ Notes: ## 4. Reduction (KV-parallel / SP combine) — 2-level reduce-to-root +**Scope: this is the DECODE-SP path** (head-replicated, KV statically +sharded, small `O`). Prefill-SP does **not** reduce — it rotates KV (Ring, +§5.5). Reduce is chosen for decode because `O = [G, d]` is tiny, so moving +`(m,ℓ,O)` beats moving the resident KV. + After its tile sweep each rank holds `(m_i, ℓ_i, O_i)` (unnormalised) over its `1/(C·P)` shard. Combine via the associative/commutative log-sum-exp merge (identical math to the baseline's fold, `_attention_mesh_mlo.py:117-122`): @@ -505,10 +523,13 @@ Because the merge is associative **and** commutative, the combine is ### 4.1 Level-2 — intra-CUBE (P PEs → CUBE root PE) -- A **reduce-to-root tree** over the `P` PEs of a CUBE, depth `⌈log₂ P⌉` - (3 for `P=8`), over the PE IPCQ `N/S/E/W` mesh - (`configure_sfr_intracube_pe_ring`). Each tree pair is a 1-hop physical - neighbour. Result lands on the CUBE's root PE. +- The CUBE's KV slice is itself **sequence-SP-split across its `P` PEs** + (decision (a): the only way decode, `T_q=1`, can use all `P` PEs — there + is no query axis to split). Each PE sweeps its `1/(C·P)` sub-shard, then a + **reduce-to-root tree** over the `P` PEs, depth `⌈log₂ P⌉` (3 for `P=8`), + over the PE IPCQ `N/S/E/W` mesh (`configure_sfr_intracube_pe_ring`). Each + tree pair is a 1-hop physical neighbour. Result lands on the CUBE's root + PE. ### 4.2 Level-1 — intra-CUBE-Group (C CUBE roots → Group root) @@ -542,12 +563,13 @@ Because the merge is associative **and** commutative, the combine is Each level's collective is **selectable** (a tuning item, §9): the defaults above (Level-2 tree, Level-1 center-root mesh) suit the latency-bound, -small-`O` decode case; a **ring** option (chunked `O`) suits bandwidth-bound -long-context prefill with large `T_q·d` payload. `C=1` degenerates to -Level-2 only (single-CUBE SP). +small-`O` decode case; a **ring** variant (chunked) is available if a +specific reduction proves bandwidth-bound. (Note: this is the *reduce* +collective for **decode**; **prefill** does not reduce at all — it uses the +Ring-KV kernel of §5.5.) `C=1` degenerates to Level-2 only (single-CUBE SP). -**Payload:** `O_i` (`[G·T_q, d]`) is heavy; `m,ℓ` are light. Sent as -separate `tl.send`s (baseline pattern). +**Payload:** `O_i` (`[G·1, d]` for decode) is the heavy part; `m,ℓ` are +light. Sent as separate `tl.send`s (baseline pattern). Kernel structure (greenlet), both levels reuse `merge`: @@ -576,31 +598,33 @@ is required** (§6). ## 5. Per-case kernels -All four cases share §3 (tile sweep) + §4 (reduction); they differ only -in the SP extents `C·P`, query-block width, and which `tile_*` predicates -fire. +All cases share §3's composite-hybrid inner tile, but split into **two +kernels** (TL;DR): the **decode-reduce** skeleton below (§5.2/§5.3) and the +**prefill-ring** kernel (§5.5). They differ in head mapping, KV strategy, +and cross-rank communication — see §10. -### 5.1 Common skeleton +### 5.1 Decode skeleton (head-replicated, static shard, reduce) ```python -def attn_kernel(q_ptr, k_ptr, v_ptr, o_ptr, counter, start_pe, start_cube, - C, P, q_block, softmax_scale, *, tl): +def gqa_decode_sp(q_ptr, k_ptr, v_ptr, o_ptr, counter, start_pe, start_cube, + C, P, *, tl): cube_id = tl.program_id(axis=1); pe_id = tl.program_id(axis=0) - kv = head_of_cube(cube_id) # CUBE coord → KV head (no for-kv loop) + kv = head_of_group(cube_id) # CUBE coord → KV head (no for-kv loop) my_len = valid_len_2level(counter, start_cube, start_pe, cube_id, pe_id, C, P) n_tiles = ceil(my_len / TILE) - q_g = load_Q_group(q_ptr, kv) # [G·T_q, d]; lazy tl.load, G folded into M (no broadcast) + q_g = load_Q_group(q_ptr, kv) # [G·1, d]; lazy tl.load, G folded into M (replicated) m, l, O = init_running() # persistent arena: -inf, 0, zeros - 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: 2 composites + softmax MATH, in tl.scratch_scope + for j in range(n_tiles): # sweep ONLY my static KV shard (resident, no move) + run_tile(j) # §3: 2 composites + softmax MATH, in tl.scratch_scope if C == 1 and P == 1: tl.store(o_base(kv), O / l) # no reduction (single rank) else: - hierarchical_reduce_and_store(m, l, O, cube_id, pe_id, C, P, o_base(kv)) # §4 + hierarchical_reduce_and_store(m, l, O, cube_id, pe_id, C, P, o_base(kv)) # §4 reduce ``` +(Prefill-no-SP is this same shape with `C=P=1`, `T_q>1`, and causal masking; +**prefill-with-SP is the separate Ring kernel of §5.5**, not this skeleton.) + ### 5.2 DECODE, no SP (`C=P=1`, one PE owns the head's KV) - `T_q = 1`, attends to all past KV ⇒ no future tiles, mask only on the @@ -644,25 +668,36 @@ def attn_kernel(q_ptr, k_ptr, v_ptr, o_ptr, counter, start_pe, start_cube, the tile op sequence adds it (`Sj + mask_j`). - GQA batches the group's query heads along the same axis as decode. -### 5.5 PREFILL, with SP (Ring KV) +### 5.5 PREFILL, with SP (Ring KV) — head-parallel, NO reduce -- KV sharded across the `C·P` ranks in **ring order**; each step delivers a - peer's KV block over IPCQ. Running `(m,ℓ,O)` is carried **across ring - steps** (Python handles, exactly as `_attention_mesh_kv` does today). The - ring is the long-context, bandwidth-bound case where §4.4's **ring** - collective option (not the tree) is the natural choice. -- IPCQ overlaps next step's KV receive with current step's compute via - `tl.recv_async`/`tl.wait` (`tl_context.py:543-560` — already exists). -- **Causal ring skip:** if the incoming step's KV block is entirely - *after* the local query block, skip its compute (don't recv-consume / - don't fold) — eliminates ≈half the steps for causal attention. -- Receive buffers ping-pong (persistent arena, not recycled). -- After the ring, `(m,ℓ,O)` is final per rank; remaining KV-parallel splits - combine via the §4 2-level reduce, then normalise + store at the Group - root. +This is the **second kernel** (Kernel 2 in the TL;DR), structurally +distinct from the decode reduce path (§4). The output `O = [T_q, d]` per +head is **big**, so reducing it across ranks would move `[T_q,d]` per rank; +instead we **shard the heads** and **move the (also big) KV**, which each +head needs in full anyway. + +- **Head-parallel placement:** within a CUBE Group, **CUBE `i` owns query + head `i`** (the `G` heads → the `C=G` CUBEs, one per CUBE) and KV slice + `i`. Each CUBE computes **its one head's** full attention. Because each + CUBE produces a *different* head, there is **no `(m,ℓ,O)` reduce** — each + CUBE normalises and writes its own head's rows. +- **Ring KV:** the `C` KV slices **rotate** around the CUBE ring; each CUBE + folds the incoming block into its head's running `(m,ℓ,O)` (online-softmax, + carried across ring steps as Python handles, as `_attention_mesh_kv` does + today). After `C` steps every head has seen all KV. **GQA reuse comes from + the rotation** — slice `j` visits all `C` CUBEs, serving all `G` heads. +- IPCQ overlaps the next step's KV receive with the current step's compute + via `tl.recv_async`/`tl.wait` (`tl_context.py:543-560` — already exists); + receive buffers ping-pong (persistent arena, not recycled). +- **Causal ring skip:** if the incoming KV block is entirely *after* this + head's query block, skip its compute — eliminates ≈half the steps. +- **Within a CUBE (P PEs):** the head's query rows `[T_q, d]` and/or the + current KV block tile across the `P` PEs (a query-axis split exists here, + unlike decode); details in §B. The baseline `_attention_mesh_kv` already implements the ring fold; this -ADR adds GQA reuse, causal step-skip, and `recv_async` overlap. +ADR adds GQA reuse, the head-parallel placement, causal step-skip, and the +composite-hybrid inner tile (§3). --- @@ -769,11 +804,11 @@ new primitive** — only the kernel restructuring in §5.2 (per KV head, ahead the kernel issues non-blocking composites before waiting, and the scheduler's per-tile streaming depth. 2. **Reduction topology per level (§4.4)** — Level-2 (intra-CUBE PE) and - Level-1 (intra-CUBE-Group CUBE-mesh) each selectable tree/ring/center-mesh; - sweep tree (latency, small-`O` decode) vs ring (bandwidth, large-`O` - prefill). Also the CUBE-Group size `C` knob (small for decode, large for - long context). Ensure each tree pair is a 1-hop physical neighbour; verify - against the SFR install. + Level-1 (intra-CUBE-Group CUBE-mesh) each selectable tree/center-mesh (and + a ring variant) — **decode reduce only**; prefill uses the §5.5 Ring-KV + kernel, not a reduce. Also the decode `C` knob (small `C` for short + context where reduction dominates). Ensure each tree pair is a 1-hop + physical neighbour; verify against the SFR install. 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`). @@ -791,21 +826,21 @@ new primitive** — only the kernel restructuring in §5.2 (per KV head, ## 10. Coverage summary -| Case | KV placement | Inner loop | Reduction | Masking | +| Case | Head map | KV strategy | Cross-rank comm | Masking | |---|---|---|---|---| -| Decode, no SP | 1 rank (`C=P=1`), all KV | one-shot or tiled, GQA-batched | none | last tile only | -| Decode, SP | 2-level RR over `C·P` ranks | tiled, GQA-batched | §4 2-level (`⌈log₂P⌉`+mesh over `C`) | last tile only | -| Prefill, no SP | resident | tiled per q-block | none | skip future / triangular boundary | -| Prefill, SP (Ring) | ring over `C·P` ranks | per ring step (recv_async overlap) | running state + §4 2-level | causal step-skip + boundary | +| Decode, no SP | `G` replicated, 1 rank | all KV resident | none | last tile only | +| **Decode, SP** | `G` replicated (M-fold) | 2-level static shard `C·P` | **§4 2-level reduce** (small `O`) | last tile only | +| Prefill, no SP | `G` replicated, 1 rank | resident | none | triangular / skip future | +| **Prefill, SP (Ring)** | **1 head per CUBE** (`C=G`) | **Ring KV rotate** | **none** (KV blocks move, not `O`) | causal step-skip + boundary | **I/O per case** (full contract §0.5): -| Case | Inputs | Output | IPCQ partials | +| Case | Inputs | Output | Cross-rank traffic | |---|---|---|---| | Decode, no SP | `Q[G,1,d]`, full `K/V[S,d]` on 1 rank | `O[G,1,d]` at `O_base` | none | -| Decode, SP | `Q[G,1,d]`, per-rank `K/V[S/(C·P),d]` | `O[G,1,d]` (Group root) | `(m,ℓ,O_i)` per rank → 2-level | +| Decode, SP | `Q[G,1,d]` (replicated), per-rank `K/V[S/(C·P),d]` | `O[G,1,d]` at Group root | `(m,ℓ,O_i)` reduce (small) | | Prefill, no SP | `Q[G,T_q,d]`, `K/V[≤end,d]` | `O[G,T_q,d]` at `O_base` | none | -| Prefill, SP (Ring) | `Q[G,T_q,d]`, ring-delivered `K/V` blocks | `O[G,T_q,d]` (Group root) | running `(m,ℓ,O)` + 2-level | +| Prefill, SP (Ring) | `Q[1,T_q,d]` per CUBE, own `K/V[S/C,d]` | `O[1,T_q,d]` per CUBE (distributed) | KV blocks rotate (Ring) | In all cases: KV-cache writes and RoPE happen **upstream**; output projection **downstream**; score `S` and probs `P` are never materialised. @@ -835,14 +870,19 @@ secondary check. 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 GEMM work scales with `G`. Asserts the lever actually fires. -3. **2-level reduction step count:** decode-SP issues `⌈log₂ P⌉` (Level-2, - intra-CUBE) + center-mesh-depth-over-`C` (Level-1) reduction rounds — not - the baseline's `C·P − 1` all-to-all; op_log `ipcq_send`/`recv` counts - match the static 2-level tree/mesh, and the result lands at exactly one - rank (CUBE-Group root). A separate assertion checks Level-1 sends use the - CUBE NOC and Level-2 the PE IPCQ. -4. **Causal skip:** prefill skips all `tile_all_future` tiles — GEMM - count equals the lower-triangular tile count, not the full grid. +3. **SP cross-rank traffic — both kernels:** + - *Decode (reduce):* issues `⌈log₂ P⌉` (Level-2, intra-CUBE) + + center-mesh-depth-over-`C` (Level-1) reduction rounds — not the + baseline's `C·P − 1` all-to-all; op_log `ipcq_send`/`recv` match the + static 2-level tree/mesh; result lands at exactly one rank (CUBE-Group + root); Level-1 sends use the CUBE NOC, Level-2 the PE IPCQ. + - *Prefill (ring):* **no** `(m,ℓ,O)` reduce — instead `C` KV-block + rotations (`ipcq_send`/`recv` of K,V per step); each CUBE writes a + **distinct** head's `O` in place (no Group root); GQA reuse shows as + each KV slice consumed by all `C` heads over the ring. +4. **Causal skip:** prefill skips all `tile_all_future` tiles (and whole + future KV blocks in the ring) — GEMM count equals the lower-triangular + tile/step 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. **Load/compute overlap:** end-to-end latency of the tiled sweep is @@ -957,3 +997,37 @@ predicted default; revise on review. sequence with no gaps/overlaps and keep each rank's local buffer dense. **Recommend:** unit-test the index math (every global token maps to exactly one rank; per-rank slices are contiguous) before the kernel. + +### Items from the decode-reduce / prefill-ring split (this revision) + +1. **Two kernels, two head mappings.** Decode-SP = head-replicated + static + KV shard + 2-level reduce (§4); Prefill-SP = head-parallel (1 head/CUBE, + `C=G`) + Ring KV + no reduce (§5.5). The principle is *move the smaller + thing* — decode's `O` is tiny (reduce it), prefill's `O` is big (move KV + instead). **Recommend:** keep them as two kernels; do not re-merge. + +2. **Output head distribution differs (downstream impact).** Decode lands + all `G` heads at the CUBE-Group root; prefill leaves one head per CUBE + (distributed). The downstream **out-projection** must consume each layout + (gather for decode-root vs in-place per-CUBE for prefill). **Recommend:** + pin the O layout per kernel in the `qkv_rope`/`out_proj` contract before + implementation (§0.5.4). + +3. **Prefill within-CUBE `P` PEs.** §5.5 splits the head's query rows + `[T_q, d]` and/or the current KV block across the `P` PEs (a query axis + exists, unlike decode). **Recommend:** default to tiling the query rows + across PEs (no intra-CUBE reduce needed — disjoint output rows); fall back + to KV-block split + intra-CUBE reduce only if `T_q < P`. + +4. **`C = G` coupling for prefill.** The head-parallel mapping assumes + `C = G = 8` (one head per CUBE). If `C ≠ G`, the mapping needs revisiting + (multiple heads per CUBE, or heads spanning a partial ring). **Recommend:** + fix `C = G` for the prefill kernel at headline scale; treat `C ≠ G` as a + separate study. + +5. **Reconcile with `_attention_mesh_mlo_2d` (current impl).** The remote + impl's 2D kernel is an **AllReduce** over cubes with **Q replicated** — + i.e. the decode-reduce family, but all-reduce (broadcast-back) not + reduce-to-root, and not yet the prefill head-parallel ring. **Recommend:** + (a) move its 2D AllReduce → reduce-to-root (drop broadcast-back) for the + decode kernel; (b) add the §5.5 head-parallel Ring-KV kernel for prefill.