gqa: S_kv tile sweep for decode_long/decode_short/prefill_short (ADR-0063 §A.2)

Replace each kernel's local one-shot partial with a Tile-0 bootstrap + a
``scratch_scope``-wrapped merge loop. Per-rank scratch is now bounded by
``TILE_S_KV = 1024`` regardless of ``S_local``, lifting the long-context S
ceiling. Backward compatible at ``S_local <= TILE_S_KV`` (loop body is
empty; op_log structurally identical to today).

Framework: extend ``memory_store.read`` to support partial reads at
offsets within a stored region. Models real Triton ``tl.load(ptr+offset,
shape=...)`` — needed because each tile loads ``k_ptr + tile_start*row_bytes``
for its sub-slice of the per-rank KV region.

prefill_long is intentionally left untiled. Its ring loop carries
full-slice ``Kc``/``Vc`` for ``tl.send`` between CUBEs, so tile-sweeping
step 0 wouldn't shrink the kernel's actual scratch footprint. Lifting
prefill_long's ceiling requires a tile-granular ring rewrite — separate
phase. Docstring + inline comments cleaned to reflect the current shape.

Docstrings across all 4 GQA kernels: drop P1a/P2a/P2b/P6a/P6b lineage
paragraphs and historical deviation lists; describe each kernel by what
it does, not how it got here. Add explicit ``# Local attention`` /
``# Communication`` section headers.

