gqa(adr): ADR-0060 hierarchical CUBE-Group SP — 2-level reduce-to-root
Placement pivot (user-approved): a CUBE Group = C CUBEs within one SIP that jointly own one KV head + its G=8 query heads. KV sequence sharded 2 levels (Level-1 inter-CUBE over C, Level-2 intra-CUBE PE over P=8 = C*P ranks); G folded into matmul M. C is a knob (8/4; C=1 = single-CUBE). Reverses the old 'a query head never spans CUBEs' non-goal (held only because the baseline was H_kv=1). device=SIP; the for-kv loop is gone (head picked by CUBE coord). Reduction is a 2-level reduce-to-root (not all-reduce): Level-2 PE tree -> Level-1 center-root CUBE-mesh reduce, adapting lrab_hierarchical_allreduce's inter-CUBE pattern as reduce-only + log-sum-exp. Data-driven (send on local P.V completion, no global barrier) + level-pipelined; per-level topology configurable (tree for decode, ring for long prefill). Rewrites SS0/2/4/5/0.5/8-11, pseudocode, and adds SSB items (4-SIP config, mesh partition, C knob, invariant-reversal check, index-math test). KO mirror updated. Topology grounded in topology.yaml (4x4 CUBE mesh/SIP, 8 PE/cube). Docs only; no production code changed. Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -37,63 +37,66 @@ 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.
|
||||
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:
|
||||
|
||||
```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)
|
||||
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
|
||||
|
||||
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)
|
||||
# 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)
|
||||
# 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
|
||||
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
|
||||
# ---- 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))
|
||||
```
|
||||
|
||||
> **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).
|
||||
> 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).
|
||||
|
||||
---
|
||||
|
||||
@@ -136,8 +139,9 @@ GEMMs onto the scheduler-managed `tl.composite` path** (see §1) — this
|
||||
(`h_q = G·h_kv`) runs with **no new primitive** — see §8.
|
||||
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`
|
||||
steps). Attention only needs `O` at the query owner → a **tree
|
||||
reduction** to root is `⌈log₂ N⌉` steps (§4).
|
||||
steps). Attention only needs `O` at the query owner → a **2-level
|
||||
reduce-to-root** (intra-CUBE tree + intra-CUBE-Group center-mesh, §4) is
|
||||
`⌈log₂ P⌉` + center-mesh-over-`C` steps.
|
||||
3. **Validation scale only.** `S = 16` because the 1 MiB scratch bump
|
||||
allocator leaks per-tile temporaries
|
||||
(`test_milestone_gqa_llama70b.py:123-148`) and there is no causal
|
||||
@@ -180,20 +184,48 @@ Hardware recap:
|
||||
**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.
|
||||
|
||||
**Two orthogonal mapping layers** (round-robin placement):
|
||||
- **Layer 1 (KV-parallel, intra-request):** KV token `i` lands on PE
|
||||
`(start_pe + i) mod N` — round-robin so one long request is split
|
||||
across `N` PEs, balanced to ≤1 token.
|
||||
- **Layer 2 (inter-request):** `start_pe = request_id mod N` rotates the
|
||||
"PE-0 role" per request.
|
||||
- Combined: `pe = (request_id + global_token_idx) mod N`.
|
||||
### CUBE Group — the placement unit (hierarchical SP)
|
||||
|
||||
`N` = number of PEs in the KV-parallel reduction group for one
|
||||
(query head, request). Reduction stays **intra-CUBE** (a query head never
|
||||
spans CUBEs — explicit non-goal).
|
||||
**A `CUBE Group` is `C` CUBEs (within one SIP) that jointly own one KV
|
||||
head and its `G` query heads.** One KV head's KV sequence is **sharded
|
||||
two levels** of sequence-parallelism (SP):
|
||||
- **Level-1 (inter-CUBE, within the CUBE Group):** the head's sequence is
|
||||
split across the `C` CUBEs of the group.
|
||||
- **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.
|
||||
|
||||
**Topology grounding** (`topology.yaml`): a SIP is a `4×4` CUBE mesh
|
||||
(16 CUBEs, ADR-0017 NOC); a CUBE has `P = 8` PEs
|
||||
(`hbm_pseudo_channels/hbm_channels_per_pe = 64/8`). With `C = 8` a SIP
|
||||
holds **2 CUBE Groups = 2 KV heads**; the full `H_kv = 8` model therefore
|
||||
spans **4 SIPs** (8/2). Each CUBE Group is **intra-SIP**, so one head's
|
||||
reduction uses the CUBE NOC (Level-1) + PE IPCQ (Level-2) — never UCIe.
|
||||
(The shipped `topology.yaml` has `sips: 2`; full scale needs a 4-SIP
|
||||
config — §B.)
|
||||
|
||||
**`--device` enumerates SIPs** (CLI semantics): one SIP-device bench
|
||||
drives its 2 CUBE Groups (2 KV heads) — the head is selected spatially by
|
||||
CUBE coordinate, **not** an in-kernel `for kv` loop. The CLI runs the
|
||||
4 SIP-devices logically in parallel to cover all 8 heads.
|
||||
|
||||
**Token → rank placement** within a CUBE Group (round-robin SP, balanced
|
||||
to ≤1 token): KV token `i` lands on Level-1 rank `(start + i) mod C`, then
|
||||
Level-2 rank `((i // C) + start_pe) mod P` — a hierarchical round-robin so
|
||||
each of the `C × P` ranks owns a dense, contiguous local slice.
|
||||
|
||||
Reduction is **hierarchical and intra-SIP**: a KV head **does** span the
|
||||
`C` CUBEs of its group (Level-1), reduced over the CUBE NOC; this
|
||||
**reverses** the earlier "a query head never spans CUBEs" non-goal, which
|
||||
existed only because the `H_kv=1` baseline never needed it. Crossing
|
||||
**SIPs** for one head remains a non-goal (a head stays in one SIP).
|
||||
|
||||
---
|
||||
|
||||
@@ -220,27 +252,28 @@ spans CUBEs — explicit non-goal).
|
||||
launches.** ⇒ this kernel is **pure read** on the KV cache.
|
||||
- **P4.** V is not rotated; V-cache holds raw projected V.
|
||||
- **P5.** RoPE position upstream is the token's **absolute global
|
||||
position** `global_idx = local_slot·N + ((pe_id − start_pe) mod N)`
|
||||
(§2.1). Round-robin placement ≠ RoPE position.
|
||||
position** `global_idx = local_slot·(C·P) + rank` where `rank =
|
||||
cube_local·P + pe_local` (§2.1). Round-robin placement ≠ RoPE position.
|
||||
|
||||
### 0.5.2 Shape symbols
|
||||
`T_q` = query length this launch (decode: 1; prefill/chunk: chunk width).
|
||||
`S` = total context length. `S_pe` = keys owned by this PE ≈ `⌈S/N⌉`.
|
||||
Storage `bf16` (numpy proxy `f16`, `memory_store.py:16`); `m,ℓ,O`
|
||||
accumulators `f32` where precision matters.
|
||||
`S` = total context length. `R = C·P` = ranks per CUBE Group;
|
||||
`S_rank` = keys owned by this rank ≈ `⌈S/R⌉`. Storage `bf16` (numpy proxy
|
||||
`f16`, `memory_store.py:16`); `m,ℓ,O` accumulators `f32` where precision
|
||||
matters.
|
||||
|
||||
### 0.5.3 INPUTS (per kernel launch)
|
||||
|
||||
| Input | Shape (per KV head) | Location | Notes |
|
||||
|---|---|---|---|
|
||||
| `Q` | `[G, T_q, d]` | per-PE HBM (loaded to TCM) | post-RoPE (P1). `G` query rows of the group batched. |
|
||||
| `K_cache` | `[S_pe, d]` | per-PE HBM, base `K_base[pe]`, contiguous | post-RoPE (P2). Read-only. Dense locally, strided in global index (§2.1). |
|
||||
| `V_cache` | `[S_pe, d]` | per-PE HBM, base `V_base[pe]` | raw V (P4). Read-only. |
|
||||
| `global_token_counter` | scalar | launch arg | kernel derives `S_pe`, slot↔global, causal bounds. |
|
||||
| `start_pe` (= `request_id mod N`) | scalar | launch arg | Layer-2 rotation. |
|
||||
| `pe_id`, `N` | scalars | launch / `tl.program_id` | reduction-group geometry. |
|
||||
| `Q` | `[G, T_q, d]` | per-rank HBM (loaded to TCM) | post-RoPE (P1). `G` query rows of the group batched. |
|
||||
| `K_cache` | `[S/(C·P), d]` | per-rank HBM, base `K_base[rank]`, contiguous | post-RoPE (P2). Read-only. Dense locally, strided by `C·P` in global index (§2.1). |
|
||||
| `V_cache` | `[S/(C·P), d]` | per-rank HBM, base `V_base[rank]` | raw V (P4). Read-only. |
|
||||
| `global_token_counter` | scalar | launch arg | kernel derives local len, slot↔global, causal bounds. |
|
||||
| `start_cube`, `start_pe` (= `f(request_id)`) | scalars | launch arg | Level-1/Level-2 rotation. |
|
||||
| `cube_id`, `pe_id`, `C`, `P` | scalars | launch / `tl.program_id` 1/0 | CUBE-Group reduction geometry (`R = C·P`). |
|
||||
| `q_block_meta` `{q_start, T_q}` | launch arg | prefill/SP causal masking & skip. |
|
||||
| `O_base` | address | launch arg | where final O is written (root PE only). |
|
||||
| `O_base` | address | launch arg | where final O is written (CUBE-Group root only). |
|
||||
| `softmax_scale` | scalar | launch arg | `1/√d`. |
|
||||
|
||||
For **Ring Attention (§5.5)** add per ring step: incoming `K_block,
|
||||
@@ -251,14 +284,15 @@ V_block` via IPCQ into ping-pong buffers (post-RoPE), plus
|
||||
|
||||
| Output | Shape | Location | Notes |
|
||||
|---|---|---|---|
|
||||
| `O` (final) | `[G, T_q, d]` | `O_base`, **root PE only** | normalised `O_acc / ℓ_acc`, cast to storage dtype. |
|
||||
| `O` (final) | `[G, T_q, d]` | `O_base`, **CUBE-Group root only** | normalised `O_acc / ℓ_acc`, cast to storage dtype. |
|
||||
|
||||
**Intermediate (non-root PE, on IPCQ — not a kernel-visible output):**
|
||||
**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 to `tree_parent` via `tl.send` (§4). `O_i` is the heavy part.
|
||||
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.
|
||||
|
||||
**No-SP (`N=1`):** no IPCQ partials; the single PE's running `(m,ℓ,O)` is
|
||||
normalised in place and written to `O_base`.
|
||||
**No-SP (`C=P=1`):** no IPCQ partials; the single rank's running `(m,ℓ,O)`
|
||||
is normalised in place and written to `O_base`.
|
||||
|
||||
**Explicitly NOT outputs:** KV-cache writes (upstream, P3), output
|
||||
projection (downstream), score `S` and probs `P` (never materialised).
|
||||
@@ -278,7 +312,7 @@ latency model:
|
||||
- **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
|
||||
per-rank 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
|
||||
@@ -348,16 +382,24 @@ permits (`kernel_runner.py`, ADR-0020 D3).
|
||||
|
||||
## 2. Memory layout & driver responsibilities
|
||||
|
||||
### 2.1 KV cache allocation
|
||||
- Per-PE KV buffers sized to per-PE max context `⌈max_context / N⌉ × d ×
|
||||
dtype`, for K and V separately. In kernbench these are deployed with a
|
||||
`DPPolicy` whose `pe` (or `cube`) axis is `row_wise` over the sequence
|
||||
dimension (`policy/placement/dp.py`; the baseline uses
|
||||
`DPPolicy(pe="row_wise")` for K/V and `replicate` for Q).
|
||||
- Within a PE, assigned tokens are **appended contiguously** (slot
|
||||
0,1,2,…). Global indices are strided (`i, i+N, …`) but the per-PE
|
||||
buffer is dense ⇒ DMA stays contiguous.
|
||||
- Slot → global: `global_idx = local_slot·N + ((pe_id − start_pe) mod N)`.
|
||||
### 2.1 KV cache allocation (2-level SP)
|
||||
|
||||
One KV head's sequence is sharded across `C × P` ranks (Level-1 inter-CUBE
|
||||
× Level-2 intra-CUBE PE, §0). In kernbench this is a `DPPolicy` whose
|
||||
**both** the `cube` axis (`row_wise` over `C`) **and** the `pe` axis
|
||||
(`row_wise` over `P`) shard the sequence dimension, with Q `replicate`
|
||||
(`policy/placement/dp.py`). The baseline's single-axis
|
||||
`DPPolicy(pe="row_wise")` becomes a two-axis `row_wise` over (cube, pe).
|
||||
|
||||
- Per-rank KV buffers sized to `⌈max_context / (C·P)⌉ × d × dtype`, K and V
|
||||
separately.
|
||||
- Within a rank, assigned tokens are **appended contiguously** (slot
|
||||
0,1,2,…); the per-rank buffer is dense ⇒ DMA stays contiguous, even
|
||||
though global indices are strided by `C·P`.
|
||||
- Slot → global (hierarchical round-robin): `rank = cube_local·P + pe_local`
|
||||
where `cube_local = (cube_id − start_cube) mod C`,
|
||||
`pe_local = (pe_id − start_pe) mod P`; then
|
||||
`global_idx = local_slot·(C·P) + rank`.
|
||||
|
||||
### 2.2 Driver per-launch duties (minimal)
|
||||
The driver supplies, per launch, the bases + counter + rotation; the
|
||||
@@ -365,23 +407,26 @@ kernel derives the rest:
|
||||
|
||||
| Launch arg | Purpose |
|
||||
|---|---|
|
||||
| `K_base[pe]`, `V_base[pe]` | per-PE KV buffer bases (from tensor VA) |
|
||||
| `K_base[rank]`, `V_base[rank]` | per-rank KV buffer bases (from tensor VA) |
|
||||
| `O_base` | attention output destination |
|
||||
| `global_token_counter` | current sequence position |
|
||||
| `start_pe = request_id mod N` | Layer-2 rotation |
|
||||
| `pe_id` (`tl.program_id`), `N` | geometry |
|
||||
| `start_cube`, `start_pe = f(request_id)` | Level-1/Level-2 rotation |
|
||||
| `cube_id`, `pe_id` (`tl.program_id` 1/0), `C`, `P` | CUBE-Group geometry |
|
||||
| `q_block_meta` | prefill/SP query block start + length |
|
||||
|
||||
Kernel-derived (plain arithmetic):
|
||||
- **my turn to write this step:** `(start_pe + counter) mod N == pe_id`
|
||||
- **my valid length:** `base = counter // N; rem = counter % N;
|
||||
my_len = base + (1 if ((pe_id − start_pe) mod N) < rem else 0)`
|
||||
Let `R = C·P`, `cube_local = (cube_id − start_cube) mod C`,
|
||||
`pe_local = (pe_id − start_pe) mod P`, and this rank's index
|
||||
`rank = cube_local·P + pe_local`. Kernel-derived (plain arithmetic):
|
||||
- **my turn to write this step:** `(start_rank + counter) mod R == rank`
|
||||
- **my valid length:** `base = counter // R; rem = counter % R;
|
||||
my_len = base + (1 if rank < rem else 0)`
|
||||
- **read range:** tiles `0 .. ⌈my_len / TILE⌉`.
|
||||
- **causal bounds / per-tile skip:** from global positions of the query
|
||||
block vs each tile.
|
||||
|
||||
Driver = bases + counter + rotation. The single formula
|
||||
`(request_id + token_idx) mod N` is the entire placement policy.
|
||||
`(request_id + token_idx) mod (C·P)` over the two SP axes is the entire
|
||||
placement policy.
|
||||
|
||||
---
|
||||
|
||||
@@ -440,84 +485,124 @@ Notes:
|
||||
|
||||
---
|
||||
|
||||
## 4. Reduction (KV-parallel / SP combine) — tree to root
|
||||
## 4. Reduction (KV-parallel / SP combine) — 2-level reduce-to-root
|
||||
|
||||
After its tile sweep each PE holds `(m_i, ℓ_i, O_i)` (unnormalised).
|
||||
Combine via the associative/commutative log-sum-exp merge (identical math
|
||||
to the baseline's fold, `_attention_mesh_mlo.py:117-122`):
|
||||
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`):
|
||||
|
||||
```python
|
||||
def merge(m_a, l_a, O_a, m_b, l_b, O_b):
|
||||
m = tl.maximum(m_a, m_b)
|
||||
sa, sb = tl.exp(m_a - m), tl.exp(m_b - m)
|
||||
return m, l_a*sa + l_b*sb, O_a*sa + O_b*sb
|
||||
# final: O = O_root / l_root # normalise once at the tree root
|
||||
# final: O = O_root / l_root # normalise once at the CUBE-Group root
|
||||
```
|
||||
|
||||
- **Timing:** exchange is **after P·V** (each PE has its final `O_i`).
|
||||
- **Topology:** mesh tree, depth `⌈log₂ N⌉`, over physical neighbours.
|
||||
For `N=8`: level-0 `(0↔1,2↔3,4↔5,6↔7)`, level-1 `(1↔3,5↔7)`, level-2
|
||||
`(3↔7)`, root = `7`. Each tree pair must be a physical `N/S/E/W`
|
||||
neighbour so `tl.send(dir)`/`tl.recv(dir)` use a real link
|
||||
(tuning item §9; the SFR install that wires neighbours is
|
||||
`configure_sfr_intracube_pe_ring`, used by the baseline).
|
||||
- **Why a tree, not the baseline fan-out:** attention needs `O` at one
|
||||
place (the PE that writes `O_base`). A tree is `⌈log₂ N⌉` steps
|
||||
(3 for `N=8`) vs the baseline all-to-all's `N−1` (7). For decode, where
|
||||
the per-PE sweep is short, the reduction is the dominant cost, so this
|
||||
is a real win. (The baseline's fan-out replicates the answer to all
|
||||
ranks — unnecessary here.)
|
||||
- **Payload:** `O_i` (`d`-vector) is heavy; `m,ℓ` are scalars per `(g,
|
||||
T_q)`. Sent as separate `tl.send`s (the baseline sends the triplet as
|
||||
three sends — same pattern).
|
||||
Because the merge is associative **and** commutative, the combine is
|
||||
**reduce-to-root** (attention needs `O` at one place — the rank that writes
|
||||
`O_base` — not an all-reduce) and proceeds in **two levels**:
|
||||
|
||||
Kernel structure (greenlet):
|
||||
### 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.
|
||||
|
||||
### 4.2 Level-1 — intra-CUBE-Group (C CUBE roots → Group root)
|
||||
|
||||
- A **center-root bidirectional CUBE-mesh reduce** over the `C` CUBEs of
|
||||
the group (PE-root-only, `program_id(axis=1)` = `cube_id`), over the
|
||||
CUBE NOC 2D mesh (ADR-0017). This **adapts the proven inter-CUBE pattern
|
||||
of `lrab_hierarchical_allreduce.py` (Phases 1–2: row-then-col converge at
|
||||
the center CUBE)** — but **reduce-only** (drop its broadcast-back Phases
|
||||
4–5; attention needs root-only) and with the **log-sum-exp `merge`**
|
||||
replacing its plain `+`. Center-root halves the critical path vs a corner
|
||||
root (`lrab_hierarchical_allreduce.py:113-116`). No inter-SIP phase — the
|
||||
CUBE Group is intra-SIP (§0).
|
||||
|
||||
### 4.3 Data-driven overlap, no global barrier
|
||||
|
||||
- A rank sends up the tree the **instant its own local P·V finishes** — it
|
||||
does **not** wait for sibling ranks; internal nodes `merge` children as
|
||||
they arrive. So the reduction **overlaps the compute of slower ranks**
|
||||
(causal prefill imbalance, batched decode tokens).
|
||||
- The two levels **pipeline**: a CUBE whose Level-2 reduction is done feeds
|
||||
Level-1 immediately, while other CUBEs are still doing Level-2 — **no
|
||||
barrier between levels**. (A rank cannot send before its *own* local
|
||||
sweep completes; sending pre-final partials would multiply the heavy `O`
|
||||
payload and is not worth it.)
|
||||
- **Why reduce-to-root, not the baseline fan-out / an all-reduce:** the
|
||||
baseline replicates the answer to every rank (`R−1` steps); attention
|
||||
needs it once. Reduce-to-root is `⌈log₂ P⌉ + (center-mesh depth over C)`
|
||||
— for decode (short sweep, reduction-dominated) this is the dominant win.
|
||||
|
||||
### 4.4 Configurable topology
|
||||
|
||||
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).
|
||||
|
||||
**Payload:** `O_i` (`[G·T_q, d]`) is heavy; `m,ℓ` are light. Sent as
|
||||
separate `tl.send`s (baseline pattern).
|
||||
|
||||
Kernel structure (greenlet), both levels reuse `merge`:
|
||||
|
||||
```python
|
||||
# leaf / internal node: fold children that send to me, then forward up
|
||||
for child_dir in tree_children_dirs(pe_id, N):
|
||||
m_c = tl.recv(child_dir, m.shape); l_c = tl.recv(child_dir, l.shape)
|
||||
O_c = tl.recv(child_dir, O.shape)
|
||||
m, l, O = merge(m, l, O, m_c, l_c, O_c)
|
||||
if not is_root(pe_id):
|
||||
tl.send(parent_dir(pe_id), m); tl.send(parent_dir(pe_id), l); tl.send(parent_dir(pe_id), O)
|
||||
else:
|
||||
tl.store(O_base, O / l)
|
||||
def hierarchical_reduce_and_store(m, l, O, cube_id, pe_id, C, P, o_base):
|
||||
# ---- Level-2: PE tree within the CUBE → CUBE root PE ----
|
||||
for child in tree_children_dirs(pe_id, P): # PE IPCQ N/S/E/W
|
||||
m, l, O = merge(m, l, O, *recv_triplet(child))
|
||||
if not is_cube_root(pe_id, P):
|
||||
send_triplet(parent_dir(pe_id, P), m, l, O); return
|
||||
# ---- Level-1: CUBE-mesh center-root reduce → Group root (cube roots only) ----
|
||||
for child in mesh_children_dirs(cube_id, C): # CUBE NOC, center-root
|
||||
m, l, O = merge(m, l, O, *recv_triplet(child))
|
||||
if is_group_root(cube_id, C):
|
||||
tl.store(o_base, O / l) # normalise once
|
||||
else:
|
||||
send_triplet(parent_dir_mesh(cube_id, C), m, l, O)
|
||||
```
|
||||
|
||||
`tl.recv` blocks until data arrives (`tl_context.py:446-499`); the merge
|
||||
order is fixed by the static tree, so no data-dependent `pop` is needed —
|
||||
which is why **no hardware/composite `pop`-as-dependency change is
|
||||
required** (§6).
|
||||
The reduction tree/mesh is **static** (fixed at compile time), so the recv
|
||||
order is data-independent — no data-dependent `pop`, plain blocking
|
||||
`tl.recv` suffices, and **no hardware/composite `pop`-as-dependency change
|
||||
is required** (§6).
|
||||
|
||||
---
|
||||
|
||||
## 5. Per-case kernels
|
||||
|
||||
All four cases share §3 (tile sweep) + §4 (reduction); they differ only
|
||||
in `N`, query-block width, and which `tile_*` predicates fire.
|
||||
in the SP extents `C·P`, query-block width, and which `tile_*` predicates
|
||||
fire.
|
||||
|
||||
### 5.1 Common skeleton
|
||||
|
||||
```python
|
||||
def attn_kernel(q_ptr, k_ptr, v_ptr, o_ptr, counter, start_pe, N,
|
||||
q_block, softmax_scale, *, tl):
|
||||
pe_id = tl.program_id(axis=rank_axis)
|
||||
my_len = valid_len(counter, start_pe, pe_id, N)
|
||||
def attn_kernel(q_ptr, k_ptr, v_ptr, o_ptr, counter, start_pe, start_cube,
|
||||
C, P, q_block, softmax_scale, *, 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)
|
||||
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) # [G·T_q, d]; lazy tl.load, G folded into M (no broadcast)
|
||||
q_g = load_Q_group(q_ptr, kv) # [G·T_q, d]; lazy tl.load, G folded into M (no broadcast)
|
||||
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
|
||||
if N == 1:
|
||||
tl.store(o_ptr, O / l) # no reduction
|
||||
if C == 1 and P == 1:
|
||||
tl.store(o_base(kv), O / l) # no reduction (single rank)
|
||||
else:
|
||||
tree_reduce_and_store(pe_id, N, o_ptr) # §4
|
||||
hierarchical_reduce_and_store(m, l, O, cube_id, pe_id, C, P, o_base(kv)) # §4
|
||||
```
|
||||
|
||||
### 5.2 DECODE, no SP (`N=1`, one PE owns the head's KV)
|
||||
### 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
|
||||
final ragged tile.
|
||||
- **GQA reuse is the whole game** (decode is KV-load-bound): the `G=8`
|
||||
@@ -530,19 +615,25 @@ def attn_kernel(q_ptr, k_ptr, v_ptr, o_ptr, counter, start_pe, N,
|
||||
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
|
||||
- If `S_rank` fits in scratch (small/medium context) this degenerates to a
|
||||
**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.
|
||||
(§3) only kicks in when `S_rank` 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`
|
||||
tokens. Each PE GQA-batches its `G=8` query rows over its local tiles.
|
||||
- `T_q=1` ⇒ short per-PE sweep ⇒ **reduction dominates** ⇒ the §4 tree
|
||||
(3 steps) is the structurally important part. Reduction latency is
|
||||
hidden by other concurrent decode tokens when batched, or by the long
|
||||
per-PE tile sweep for long single-stream context.
|
||||
### 5.3 DECODE, with SP / KV-parallel (`C × P` ranks)
|
||||
|
||||
- One request's KV is hierarchical-round-robin across the `C·P` ranks of
|
||||
the CUBE Group (Level-1 over `C` CUBEs, Level-2 over `P` PEs); each rank
|
||||
owns ≈`my_len` tokens and GQA-batches its `G=8` query rows over its local
|
||||
tiles.
|
||||
- `T_q=1` ⇒ short per-rank sweep ⇒ **reduction dominates** ⇒ the §4 2-level
|
||||
reduce (`⌈log₂ P⌉` + center-mesh over `C`) is the structurally important
|
||||
part. Reduction latency is hidden by other concurrent decode tokens when
|
||||
batched, by the long per-rank sweep for long single-stream context, and
|
||||
by the §4.3 level-pipelining. For short decode prefer a **small `C`**
|
||||
(less inter-CUBE reduction); push `C` up only when context length demands
|
||||
the extra KV-parallelism.
|
||||
|
||||
### 5.4 PREFILL, no SP
|
||||
- Whole prompt resident; query is a block of `T_q` tokens (chunked).
|
||||
@@ -554,17 +645,21 @@ def attn_kernel(q_ptr, k_ptr, v_ptr, o_ptr, counter, start_pe, N,
|
||||
- GQA batches the group's query heads along the same axis as decode.
|
||||
|
||||
### 5.5 PREFILL, with SP (Ring KV)
|
||||
- KV sharded across `N` PEs (the ring); 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).
|
||||
|
||||
- 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 query-owning PE; if KV-parallel
|
||||
splits remain, §4 tree-merges them, then normalise + store.
|
||||
- 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.
|
||||
|
||||
The baseline `_attention_mesh_kv` already implements the ring fold; this
|
||||
ADR adds GQA reuse, causal step-skip, and `recv_async` overlap.
|
||||
@@ -575,11 +670,12 @@ ADR adds GQA reuse, causal step-skip, and `recv_async` overlap.
|
||||
|
||||
- The only place a HW/composite `pop`-as-dependency would help is a lone
|
||||
reduction with no other work to hide a poll — i.e. batch=1 **and** short
|
||||
context. The target is **agentic = low batch, long context** ⇒ each PE
|
||||
has many KV tiles; the §4 tree's `tl.recv` blocks are covered by the
|
||||
sweep work of concurrent tokens / long context.
|
||||
- The §4 reduction uses a **static** tree, so the recv order is fixed at
|
||||
compile time — no data-dependent pop. `tl.recv` (blocking) suffices.
|
||||
context. The target is **agentic = low batch, long context** ⇒ each rank
|
||||
has many KV tiles; the §4 reduce's `tl.recv` blocks are covered by the
|
||||
sweep work of concurrent tokens / long context / the §4.3 level-pipeline.
|
||||
- The §4 reduction uses a **static** 2-level tree/mesh, so the recv order is
|
||||
fixed at compile time — no data-dependent pop. `tl.recv` (blocking)
|
||||
suffices.
|
||||
- **Decision: ship with the greenlet `tl.send`/`tl.recv` collective.**
|
||||
Revisit a HW `pop` only if a short-context, single-stream,
|
||||
latency-critical target emerges.
|
||||
@@ -592,7 +688,7 @@ 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) |
|
||||
| reduction scheduling, IPCQ send/recv ordering | **kernel** (static 2-level tree/mesh) |
|
||||
| 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) |
|
||||
|
||||
@@ -620,12 +716,15 @@ new primitive** — only the kernel restructuring in §5.2 (per KV head,
|
||||
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 —
|
||||
- **2-level reduce-to-root** (§4) replacing the baseline all-to-all fan-out
|
||||
— Level-2 PE tree (intra-CUBE) + Level-1 center-root CUBE-mesh reduce
|
||||
(intra-CUBE-Group, adapting the `lrab_hierarchical_allreduce` inter-CUBE
|
||||
pattern as reduce-only + log-sum-exp), data-driven and level-pipelined,
|
||||
pure kernel control flow over `tl.send`/`tl.recv`.
|
||||
- Causal tile skip + additive boundary mask (§3/§5.4) — kernel `if` +
|
||||
a mask tensor added with `+`.
|
||||
- Round-robin KV placement / valid-length arithmetic (§2) — launch-arg
|
||||
arithmetic + `DPPolicy(pe="row_wise")`.
|
||||
- 2-level round-robin KV placement / valid-length arithmetic (§0, §2) —
|
||||
launch-arg arithmetic + `DPPolicy` `row_wise` over both `cube` and `pe`.
|
||||
|
||||
**New primitives for efficiency / scale (each has a supporting ADR):**
|
||||
|
||||
@@ -669,14 +768,17 @@ new primitive** — only the kernel restructuring in §5.2 (per KV head,
|
||||
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.
|
||||
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.
|
||||
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`
|
||||
5. **One-shot vs tiled crossover for decode** (§5.2) — the `S_rank`
|
||||
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.
|
||||
@@ -691,19 +793,19 @@ new primitive** — only the kernel restructuring in §5.2 (per KV head,
|
||||
|
||||
| Case | KV placement | Inner loop | Reduction | Masking |
|
||||
|---|---|---|---|---|
|
||||
| Decode, no SP | 1 PE, all KV | one-shot or tiled, GQA-batched | none | last tile only |
|
||||
| Decode, SP | round-robin `N` PEs | tiled, GQA-batched | §4 tree (`⌈log₂ N⌉`) | last tile only |
|
||||
| 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-sharded | per ring step (recv_async overlap) | running state + §4 tree | causal step-skip + boundary |
|
||||
| Prefill, SP (Ring) | ring over `C·P` ranks | per ring step (recv_async overlap) | running state + §4 2-level | causal step-skip + boundary |
|
||||
|
||||
**I/O per case** (full contract §0.5):
|
||||
|
||||
| Case | Inputs | Output | IPCQ partials |
|
||||
|---|---|---|---|
|
||||
| Decode, no SP | `Q[G,1,d]`, full `K/V[S,d]` on 1 PE | `O[G,1,d]` at `O_base` | none |
|
||||
| Decode, SP | `Q[G,1,d]`, per-PE `K/V[S_pe,d]` | `O[G,1,d]` (root) | `(m,ℓ,O_i)` per PE → tree |
|
||||
| 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 |
|
||||
| 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]` (root) | running `(m,ℓ,O)` + tree |
|
||||
| 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 |
|
||||
|
||||
In all cases: KV-cache writes and RoPE happen **upstream**; output
|
||||
projection **downstream**; score `S` and probs `P` are never materialised.
|
||||
@@ -733,8 +835,12 @@ 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. **Tree reduction step count:** decode-SP issues `⌈log₂ N⌉` reduction
|
||||
rounds (not `N−1`); op_log `ipcq_send`/`recv` counts match the tree.
|
||||
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.
|
||||
5. **Long context (ADR-0063):** a sweep at `S` that overflows 1 MiB
|
||||
@@ -769,10 +875,10 @@ predicted default; revise on review.
|
||||
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
|
||||
[d, TILE]`, but the KV cache stores K as `[S_rank, 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
|
||||
data. **Recommend:** store K **pre-transposed** `[d, S_rank]` 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).
|
||||
@@ -815,3 +921,39 @@ predicted default; revise on review.
|
||||
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.
|
||||
|
||||
### Items from the hierarchical CUBE-Group SP pivot (this revision)
|
||||
|
||||
1. **A 4-SIP topology config is needed for full scale.** The shipped
|
||||
`topology.yaml` has `sips: count: 2`; the full `H_kv=8` model at `C=8`
|
||||
needs **4 SIPs** (2 KV heads each, §0). **Recommend:** add a 4-SIP
|
||||
topology config (16-CUBE 4×4 mesh per SIP, 8 PE/CUBE — already the per-SIP
|
||||
shape) for the headline eval; validation-scale runs can use fewer
|
||||
CUBEs/SIPs. The bench must remain single-device per launch (one SIP).
|
||||
|
||||
2. **`head_of_cube` / CUBE-Group partition of the 4×4 mesh.** With `C=8`,
|
||||
each SIP's 16-CUBE mesh splits into 2 CUBE Groups of 8. The exact sub-mesh
|
||||
shape (4×2 vs 2×4) sets the Level-1 center-root mesh hop counts and which
|
||||
CUBEs are 1-hop neighbours. **Recommend:** pick the partition that keeps
|
||||
each CUBE Group a contiguous rectangular sub-mesh with a true center CUBE
|
||||
(so `lrab`-style center-root applies); verify against the CUBE NOC routing
|
||||
(ADR-0017).
|
||||
|
||||
3. **`C` as a per-operating-point knob.** `C` should differ by case (small
|
||||
for decode where reduction dominates, large for long-context prefill).
|
||||
**Recommend:** expose `C` (and per-level topology, §4.4) as launch/bench
|
||||
config, sweep it in the milestone bench, and report the decode-vs-prefill
|
||||
optimum. Default `C` TBD by measurement.
|
||||
|
||||
4. **Reverses the §0 "never spans CUBEs" non-goal — confirm no downstream
|
||||
assumption depends on it.** Earlier text and possibly other ADRs treated
|
||||
intra-CUBE-only reduction as invariant. **Recommend:** grep for that
|
||||
assumption (SFR installs, address policy, diagrams) before implementation;
|
||||
none found in this kernel's scope, but the `configure_sfr_*` neighbour
|
||||
wiring must expose CUBE-NOC `N/S/E/W` for Level-1, not just PE IPCQ.
|
||||
|
||||
5. **`valid_len_2level` / hierarchical round-robin correctness.** The
|
||||
two-axis `(cube_local·P + pe_local)` placement (§2.1) must tile the
|
||||
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.
|
||||
|
||||
Reference in New Issue
Block a user