Tests: new ``tests/attention/test_gqa_tile_sweep.py`` with 5 tests
covering single-tile op_log stability, multi-tile ``copy_to`` emission
across all 3 refactored kernels, and the headline ceiling-lift
(decode_long at S_kv=256K). Full focused regression green: 85/85 across
``tests/attention/`` + Phase E + TL discipline tests.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
This commit is contained in:
2026-06-10 15:33:29 -07:00
parent 5a76ed4f6a
commit 91cdfebb67
6 changed files with 564 additions and 203 deletions
+65 -71
View File
@@ -1,64 +1,32 @@
"""GQA fused-attention decode kernel — P1a + P2a + P2b (2-level SP).
"""GQA fused-attention decode kernel — long context (ADR-0060).
Lineage (DDD-0060 §7 phase plan):
P1a : real GQA via M-fold using ``tl.dot``; one-shot per rank.
P2a : intra-CUBE PE-level chain reduce-to-root (single-CUBE SP).
P2b : adds inter-CUBE chain reduce-to-root (multi-CUBE SP);
switches to the canonical full SFR install
``configure_sfr_intercube_multisip`` with disjoint namespaces
(``intra_*`` for PE, ``N/S/E/W`` for CUBE, ``global_*`` for SIP).
P3b : tile S_kv sweep + ``tl.scratch_scope`` (deferred).
Later : P1b composite swap; P4 lazy load; P5 opt3 pipelining; P6 prefill.
Each rank holds an ``S_local = S_kv / (C·P)`` slice of K, V and the full
Q (replicated). The local attention is computed via an S_kv-axis tile
sweep (ADR-0063 §A.2) so per-rank scratch is bounded by ``TILE_S_KV``
regardless of ``S_local``. The partial ``(m, , O)`` is then reduced
to PE 0 of CUBE 0 via a 2-level chain (intra-CUBE row+col, then
inter-CUBE), and the root writes the final output.
P2b SFR + topology assumptions:
- SFR install: ``configure_sfr_intercube_multisip`` is required when
P > 1 or C > 1 (provides ``intra_*`` and ``E/W/N/S`` namespaces).
- Intra-CUBE PE layout: logical 2×4 grid (no wrap):
Row 0: PE 0, 1, 2, 3
Row 1: PE 4, 5, 6, 7
- Inter-CUBE layout: CUBEs of one CUBE Group are laid out as a 1D row
(single row of C CUBEs, no wrap). Multi-row CUBE Group placement
(e.g. 2×4) is future work — head_of_group / cube_start dispatch is
a P7 concern (DDD-0060 §4.1).
Topology / SFR:
- Requires ``configure_sfr_intercube_multisip`` when ``P > 1`` or
``C > 1`` (provides disjoint ``intra_*`` and ``E/W/N/S`` namespaces).
- Intra-CUBE PEs are arranged as a 2×4 grid (no wrap).
- Inter-CUBE CUBEs are arranged as a 1D row (no wrap).
Reduce strategy — chain reduce-to-root at PE 0 of CUBE 0:
Level-2 (intra-CUBE, row-then-col chain on 2×4 grid):
Row chain along ``intra_W``: rightmost-col PEs send leftward;
leftmost-col PE of each row holds its row's partial.
Col bridge along ``intra_N``: PE 4 (col-0, row-1) sends to PE 0
(col-0, row-0). Only relevant when P > 4.
Result: PE 0 of each CUBE holds the CUBE's partial.
Level-1 (inter-CUBE, only PE 0 of each CUBE participates):
Chain along ``W``: rightmost CUBE sends leftward; CUBE 0 of the
CUBE Group ends with the final answer.
Final store: PE 0 of CUBE 0 normalises (O / ) and writes.
Chain step counts (per ADR-0060 §A.2 root-only output, §4 chain
deviation noted): for (C, P)=(2, 8), 7 intra-cube × 2 cubes + (C-1)
inter-cube = 14 + 1 = 15 chain steps; each step ships 3 handles
(m, , O) ⇒ 45 ``ipcq_copy`` total.
Three deliberate deviations from ADR-0060, addressed in later phases:
1. GEMMs use ``tl.dot``, not ``tl.composite`` (P1b).
2. ``softmax_scale`` omitted — needs composite epilogue mechanism (P1b/P5).
3. K loaded as ``[d, S_local]`` via byte-conserving reshape of the
deployed ``[S_local, h_kv·d]`` slice (ADR-0060 §3 / §B item 2,
reshape-not-transpose caveat; correct for zero / symmetric inputs).
Chain-vs-tree deviation from DDD-0060 §7 P2 gate (``⌈log₂ P⌉``):
P2b uses linear chain reduce-to-root (P-1 + C-1 hops). True tree on
the 2×4 PE grid requires a different SFR install — separate ADR.
Architectural intent (root-only output replacing baseline's
bidirectional fan-out) is preserved.
Layout caveats:
- GEMMs use ``tl.dot`` (no composite epilogue / ``softmax_scale``).
- K is loaded as ``(d_head, S_local)`` via byte-conserving reshape of
the deployed ``(S_local, h_kv·d_head)`` slice — correct for zero /
symmetric inputs (ADR-0060 §3 reshape-not-transpose caveat).
"""
from __future__ import annotations
TILE_S_KV = 1024 # ADR-0063 §A.2 S_kv-axis tile sweep (per-tile width).
def _merge_running(m_local, l_local, O_local, m_other, l_other, O_other, *, tl):
"""Online-softmax merge (ADR-0060 §4 / _attention_mesh_mlo baseline)."""
"""Online-softmax merge of two partial ``(m, , O)`` triples."""
m_new = tl.maximum(m_local, m_other)
scale_old = tl.exp(m_local - m_new)
scale_new = tl.exp(m_other - m_new)
@@ -82,7 +50,7 @@ def gqa_decode_long_kernel(
*,
tl,
) -> None:
"""GQA decode with M-fold + 2-level chain reduce-to-root.
"""GQA decode with M-fold + S_kv tile sweep + 2-level chain reduce-to-root.
Tensor layout:
Q : (T_q, h_q · d_head) replicated on every rank; loaded as
@@ -99,10 +67,15 @@ def gqa_decode_long_kernel(
pe_id = tl.program_id(axis=0)
cube_id = tl.program_id(axis=1)
# ── Local one-shot partial attention (M-fold on the rank's slice) ──
# ── Local attention (S_kv-axis tile sweep, ADR-0063 §A.2) ──
Q = tl.load(q_ptr, shape=(G * T_q, d_head), dtype="f16")
K_T = tl.load(k_ptr, shape=(d_head, S_local), dtype="f16")
V = tl.load(v_ptr, shape=(S_local, d_head), dtype="f16")
n_tiles = (S_local + TILE_S_KV - 1) // TILE_S_KV
KV_ROW_BYTES = d_head * 2 # f16
# Tile 0: establishes persistent (m_local, l_local, O_local).
tile_s0 = min(TILE_S_KV, S_local)
K_T = tl.load(k_ptr, shape=(d_head, tile_s0), dtype="f16")
V = tl.load(v_ptr, shape=(tile_s0, d_head), dtype="f16")
scores = tl.dot(Q, K_T)
m_local = tl.max(scores, axis=-1)
centered = scores - m_local
@@ -110,19 +83,40 @@ def gqa_decode_long_kernel(
l_local = tl.sum(exp_scores, axis=-1)
O_local = tl.dot(exp_scores, V)
# ── Level-2: intra-CUBE row-then-col chain reduce-to-(PE 0) ──
# Tiles 1..n_tiles-1: fold into running state via online-softmax merge.
# Triton port: drop the ``with tl.scratch_scope():`` line and replace
# each ``copy_to`` with a Python rebind.
for tile_idx in range(1, n_tiles):
tile_start = tile_idx * TILE_S_KV
tile_s = min(TILE_S_KV, S_local - tile_start)
with tl.scratch_scope():
K_T_t = tl.load(k_ptr + tile_start * KV_ROW_BYTES,
shape=(d_head, tile_s), dtype="f16")
V_t = tl.load(v_ptr + tile_start * KV_ROW_BYTES,
shape=(tile_s, d_head), dtype="f16")
scores_t = tl.dot(Q, K_T_t)
m_tile = tl.max(scores_t, axis=-1)
centered_t = scores_t - m_tile
exp_scores_t = tl.exp(centered_t)
l_tile = tl.sum(exp_scores_t, axis=-1)
O_tile = tl.dot(exp_scores_t, V_t)
m_new, l_new, O_new = _merge_running(
m_local, l_local, O_local, m_tile, l_tile, O_tile, tl=tl,
)
tl.copy_to(m_local, m_new)
tl.copy_to(l_local, l_new)
tl.copy_to(O_local, O_new)
# ── Communication: chain reduce-to-root at PE 0 of CUBE 0 ──
PE_GRID_COLS = 4
pe_col = pe_id % PE_GRID_COLS
pe_row = pe_id // PE_GRID_COLS
pe_cols_used = min(PE_GRID_COLS, P)
pe_rows_used = (P + PE_GRID_COLS - 1) // PE_GRID_COLS
# Row chain (along intra_W within each row, gathering leftward).
# Each merge step's intermediates are wrapped in tl.scratch_scope and
# the new running (m, , O) is persisted back to the outside-scope
# (persistent) m_local/l_local/O_local via tl.copy_to (ADR-0063 §D3/D3.1).
# Level-2 row chain (intra-CUBE, along intra_W, leftward).
if pe_cols_used > 1:
if pe_col < pe_cols_used - 1: # not rightmost: receive E
if pe_col < pe_cols_used - 1:
with tl.scratch_scope():
m_other = tl.recv(dir="intra_E", shape=m_local.shape, dtype="f16")
l_other = tl.recv(dir="intra_E", shape=l_local.shape, dtype="f16")
@@ -133,14 +127,14 @@ def gqa_decode_long_kernel(
tl.copy_to(m_local, m_new)
tl.copy_to(l_local, l_new)
tl.copy_to(O_local, O_new)
if pe_col > 0: # not leftmost: send W
if pe_col > 0:
tl.send(dir="intra_W", src=m_local)
tl.send(dir="intra_W", src=l_local)
tl.send(dir="intra_W", src=O_local)
# Col bridge (intra_N from row 1 col 0 → row 0 col 0). Only at col 0.
# Level-2 col bridge (intra-CUBE, along intra_N, row-1 → row-0).
if pe_col == 0 and pe_rows_used > 1:
if pe_row < pe_rows_used - 1: # row 0 receives from S
if pe_row < pe_rows_used - 1:
with tl.scratch_scope():
m_other = tl.recv(dir="intra_S", shape=m_local.shape, dtype="f16")
l_other = tl.recv(dir="intra_S", shape=l_local.shape, dtype="f16")
@@ -151,14 +145,14 @@ def gqa_decode_long_kernel(
tl.copy_to(m_local, m_new)
tl.copy_to(l_local, l_new)
tl.copy_to(O_local, O_new)
if pe_row > 0: # row >0 sends to N
if pe_row > 0:
tl.send(dir="intra_N", src=m_local)
tl.send(dir="intra_N", src=l_local)
tl.send(dir="intra_N", src=O_local)
# ── Level-1: inter-CUBE chain reduce (only PE 0 of each CUBE) ──
# Level-1 inter-CUBE chain (along W, leftward; only PE 0 of each CUBE).
if pe_id == 0 and C > 1:
if cube_id < C - 1: # not rightmost CUBE: recv E
if cube_id < C - 1:
with tl.scratch_scope():
m_other = tl.recv(dir="E", shape=m_local.shape, dtype="f16")
l_other = tl.recv(dir="E", shape=l_local.shape, dtype="f16")
@@ -169,12 +163,12 @@ def gqa_decode_long_kernel(
tl.copy_to(m_local, m_new)
tl.copy_to(l_local, l_new)
tl.copy_to(O_local, O_new)
if cube_id > 0: # non-root CUBE: send W
if cube_id > 0:
tl.send(dir="W", src=m_local)
tl.send(dir="W", src=l_local)
tl.send(dir="W", src=O_local)
# ── Final normalise + store (only at PE 0 of CUBE 0) ──
# ── Final normalise + store (root only) ──
if pe_id == 0 and cube_id == 0:
O_final = O_local / l_local
tl.store(o_ptr, O_final)
+65 -56
View File
@@ -1,26 +1,13 @@
"""GQA fused-attention SHORT-CONTEXT decode kernel (ADR-0060 §B.split.2).
"""GQA fused-attention decode kernel — short context (ADR-0060 §B.split.2).
Short context (S_kv < 256K, per ADR-0060 §B.split.1): each CUBE owns
``kv_per_cube`` whole KV heads, no S_kv sharding across CUBEs, no
inter-CUBE reduce. PE-SP within each CUBE: the P PEs split into
``kv_per_cube`` groups, each group does PE-SP across (P/kv_per_cube)
PEs for ONE owned head.
Short context (``S_kv < 256K``): each CUBE owns ``kv_per_cube`` whole
KV heads, with no S_kv sharding across CUBEs and no inter-CUBE reduce.
PE-SP within each CUBE: the ``P`` PEs split into ``kv_per_cube`` groups
of ``P/kv_per_cube`` PEs each; each group does PE-SP across the group
for one owned head, then the group's root PE stores its head's output.
Layout (after design iteration during Phase D — see ADR-0060 §B.split.2):
- K, V: shape ``(h_kv·S_kv, d_head)`` head-stacked, with the bench
deploying ``dp = (cube=row_wise, pe=row_wise)`` so each PE's chunk
is exactly ``(S_local, d_head)`` contiguous at its own addressable
shard. The kernel just loads at its ``k_ptr`` / ``v_ptr`` — no
offset arithmetic needed.
- Q: replicated ``(T_q, h_q·d_head)``; the kernel reshapes
byte-conservingly to ``(h_q·T_q, d_head)`` and operates on the
full stack. Other heads' rows are computed too (semantic noise);
with zero/symmetric inputs the math is unchanged. A proper
per-head Q slice would require runtime support for partial reads
of stored tensors (deferred).
- O: replicated; each group root writes the full byte-conserving
``(h_q·T_q, d_head)`` result. Multiple roots within a CUBE write
to disjoint PE-local addresses (no overwrite collision).
The local attention uses an S_kv-axis tile sweep (ADR-0063 §A.2) so
per-rank scratch is bounded by ``TILE_S_KV``.
Group layout on the 2×4 PE grid:
kv_per_cube=1, group=8 PEs (full 2×4): row chain + col bridge.
@@ -28,24 +15,27 @@ Group layout on the 2×4 PE grid:
kv_per_cube=4, group=2 PEs (adj cols): 1-step chain.
kv_per_cube=8, group=1 PE: no chain — direct write.
Chain reduce within group via the existing ``intra_E/W/N/S`` SFR
namespace (configure_sfr_intercube_multisip). After chain reduce, the
group's root PE (pe_in_group == 0) writes its working state to HBM.
Deviations from ADR-0060 (deliberate, documented):
1. GEMMs use ``tl.dot``, not ``tl.composite``.
2. ``softmax_scale`` omitted.
3. K loaded as ``[d, S_local]`` via byte-conserving reshape
(reshape-not-transpose caveat — correct for zero / symmetric inputs).
4. Q byte-conserving reshape: kernel computes attention for ALL Q
rows against the group's owned K head; only the rows for my head
are semantically meaningful. Correct for zero / symmetric inputs.
Layout caveats:
- K, V: ``(h_kv·S_kv, d_head)`` head-stacked, deployed with
``dp = (cube=row_wise, pe=row_wise)`` so each PE's chunk is
contiguously ``(S_local, d_head)`` at its own shard. K loaded as
``(d_head, S_local)`` via byte-conserving reshape (ADR-0060 §3).
- Q: replicated ``(T_q, h_q·d_head)``, reshaped byte-conservingly to
``(h_q·T_q, d_head)``. Kernel computes attention for ALL Q rows
against the group's owned K head; only the group's owned head rows
are semantically meaningful (correct for zero / symmetric inputs).
- O: replicated; each group root writes its head's
``(h_q·T_q, d_head)`` result at disjoint PE-local addresses.
- GEMMs use ``tl.dot`` (no composite epilogue / ``softmax_scale``).
"""
from __future__ import annotations
TILE_S_KV = 1024 # ADR-0063 §A.2 S_kv-axis tile sweep (per-tile width).
def _merge_running(m_local, l_local, O_local, m_other, l_other, O_other, *, tl):
"""Online-softmax merge — identical to long kernel."""
"""Online-softmax merge of two partial ``(m, , O)`` triples."""
m_new = tl.maximum(m_local, m_other)
scale_old = tl.exp(m_local - m_new)
scale_new = tl.exp(m_other - m_new)
@@ -71,25 +61,20 @@ def gqa_decode_short_kernel(
tl,
) -> None:
"""Short-context GQA decode with PE-parallel heads + intra-group PE-SP."""
group_size = P // kv_per_cube # PEs per head group
group_size = P // kv_per_cube
pe_id = tl.program_id(axis=0)
pe_in_group = pe_id % group_size
# PE-SP within group: shard S_kv across group_size PEs
S_local = S_kv // group_size
# ── Loads (DP layout already places each PE at its own shard) ──
# Q replicated → byte-conserving reshape to (h_q·T_q, d_head).
Q = tl.load(q_ptr,
shape=(h_q * T_q, d_head), dtype="f16")
# K, V row_wise per (cube, pe) → each PE has (S_local, d_head) at k_ptr.
K_T = tl.load(k_ptr,
shape=(d_head, S_local), dtype="f16")
V = tl.load(v_ptr,
shape=(S_local, d_head), dtype="f16")
# ── Local attention (S_kv-axis tile sweep, ADR-0063 §A.2) ──
Q = tl.load(q_ptr, shape=(h_q * T_q, d_head), dtype="f16")
n_tiles = (S_local + TILE_S_KV - 1) // TILE_S_KV
KV_ROW_BYTES = d_head * 2 # f16
# ── Local one-shot partial attention ──
# Tile 0: establishes persistent (m_local, l_local, O_local).
tile_s0 = min(TILE_S_KV, S_local)
K_T = tl.load(k_ptr, shape=(d_head, tile_s0), dtype="f16")
V = tl.load(v_ptr, shape=(tile_s0, d_head), dtype="f16")
scores = tl.dot(Q, K_T)
m_local = tl.max(scores, axis=-1)
centered = scores - m_local
@@ -97,15 +82,39 @@ def gqa_decode_short_kernel(
l_local = tl.sum(exp_scores, axis=-1)
O_local = tl.dot(exp_scores, V)
# ── Within-group chain reduce-to-root (Level-2 only) ──
# Tiles 1..n_tiles-1: fold into running state via online-softmax merge.
# Triton port: drop the ``with tl.scratch_scope():`` line and replace
# each ``copy_to`` with a Python rebind.
for tile_idx in range(1, n_tiles):
tile_start = tile_idx * TILE_S_KV
tile_s = min(TILE_S_KV, S_local - tile_start)
with tl.scratch_scope():
K_T_t = tl.load(k_ptr + tile_start * KV_ROW_BYTES,
shape=(d_head, tile_s), dtype="f16")
V_t = tl.load(v_ptr + tile_start * KV_ROW_BYTES,
shape=(tile_s, d_head), dtype="f16")
scores_t = tl.dot(Q, K_T_t)
m_tile = tl.max(scores_t, axis=-1)
centered_t = scores_t - m_tile
exp_scores_t = tl.exp(centered_t)
l_tile = tl.sum(exp_scores_t, axis=-1)
O_tile = tl.dot(exp_scores_t, V_t)
m_new, l_new, O_new = _merge_running(
m_local, l_local, O_local, m_tile, l_tile, O_tile, tl=tl,
)
tl.copy_to(m_local, m_new)
tl.copy_to(l_local, l_new)
tl.copy_to(O_local, O_new)
# ── Communication: within-group chain reduce-to-root (Level-2 only) ──
group_cols = min(4, group_size)
group_rows = (group_size + group_cols - 1) // group_cols
pe_col_in_group = pe_in_group % group_cols
pe_row_in_group = pe_in_group // group_cols
# Row chain along intra_W (within group's row).
# Row chain (within group's row, along intra_W, leftward).
if group_cols > 1:
if pe_col_in_group < group_cols - 1: # receive E (in-group)
if pe_col_in_group < group_cols - 1:
with tl.scratch_scope():
m_other = tl.recv(dir="intra_E", shape=m_local.shape, dtype="f16")
l_other = tl.recv(dir="intra_E", shape=l_local.shape, dtype="f16")
@@ -116,14 +125,14 @@ def gqa_decode_short_kernel(
tl.copy_to(m_local, m_new)
tl.copy_to(l_local, l_new)
tl.copy_to(O_local, O_new)
if pe_col_in_group > 0: # send W (in-group)
if pe_col_in_group > 0:
tl.send(dir="intra_W", src=m_local)
tl.send(dir="intra_W", src=l_local)
tl.send(dir="intra_W", src=O_local)
# Col bridge along intra_N (only if group spans 2 grid rows).
# Col bridge (within group, along intra_N, row-1 → row-0).
if pe_col_in_group == 0 and group_rows > 1:
if pe_row_in_group < group_rows - 1: # receive S (in-group)
if pe_row_in_group < group_rows - 1:
with tl.scratch_scope():
m_other = tl.recv(dir="intra_S", shape=m_local.shape, dtype="f16")
l_other = tl.recv(dir="intra_S", shape=l_local.shape, dtype="f16")
@@ -134,12 +143,12 @@ def gqa_decode_short_kernel(
tl.copy_to(m_local, m_new)
tl.copy_to(l_local, l_new)
tl.copy_to(O_local, O_new)
if pe_row_in_group > 0: # send N (in-group)
if pe_row_in_group > 0:
tl.send(dir="intra_N", src=m_local)
tl.send(dir="intra_N", src=l_local)
tl.send(dir="intra_N", src=O_local)
# ── Group root writes its owned head's output ──
# ── Final normalise + store (group root only) ──
if pe_in_group == 0:
O_final = O_local / l_local
tl.store(o_ptr, O_final)
+30 -37
View File
@@ -1,26 +1,26 @@
"""GQA fused-attention prefill kernel — P6a + P6b (head-parallel + Ring KV).
"""GQA fused-attention prefill kernel — long context (ADR-0060 §5.5).
Lineage (DDD-0060 §7 phase plan):
P6a : head-parallel structure (one Q head per CUBE), C=1 baseline.
P6b : add Ring KV rotation across C CUBEs (ADR-0060 §5.5). Each
CUBE rotates its KV block to its W neighbour and receives
from E; over C-1 steps every CUBE sees every block. Online-
softmax merge folds each step into running (m, , O). No
reduce — each CUBE writes its own head's output.
Requires ``configure_sfr_intercube_ring(ring_size=C)`` SFR.
P6c (later): tile T_q across the P PEs inside a CUBE for intra-CUBE
parallelism (ADR-0060 §B item 3).
Head-parallel: each CUBE owns one Q head and one KV slice. Over C ring
steps the KV blocks rotate around the C CUBEs (W-send / E-recv) so every
CUBE sees every block; the online-softmax merge folds each step into the
running ``(m, , O)``. No inter-CUBE reduce — each CUBE writes its own
head's output.
Deviations from ADR-0060 §5.5 — deferred to later phases:
1. GEMMs use ``tl.dot`` not ``tl.composite`` (parallel to P1a; lifts
when P1b decides the composite-output-handle question).
2. ``softmax_scale`` omitted — same composite-epilogue deferral.
3. K loaded as ``[d, S_local]`` via byte-conserving reshape of
``[S_local, d]`` (ADR-0060 §3 / §B item 2 reshape-not-transpose
caveat; correct for zero / symmetric inputs).
4. No causal masking / step-skip — future P6c.
5. Blocking ``tl.recv`` (not ``recv_async``) — overlap via lazy
``tl.load`` lands in P4.
Topology / SFR:
- Requires ``configure_sfr_intercube_ring(ring_size=C)`` (1D ring of
C CUBEs with wrap at the CUBE level).
- Only PE 0 of each CUBE participates (head-parallel; intra-CUBE PE
parallelism is a separate phase).
Layout caveats:
- GEMMs use ``tl.dot`` (no composite epilogue / ``softmax_scale``).
- K loaded as ``(d_head, S_local)`` via byte-conserving reshape of
the deployed ``(S_local, d_head)`` slice (ADR-0060 §3
reshape-not-transpose caveat).
- No causal masking / step-skip; blocking ``tl.recv``.
- Step-0 local partial is NOT tile-swept — its scratch bound is set
by the ring's full-slice ``Kc``/``Vc`` carry, which dominates the
score-stack regardless. Tile-granular ring is a separate phase.
"""
from __future__ import annotations
@@ -54,8 +54,7 @@ def gqa_prefill_long_kernel(
while online-softmax merges each step into running (m, , O).
"""
pe_id = tl.program_id(axis=0)
# Head-parallel: only PE 0 of each CUBE participates. P6c (future)
# will tile T_q across the 8 PEs for intra-CUBE parallelism.
# Head-parallel: only PE 0 of each CUBE participates.
if pe_id != 0:
return
@@ -64,20 +63,19 @@ def gqa_prefill_long_kernel(
Kc = tl.load(k_ptr, shape=(d_head, S_local), dtype="f16")
Vc = tl.load(v_ptr, shape=(S_local, d_head), dtype="f16")
# ── Step 0: initial partial against own KV block — establishes the
# persistent (m, , O) arena. Intermediates (scores, exp_scores) stay
# allocated; ring steps below recycle per-step intermediates inside
# tl.scratch_scope to keep peak scratch O(one step) (ADR-0063 §D3).
# ── Local attention: initial partial against own KV block ──
# Establishes the persistent (m, , O) running state.
scores = tl.dot(Q, Kc)
m = tl.max(scores, axis=-1)
exp_scores = tl.exp(scores - m)
l = tl.sum(exp_scores, axis=-1)
O = tl.dot(exp_scores, Vc)
# ── Steps 1..C-1: Ring KV rotation + online-softmax merge ──
# Per-step intermediates wrapped in tl.scratch_scope; the merged
# running (m, , O) is persisted to the outside-scope handles via
# tl.copy_to (ADR-0063 §D3.1) so its bytes survive __exit__.
# ── Communication: Ring KV rotation + online-softmax merge ──
# Each step sends K, V to W and receives from E. Per-step
# intermediates are scope-recycled; the merged (m, , O) is
# persisted via tl.copy_to. Triton port: drop the scope and
# replace each copy_to with a Python rebind.
for _ in range(1, C):
tl.send(dir="W", src=Kc)
tl.send(dir="W", src=Vc)
@@ -85,27 +83,22 @@ def gqa_prefill_long_kernel(
Vc = tl.recv(dir="E", shape=(S_local, d_head), dtype="f16")
with tl.scratch_scope():
# Partial on rotated KV
scores = tl.dot(Q, Kc)
m_step = tl.max(scores, axis=-1)
exp_scores = tl.exp(scores - m_step)
l_step = tl.sum(exp_scores, axis=-1)
O_step = tl.dot(exp_scores, Vc)
# Online-softmax merge into new running (m, , O)
m_new = tl.maximum(m, m_step)
scale_old = tl.exp(m - m_new)
scale_step = tl.exp(m_step - m_new)
l_new = l * scale_old + l_step * scale_step
O_new = O * scale_old + O_step * scale_step
# Persist new running state back to the outside-scope arena.
tl.copy_to(m, m_new)
tl.copy_to(l, l_new)
tl.copy_to(O, O_new)
# __exit__: scoped intermediates gone; persistent m, l, O carry
# the new running state into the next ring iteration.
# Final normalise + store each CUBE writes its own head's rows.
# ── Final normalise + store (each CUBE writes its own head) ──
O_final = O / l
tl.store(o_ptr, O_final)
+50 -28
View File
@@ -1,24 +1,24 @@
"""GQA fused-attention SHORT-CONTEXT prefill kernel (ADR-0060 §B.split.2).
"""GQA fused-attention prefill kernel — short context (ADR-0060 §B.split.2).
Prefill analogue of ``_gqa_decode_short.py``. Same layout decisions:
- K, V head-stacked (h_kv·S_kv, d_head) with row_wise DP so each PE
has a contiguous (S_local, d_head) shard at its own address.
- Q replicated; kernel uses byte-conserving reshape.
- O replicated; group root writes the full byte-conserving result.
Prefill analogue of ``_gqa_decode_short.py`` — same CUBE/PE layout
(``kv_per_cube`` heads per CUBE, group-PE-SP, within-group chain reduce,
no inter-CUBE reduce). The only structural difference from short decode:
``T_q`` may be > 1 (prefill processes multiple query tokens) and Q is
shaped ``(T_q, h_kv·d_head)`` — one Q head per KV head, no GQA M-fold.
No Ring KV (each owned head fully resident at its CUBE). Test
``test_short_prefill_no_ring_KV_traffic`` asserts no inter-CUBE IPCQ.
The local attention uses an S_kv-axis tile sweep (ADR-0063 §A.2) so
per-rank scratch is bounded by ``TILE_S_KV``.
The only structural difference from short decode:
- ``T_q`` may be > 1 (prefill processes multiple query tokens).
- Q is shaped (T_q, h_kv·d_head) — one Q head per KV head (no GQA
M-fold here; head-parallel prefill in ADR-0060 §5.5 is 1:1).
No Ring KV here — each owned head is fully resident at its CUBE.
"""
from __future__ import annotations
TILE_S_KV = 1024 # ADR-0063 §A.2 S_kv-axis tile sweep (per-tile width).
def _merge_running(m_local, l_local, O_local, m_other, l_other, O_other, *, tl):
"""Online-softmax merge — identical to short decode kernel."""
"""Online-softmax merge of two partial ``(m, , O)`` triples."""
m_new = tl.maximum(m_local, m_other)
scale_old = tl.exp(m_local - m_new)
scale_new = tl.exp(m_other - m_new)
@@ -42,26 +42,21 @@ def gqa_prefill_short_kernel(
*,
tl,
) -> None:
"""Short-context prefill with PE-parallel heads + intra-group PE-SP.
NO Ring KV (each CUBE owns its KV heads fully).
"""
"""Short-context prefill with PE-parallel heads + intra-group PE-SP."""
group_size = P // kv_per_cube
pe_id = tl.program_id(axis=0)
pe_in_group = pe_id % group_size
S_local = S_kv // group_size
# Q: replicated (T_q, h_kv·d_head) → byte-conserving reshape to
# (h_kv·T_q, d_head) — one Q head per KV head, stacked.
Q = tl.load(q_ptr,
shape=(h_kv * T_q, d_head), dtype="f16")
K_T = tl.load(k_ptr,
shape=(d_head, S_local), dtype="f16")
V = tl.load(v_ptr,
shape=(S_local, d_head), dtype="f16")
# ── Local attention (S_kv-axis tile sweep, ADR-0063 §A.2) ──
Q = tl.load(q_ptr, shape=(h_kv * T_q, d_head), dtype="f16")
n_tiles = (S_local + TILE_S_KV - 1) // TILE_S_KV
KV_ROW_BYTES = d_head * 2 # f16
# Tile 0: establishes persistent (m_local, l_local, O_local).
tile_s0 = min(TILE_S_KV, S_local)
K_T = tl.load(k_ptr, shape=(d_head, tile_s0), dtype="f16")
V = tl.load(v_ptr, shape=(tile_s0, d_head), dtype="f16")
scores = tl.dot(Q, K_T)
m_local = tl.max(scores, axis=-1)
centered = scores - m_local
@@ -69,12 +64,37 @@ def gqa_prefill_short_kernel(
l_local = tl.sum(exp_scores, axis=-1)
O_local = tl.dot(exp_scores, V)
# Within-group chain reduce-to-root (same machinery as short decode).
# Tiles 1..n_tiles-1: fold into running state via online-softmax merge.
# Triton port: drop the ``with tl.scratch_scope():`` line and replace
# each ``copy_to`` with a Python rebind.
for tile_idx in range(1, n_tiles):
tile_start = tile_idx * TILE_S_KV
tile_s = min(TILE_S_KV, S_local - tile_start)
with tl.scratch_scope():
K_T_t = tl.load(k_ptr + tile_start * KV_ROW_BYTES,
shape=(d_head, tile_s), dtype="f16")
V_t = tl.load(v_ptr + tile_start * KV_ROW_BYTES,
shape=(tile_s, d_head), dtype="f16")
scores_t = tl.dot(Q, K_T_t)
m_tile = tl.max(scores_t, axis=-1)
centered_t = scores_t - m_tile
exp_scores_t = tl.exp(centered_t)
l_tile = tl.sum(exp_scores_t, axis=-1)
O_tile = tl.dot(exp_scores_t, V_t)
m_new, l_new, O_new = _merge_running(
m_local, l_local, O_local, m_tile, l_tile, O_tile, tl=tl,
)
tl.copy_to(m_local, m_new)
tl.copy_to(l_local, l_new)
tl.copy_to(O_local, O_new)
# ── Communication: within-group chain reduce-to-root (Level-2 only) ──
group_cols = min(4, group_size)
group_rows = (group_size + group_cols - 1) // group_cols
pe_col_in_group = pe_in_group % group_cols
pe_row_in_group = pe_in_group // group_cols
# Row chain (within group's row, along intra_W, leftward).
if group_cols > 1:
if pe_col_in_group < group_cols - 1:
with tl.scratch_scope():
@@ -92,6 +112,7 @@ def gqa_prefill_short_kernel(
tl.send(dir="intra_W", src=l_local)
tl.send(dir="intra_W", src=O_local)
# Col bridge (within group, along intra_N, row-1 → row-0).
if pe_col_in_group == 0 and group_rows > 1:
if pe_row_in_group < group_rows - 1:
with tl.scratch_scope():
@@ -109,6 +130,7 @@ def gqa_prefill_short_kernel(
tl.send(dir="intra_N", src=l_local)
tl.send(dir="intra_N", src=O_local)
# ── Final normalise + store (group root only) ──
if pe_in_group == 0:
O_final = O_local / l_local
tl.store(o_ptr, O_final)
+64 -11
View File
@@ -52,26 +52,79 @@ class MemoryStore:
dtype: str | None = None) -> np.ndarray:
"""Read tensor from (space, addr). Returns reference, no copy.
If shape/dtype match stored tensor, returns as-is.
If dtype differs, performs reinterpret cast (view).
If shape differs but nbytes match, reshapes.
Models real Triton ``tl.load(ptr + offset, shape=...)`` semantics:
- exact base hit, exact bytes: returns the stored array (optionally
reinterpreted/reshaped — byte-conserving).
- exact base hit, fewer bytes: returns a prefix slice.
- non-base addr falling inside a stored region: returns the
slice at ``[addr - base, addr - base + requested_nbytes)``.
Out-of-bounds, no-containing-region, and dtype mismatches still
raise. Used by tile-sweep kernels that load sub-slices at
``ptr + tile_start * row_bytes``.
"""
store = self._storage.get(space)
if store is None or addr not in store:
if store is None:
raise KeyError(f"No data at ({space}, 0x{addr:x})")
arr = store[addr]
arr = store.get(addr)
base_addr = addr
if arr is None:
base_addr, arr = self._find_containing_region(store, addr)
if arr is None:
raise KeyError(f"No data at ({space}, 0x{addr:x})")
if dtype is not None:
np_dtype = _resolve_dtype(dtype)
if arr.dtype != np_dtype:
arr = arr.view(np_dtype)
if shape is not None and arr.shape != shape:
if arr.nbytes != np.prod(shape) * arr.dtype.itemsize:
if shape is None:
if base_addr != addr:
raise ValueError(
f"Shape mismatch: stored {arr.shape} ({arr.nbytes}B) "
f"vs requested {shape} ({np.prod(shape) * arr.dtype.itemsize}B)"
f"Partial read at ({space}, 0x{addr:x}) requires explicit shape"
)
arr = arr.reshape(shape)
return arr
return arr
requested_nbytes = int(np.prod(shape)) * arr.dtype.itemsize
byte_offset = addr - base_addr
if byte_offset + requested_nbytes > arr.nbytes:
raise ValueError(
f"Out-of-bounds read at ({space}, 0x{addr:x}): "
f"offset {byte_offset}B + {requested_nbytes}B > "
f"region {arr.nbytes}B"
)
if byte_offset == 0 and requested_nbytes == arr.nbytes:
return arr if arr.shape == shape else arr.reshape(shape)
itemsize = arr.dtype.itemsize
if byte_offset % itemsize != 0 or requested_nbytes % itemsize != 0:
raise ValueError(
f"Unaligned partial read at ({space}, 0x{addr:x}): "
f"offset {byte_offset}B / size {requested_nbytes}B "
f"not a multiple of {itemsize}B"
)
flat = arr.reshape(-1)
start = byte_offset // itemsize
end = start + requested_nbytes // itemsize
return flat[start:end].reshape(shape)
@staticmethod
def _find_containing_region(
store: dict[int, np.ndarray], addr: int,
) -> tuple[int, np.ndarray | None]:
"""Find the stored region whose ``[base, base+nbytes)`` contains ``addr``.
Linear scan over stored bases. Acceptable since the store is
small (one entry per deployed tensor); upgrade to a sorted
interval map only if profiling shows it dominates.
"""
for base, arr in store.items():
if base <= addr < base + arr.nbytes:
return base, arr
return 0, None
def has(self, space: str, addr: int) -> bool:
return addr in self._storage.get(space, {})
+290
View File
@@ -0,0 +1,290 @@
"""Phase 1 spec test for P3b: S_kv tile sweep in the GQA kernels
(ADR-0063 §A.2 + ADR-0060 §B "long/short context split").
The local one-shot partial in three kernels — ``_gqa_decode_long.py``,
``_gqa_decode_short.py``, ``_gqa_prefill_short.py`` — loads its rank's
entire ``(d_head, S_local)`` / ``(S_local, d_head)`` KV slice in one
shot, so per-rank scratch grows linearly with ``S_local``. The
``_attention_*`` baselines hit a TCM ceiling once ``S_local`` exceeds
the 1 MiB pool divided by the per-rank intermediate footprint.
P3b replaces the one-shot partial with a tile sweep:
- Module-level ``TILE_S_KV = 1024`` per kernel file.
- Tile 0 establishes the persistent ``(m_local, l_local, O_local)``.
- Tiles 1..n_tiles-1 wrap their intermediates (K_T tile, V tile,
scores, centered, exp_scores, partials) in ``tl.scratch_scope()``
and persist the merged ``(m, , O)`` to the persistent handles via
``tl.copy_to`` (ADR-0063 §D3 / §D3.1).
When ``S_local ≤ TILE_S_KV`` the loop body never runs and the op_log
is structurally identical to today's one-shot path — every existing
validation-scale test continues to pass.
``_gqa_prefill_long.py`` is **out of scope** for P3b. Its inner step
is IPCQ partial-recv (Ring KV rotation), not an HBM load; tiling it
intersects the ring-step structure and is a separate phase.
Phase 1 (this commit): tests only — production code lands in Phase 2.
The four multi-tile tests fail today because the kernels never call
``copy_to`` outside the chain-reduce merges (so isolating to a
chain-free config gives ``copy_to == 0`` today). The 128K test fails
today with a ``TLContext`` scratch overflow.
"""
from __future__ import annotations
from pathlib import Path
from kernbench.benches._gqa_decode_long import gqa_decode_long_kernel # noqa: F401
from kernbench.benches._gqa_decode_short import gqa_decode_short_kernel # noqa: F401
from kernbench.benches._gqa_prefill_short import gqa_prefill_short_kernel # noqa: F401
from kernbench.ccl.install import load_ccl_config, resolve_algorithm_config
from kernbench.ccl.sfr_config import configure_sfr_intercube_multisip
from kernbench.policy.placement.dp import DPPolicy
from kernbench.runtime_api.bench_runner import run_bench
from kernbench.runtime_api.types import resolve_device
from kernbench.sim_engine.engine import GraphEngine
from kernbench.topology.builder import resolve_topology
TOPOLOGY_DEFAULT = Path(__file__).resolve().parents[2] / "topology.yaml"
D_HEAD = 64
DTYPE = "f16"
def _ccl_cfg():
return resolve_algorithm_config(
load_ccl_config(), name="lrab_hierarchical_allreduce",
)
def _engine_factory(t, d):
return GraphEngine(getattr(t, "topology_obj", t), enable_data=True)
def _count(op_log, name: str) -> int:
return sum(1 for r in op_log if r.op_name == name)
# ── decode_long runner (chain-free configs isolate tile-sweep behavior) ──
def _run_decode_long(*, C: int, P: int, S_kv: int, h_q: int = 1, h_kv: int = 1):
"""Run the long-context decode kernel.
For tile-sweep isolation, callers pass C=1, P=1 — this leaves both
chain-reduce levels inactive so any ``copy_to`` in op_log comes
solely from the tile-merge body.
"""
topo = resolve_topology(str(TOPOLOGY_DEFAULT))
def _bench_fn(ctx):
configure_sfr_intercube_multisip(ctx.engine, ctx.spec, _ccl_cfg())
dp_full = DPPolicy(cube="replicate", pe="replicate",
num_cubes=C, num_pes=P)
dp_kv = DPPolicy(cube="row_wise" if C > 1 else "replicate",
pe="row_wise" if P > 1 else "replicate",
num_cubes=C, num_pes=P)
q = ctx.zeros((1, h_q * D_HEAD),
dtype=DTYPE, dp=dp_full, name=f"q_tl_c{C}_p{P}_s{S_kv}")
k = ctx.zeros((S_kv, h_kv * D_HEAD),
dtype=DTYPE, dp=dp_kv, name=f"k_tl_c{C}_p{P}_s{S_kv}")
v = ctx.zeros((S_kv, h_kv * D_HEAD),
dtype=DTYPE, dp=dp_kv, name=f"v_tl_c{C}_p{P}_s{S_kv}")
o = ctx.empty((1, h_q * D_HEAD),
dtype=DTYPE, dp=dp_full, name=f"o_tl_c{C}_p{P}_s{S_kv}")
ctx.launch(
f"gqa_decode_long_tile_c{C}_p{P}_s{S_kv}",
gqa_decode_long_kernel,
q, k, v, o,
1, S_kv, h_q, h_kv, D_HEAD, C, P,
_auto_dim_remap=False,
)
return run_bench(
topology=topo, bench_fn=_bench_fn,
device=resolve_device(None), engine_factory=_engine_factory,
)
# ── decode_short / prefill_short runners ─────────────────────────────
def _run_decode_short(*, kv_per_cube: int, C: int, P: int, S_kv: int,
h_q: int = 8, h_kv: int = 8):
"""For tile-sweep isolation: kv_per_cube=P → group_size=1 (no chain)."""
topo = resolve_topology(str(TOPOLOGY_DEFAULT))
def _bench_fn(ctx):
configure_sfr_intercube_multisip(ctx.engine, ctx.spec, _ccl_cfg())
dp_full = DPPolicy(cube="replicate", pe="replicate",
num_cubes=C, num_pes=P)
dp_kv = DPPolicy(cube="row_wise" if C > 1 else "replicate",
pe="row_wise", num_cubes=C, num_pes=P)
q = ctx.zeros((1, h_q * D_HEAD),
dtype=DTYPE, dp=dp_full, name=f"q_dsh_s{S_kv}")
k = ctx.zeros((h_kv * S_kv, D_HEAD),
dtype=DTYPE, dp=dp_kv, name=f"k_dsh_s{S_kv}")
v = ctx.zeros((h_kv * S_kv, D_HEAD),
dtype=DTYPE, dp=dp_kv, name=f"v_dsh_s{S_kv}")
o = ctx.empty((1, h_q * D_HEAD),
dtype=DTYPE, dp=dp_full, name=f"o_dsh_s{S_kv}")
ctx.launch(
f"gqa_decode_short_tile_s{S_kv}",
gqa_decode_short_kernel,
q, k, v, o,
1, S_kv, h_q, h_kv, D_HEAD, C, P, kv_per_cube,
_auto_dim_remap=False,
)
return run_bench(
topology=topo, bench_fn=_bench_fn,
device=resolve_device(None), engine_factory=_engine_factory,
)
def _run_prefill_short(*, kv_per_cube: int, C: int, P: int,
T_q: int, S_kv: int, h_kv: int = 8):
topo = resolve_topology(str(TOPOLOGY_DEFAULT))
def _bench_fn(ctx):
configure_sfr_intercube_multisip(ctx.engine, ctx.spec, _ccl_cfg())
dp_full = DPPolicy(cube="replicate", pe="replicate",
num_cubes=C, num_pes=P)
dp_kv = DPPolicy(cube="row_wise" if C > 1 else "replicate",
pe="row_wise", num_cubes=C, num_pes=P)
q = ctx.zeros((T_q, h_kv * D_HEAD),
dtype=DTYPE, dp=dp_full, name=f"q_psh_s{S_kv}")
k = ctx.zeros((h_kv * S_kv, D_HEAD),
dtype=DTYPE, dp=dp_kv, name=f"k_psh_s{S_kv}")
v = ctx.zeros((h_kv * S_kv, D_HEAD),
dtype=DTYPE, dp=dp_kv, name=f"v_psh_s{S_kv}")
o = ctx.empty((T_q, h_kv * D_HEAD),
dtype=DTYPE, dp=dp_full, name=f"o_psh_s{S_kv}")
ctx.launch(
f"gqa_prefill_short_tile_s{S_kv}",
gqa_prefill_short_kernel,
q, k, v, o,
T_q, S_kv, h_kv, D_HEAD, C, P, kv_per_cube,
_auto_dim_remap=False,
)
return run_bench(
topology=topo, bench_fn=_bench_fn,
device=resolve_device(None), engine_factory=_engine_factory,
)
# ── T1: single-tile path is op_log-stable (regression guard) ─────────
def test_decode_long_single_tile_path_unchanged():
"""ADR-0063 §A.2: when S_local <= TILE_S_KV the tile-sweep loop
body never runs and op_log is structurally identical to today's
one-shot path.
Passes today (one-shot path) AND after Phase 2 (n_tiles=1 skips
the merge loop). The witness: with chain reduce inactive (C=1,
P=1), there must be zero ``copy_to`` entries — neither today nor
after Phase 2 should the kernel emit tile-merge writebacks here.
"""
result = _run_decode_long(C=1, P=1, S_kv=64)
assert result.completion.ok, (
f"single-tile decode_long must complete; got {result.completion}"
)
n_copy = _count(result.engine.op_log, "copy")
assert n_copy == 0, (
f"S_local <= TILE_S_KV must not emit tile-merge copy_to; got {n_copy}"
)
# ── T2: decode_long multi-tile emits tile-merge copy_to ──────────────
def test_decode_long_multi_tile_emits_copy_to_merges():
"""ADR-0063 §A.2 + §D3.1: when S_local > TILE_S_KV the kernel must
wrap each subsequent tile's intermediates in ``scratch_scope`` and
persist the merged ``(m, , O)`` via ``tl.copy_to``.
Chain-free config (C=1, P=1) isolates the tile-merge body — any
``copy_to`` in op_log comes from the tile sweep, not chain reduce.
S_kv=2048, C=1, P=1 → S_local=2048 → 2 tiles → 1 merge × 3 handles
(m, , O) ⇒ 3 ``copy`` entries.
Currently 0 because the kernel performs a one-shot partial.
"""
result = _run_decode_long(C=1, P=1, S_kv=2048)
assert result.completion.ok, (
f"multi-tile decode_long must complete; got {result.completion}"
)
n_copy = _count(result.engine.op_log, "copy")
assert n_copy >= 3, (
f"decode_long multi-tile must emit >= 3 copy_to entries "
f"(1 merge × 3 handles); got {n_copy}"
)
# ── T3: decode_short multi-tile emits tile-merge copy_to ─────────────
def test_decode_short_multi_tile_emits_copy_to_merges():
"""Same property as T2 for the short decode kernel.
kv_per_cube=8, P=8, C=1 → group_size=1 (no chain reduce); S_kv=2048
→ S_local=2048 → 2 tiles per PE → 3 ``copy`` entries per PE × 8 PEs
= 24 total.
Currently 0 because the kernel performs a one-shot partial.
"""
result = _run_decode_short(kv_per_cube=8, C=1, P=8, S_kv=2048)
assert result.completion.ok, (
f"multi-tile decode_short must complete; got {result.completion}"
)
n_copy = _count(result.engine.op_log, "copy")
assert n_copy >= 3 * 8, (
f"decode_short multi-tile must emit >= 24 copy_to entries "
f"(1 merge × 3 handles × 8 PEs); got {n_copy}"
)
# ── T4: prefill_short multi-tile emits tile-merge copy_to ────────────
def test_prefill_short_multi_tile_emits_copy_to_merges():
"""Same property as T3 for the short prefill kernel.
kv_per_cube=8, P=8, C=1, T_q=4, S_kv=2048 → group_size=1 (no chain),
S_local=2048 → 2 tiles per PE → 3 ``copy`` entries per PE × 8 PEs
= 24 total.
Currently 0 because the kernel performs a one-shot partial.
"""
result = _run_prefill_short(
kv_per_cube=8, C=1, P=8, T_q=4, S_kv=2048,
)
assert result.completion.ok, (
f"multi-tile prefill_short must complete; got {result.completion}"
)
n_copy = _count(result.engine.op_log, "copy")
assert n_copy >= 3 * 8, (
f"prefill_short multi-tile must emit >= 24 copy_to entries "
f"(1 merge × 3 handles × 8 PEs); got {n_copy}"
)
# ── T5: decode_long at S_kv=256K completes (TCM ceiling lifted) ──────
def test_decode_long_context_256k_completes():
"""ADR-0063 §A.2 (test req 3): headline ceiling-lift. C=1 P=8
decode at S_kv=256K → S_local=32K. Today the per-rank score stack
(3 ×M·S_local·2 bytes, M=G·T_q=8) is ~1.5 MB → exceeds the 1 MiB
scratch pool → TLContext scratch overflow. After Phase 2 the tile
sweep bounds per-tile score stack at 3·M·TILE_S_KV·2 ≈ 48 KB and
the run completes regardless of S_kv.
"""
result = _run_decode_long(C=1, P=8, S_kv=262144, h_q=8, h_kv=1)
assert result.completion.ok, (
f"decode_long at S_kv=256K must complete after tile sweep; "
f"got {result.completion}"
)