7fad0371c5
Three logically distinct changes, bundled for atomic test green:
1. **P3c — prefill_long tile-granular Ring KV** (ADR-0060 §5.5.1 amendment).
Convert the ring from slice-granular (one full ``(d_head, S_local)``
KV slice per step) to tile-granular (``n_tiles`` tiles of
``TILE_S_KV`` per step). Nested loop with outer tile, inner ring step:
each tile propagates through all C ring positions before the next
tile starts, so IPCQ in-flight depth stays at 1 per direction.
Bootstrap at ``(t=0, k=0)`` outside the scratch_scope establishes the
persistent ``(m, ℓ, O)``; every other iteration scope-wraps + persists
via ``copy_to``. Per-rank persistent scratch shrinks to ~1 KB; per-tile
scope bounded by TILE_S_KV regardless of S_local. Headline:
prefill_long now completes at S_kv=128K (previously overflowed).
New: ``tests/attention/test_gqa_prefill_long_tile_ring.py``
(3 tests — ceiling-lift + tile-granular ipcq_copy count +
per-CUBE distributed output regression guard).
2. **Rename ``gqa_*`` → ``gqa_attention_*``** across kernel files,
function names, and importers. The "attention" name makes the role
explicit (GQA is grouped-query attention) and matches upstream Triton
FlashAttention naming conventions. Renames:
_gqa_decode_long.py -> _gqa_attention_decode_long.py
_gqa_decode_short.py -> _gqa_attention_decode_short.py
_gqa_prefill_long.py -> _gqa_attention_prefill_long.py
_gqa_prefill_short.py -> _gqa_attention_prefill_short.py
And function names ``gqa_<phase>_<context>_kernel`` →
``gqa_attention_<phase>_<context>_kernel``. Updated 1 bench file
(milestone_gqa_headline.py) and 10 test files.
3. **ADR-0060 / 0062 / 0063 / 0064: Proposed → Accepted**.
All four are reflected in production code and covered by tests:
- ADR-0060 (GQA fused attention): 4 kernels deployed; §5.5.1
amendment added for the tile-granular Ring KV introduced by P3c
(EN + KO mirror).
- ADR-0062 (lazy tl.load): LoadFuture + _await_pending live in
tl_context.py.
- ADR-0063 (tl.scratch_scope + tl.copy_to): used in every chain
reduce + tile sweep + ring step. EN-only previously; KO
translation authored as part of this commit (CLAUDE.md
bidirectional rule).
- ADR-0064 (per-op-type CPU issue cost): cpu_issue_cost.py +
issue_cost_table wiring in tl_context.py (Phase E).
Files git mv'd from docs/adr-proposed/ to docs/adr/ (EN) and
docs/adr-ko/ (KO). ADR-0061 (tl.broadcast) stays Proposed — no
implementation; documented as optional convenience primitive in
the ADR itself.
Tests: 88/88 focused regression green
(tests/attention/ + Phase E + TL discipline).
ADR pair verification: ``python tools/verify_adr_lang_pairs.py`` OK.
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
261 lines
10 KiB
Python
261 lines
10 KiB
Python
"""Phase 1 spec test for Phase D: short-context GQA kernels
|
||
(ADR-0060 §B "Items from the long/short context split", item B.split.2).
|
||
|
||
The long-context kernels (§5.2 decode / §5.5 prefill) shard each KV
|
||
head row-wise across all CUBEs. At short context (S_kv < 256K), that
|
||
shard is too thin to feed the engines and the cube-level collective
|
||
overhead dominates. The short-context kernels drop cube-SP entirely:
|
||
each CUBE owns ``kv_per_cube`` *whole* KV heads, no S_kv sharding across
|
||
CUBEs.
|
||
|
||
Design (per AskUserQuestion answers in this session):
|
||
- PE-parallel heads: P PEs split into ``kv_per_cube`` groups, each
|
||
group does PE-SP across (P/kv_per_cube) PEs for one owned head.
|
||
- scratch_scope + tl.copy_to discipline mirrors the long kernels.
|
||
- One short kernel handles kv_per_cube ∈ {1, 2, 4, 8} via a parameter.
|
||
|
||
Group layout on the 2×4 PE grid:
|
||
kv_per_cube=1, group=8 PEs (full 2×4): row chain + col bridge.
|
||
kv_per_cube=2, group=4 PEs (one row): row chain only.
|
||
kv_per_cube=4, group=2 PEs (adj cols): 1-step chain.
|
||
kv_per_cube=8, group=1 PE: no chain.
|
||
|
||
After chain reduce-to-group-root, the group's root PE writes its
|
||
owned head's output. No inter-CUBE reduce.
|
||
|
||
Phase 1 (this commit): tests only — production code lands in Phase 2.
|
||
All tests fail because the short kernels (``_gqa_attention_decode_short.py`` and
|
||
``_gqa_attention_prefill_short.py``) do not exist yet.
|
||
"""
|
||
from __future__ import annotations
|
||
|
||
from pathlib import Path
|
||
|
||
import pytest
|
||
|
||
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 short-context kernel ──────────────────────────────────────
|
||
|
||
|
||
def _run_decode_short(*, h_q: int, h_kv: int, kv_per_cube: int,
|
||
C: int, P: int, S_kv: int):
|
||
"""Run the short-context decode kernel with PE-parallel heads.
|
||
|
||
Layout (after design iteration — see Phase D failure-recovery):
|
||
Q: (T_q, h_q·D_HEAD) replicated; kernel reshapes byte-conservingly.
|
||
K, V: (h_kv·S_kv, D_HEAD) head-stacked, ``cube=row_wise, pe=row_wise``
|
||
so each PE gets a contiguous (S_local, D_HEAD) chunk at its own
|
||
addressable shard (no partial reads needed).
|
||
O: replicated; each group root writes the full byte-conserving
|
||
(h_q·T_q, D_HEAD) result.
|
||
"""
|
||
from kernbench.benches._gqa_attention_decode_short import gqa_attention_decode_short_kernel # Phase 2
|
||
|
||
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)
|
||
# Head-stacked KV with row_wise sharding so each PE's chunk is
|
||
# contiguous and exactly (S_local, D_HEAD) addressable.
|
||
dp_kv = DPPolicy(cube="row_wise", 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_short_{kv_per_cube}_{C}")
|
||
k = ctx.zeros((h_kv * S_kv, D_HEAD),
|
||
dtype=DTYPE, dp=dp_kv, name=f"k_short_{kv_per_cube}_{C}")
|
||
v = ctx.zeros((h_kv * S_kv, D_HEAD),
|
||
dtype=DTYPE, dp=dp_kv, name=f"v_short_{kv_per_cube}_{C}")
|
||
o = ctx.empty((1, h_q * D_HEAD),
|
||
dtype=DTYPE, dp=dp_full, name=f"o_short_{kv_per_cube}_{C}")
|
||
ctx.launch(
|
||
f"gqa_decode_short_{kv_per_cube}_{C}",
|
||
gqa_attention_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 test_short_decode_smoke_kv_per_cube_2_C_4():
|
||
"""ADR-0060 §B.split.2 headline: kv_per_cube=2, C=4 — each CUBE
|
||
owns 2 heads (8 heads / 4 CUBEs). PE-parallel heads splits the 8
|
||
PEs into 2 groups of 4, each group does PE-SP for one owned head.
|
||
Smoke: kernel completes."""
|
||
result = _run_decode_short(
|
||
h_q=8, h_kv=8, kv_per_cube=2, C=4, P=8, S_kv=64,
|
||
)
|
||
assert result.completion.ok, (
|
||
f"short decode kv_per_cube=2 must complete; got {result.completion}"
|
||
)
|
||
|
||
|
||
def test_short_decode_smoke_kv_per_cube_4_C_2():
|
||
"""kv_per_cube=4, C=2 — half the CUBEs participate, each owns 4
|
||
heads, 4 PE groups of 2 PEs each."""
|
||
result = _run_decode_short(
|
||
h_q=8, h_kv=8, kv_per_cube=4, C=2, P=8, S_kv=64,
|
||
)
|
||
assert result.completion.ok, (
|
||
f"short decode kv_per_cube=4 must complete; got {result.completion}"
|
||
)
|
||
|
||
|
||
def test_short_decode_smoke_kv_per_cube_8_C_1():
|
||
"""kv_per_cube=8, C=1 — all heads on one CUBE. 8 PE groups of 1 PE
|
||
each → no chain reduce, each PE writes its head's output."""
|
||
result = _run_decode_short(
|
||
h_q=8, h_kv=8, kv_per_cube=8, C=1, P=8, S_kv=64,
|
||
)
|
||
assert result.completion.ok, (
|
||
f"short decode kv_per_cube=8 must complete; got {result.completion}"
|
||
)
|
||
|
||
|
||
def test_short_decode_dma_write_count_equals_h_kv():
|
||
"""ADR-0060 §B.split.2: each owned head produces exactly one
|
||
output (no inter-CUBE reduce). Total dma_writes = h_kv across all
|
||
participating CUBEs and groups.
|
||
|
||
For h_kv=8, kv_per_cube=2, C=4: each CUBE writes 2 outputs →
|
||
4 × 2 = 8 dma_writes total.
|
||
"""
|
||
result = _run_decode_short(
|
||
h_q=8, h_kv=8, kv_per_cube=2, C=4, P=8, S_kv=64,
|
||
)
|
||
assert result.completion.ok
|
||
n_writes = _count(result.engine.op_log, "dma_write")
|
||
assert n_writes == 8, (
|
||
f"short decode: expected 8 dma_writes (h_kv); got {n_writes}"
|
||
)
|
||
|
||
|
||
def test_short_decode_no_inter_cube_traffic():
|
||
"""ADR-0060 §B.split.2: each head is fully owned by one CUBE → no
|
||
inter-CUBE reduce. The kernel must not invoke CUBE-level E/W IPCQ.
|
||
|
||
Today's long-context kernel at C=4 emits ~12 inter-CUBE ipcq_copy
|
||
via direction "E"/"W". The short kernel must emit zero of those,
|
||
keeping all IPCQ traffic on the ``intra_*`` directions.
|
||
"""
|
||
result = _run_decode_short(
|
||
h_q=8, h_kv=8, kv_per_cube=2, C=4, P=8, S_kv=64,
|
||
)
|
||
assert result.completion.ok
|
||
# Count IPCQ traffic that targeted the CUBE-level "E"/"W" directions.
|
||
inter_cube_ipcq = sum(
|
||
1 for r in result.engine.op_log
|
||
if r.op_name == "ipcq_copy"
|
||
and r.params.get("direction") in ("E", "W")
|
||
)
|
||
assert inter_cube_ipcq == 0, (
|
||
f"short decode must have no inter-CUBE E/W IPCQ; got {inter_cube_ipcq}"
|
||
)
|
||
|
||
|
||
# ── Prefill short-context kernel ─────────────────────────────────────
|
||
|
||
|
||
def _run_prefill_short(*, h_kv: int, kv_per_cube: int,
|
||
C: int, P: int, T_q: int, S_kv: int):
|
||
"""Run the short-context prefill kernel with PE-parallel heads.
|
||
|
||
Layout: same head-stacked K/V scheme as decode short, with Q/O
|
||
replicated and byte-conserving reshape inside the kernel.
|
||
"""
|
||
from kernbench.benches._gqa_attention_prefill_short import gqa_attention_prefill_short_kernel # Phase 2
|
||
|
||
topo = resolve_topology(str(TOPOLOGY_DEFAULT))
|
||
|
||
def _bench_fn(ctx):
|
||
configure_sfr_intercube_multisip(ctx.engine, ctx.spec, _ccl_cfg())
|
||
dp_q = DPPolicy(cube="replicate", pe="replicate",
|
||
num_cubes=C, num_pes=P)
|
||
dp_kv = DPPolicy(cube="row_wise", pe="row_wise",
|
||
num_cubes=C, num_pes=P)
|
||
dp_o = DPPolicy(cube="replicate", pe="replicate",
|
||
num_cubes=C, num_pes=P)
|
||
q = ctx.zeros((T_q, h_kv * D_HEAD),
|
||
dtype=DTYPE, dp=dp_q, name=f"q_pre_short_{kv_per_cube}_{C}")
|
||
k = ctx.zeros((h_kv * S_kv, D_HEAD),
|
||
dtype=DTYPE, dp=dp_kv, name=f"k_pre_short_{kv_per_cube}_{C}")
|
||
v = ctx.zeros((h_kv * S_kv, D_HEAD),
|
||
dtype=DTYPE, dp=dp_kv, name=f"v_pre_short_{kv_per_cube}_{C}")
|
||
o = ctx.empty((T_q, h_kv * D_HEAD),
|
||
dtype=DTYPE, dp=dp_o, name=f"o_pre_short_{kv_per_cube}_{C}")
|
||
ctx.launch(
|
||
f"gqa_prefill_short_{kv_per_cube}_{C}",
|
||
gqa_attention_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,
|
||
)
|
||
|
||
|
||
def test_short_prefill_smoke_kv_per_cube_2_C_4():
|
||
"""Short prefill kv_per_cube=2, C=4. Smoke: completes."""
|
||
result = _run_prefill_short(
|
||
h_kv=8, kv_per_cube=2, C=4, P=8, T_q=4, S_kv=64,
|
||
)
|
||
assert result.completion.ok, (
|
||
f"short prefill kv_per_cube=2 must complete; got {result.completion}"
|
||
)
|
||
|
||
|
||
def test_short_prefill_no_ring_KV_traffic():
|
||
"""ADR-0060 §B.split.2: short prefill DOES NOT use Ring KV — each
|
||
CUBE owns its KV heads fully, no rotation. The kernel must not
|
||
emit any KV-rotation IPCQ traffic at the CUBE level.
|
||
"""
|
||
result = _run_prefill_short(
|
||
h_kv=8, kv_per_cube=2, C=4, P=8, T_q=4, S_kv=64,
|
||
)
|
||
assert result.completion.ok
|
||
inter_cube_ipcq = sum(
|
||
1 for r in result.engine.op_log
|
||
if r.op_name == "ipcq_copy"
|
||
and r.params.get("direction") in ("E", "W")
|
||
)
|
||
assert inter_cube_ipcq == 0, (
|
||
f"short prefill must have no Ring KV (no inter-CUBE E/W IPCQ); "
|
||
f"got {inter_cube_ipcq}"
|
||
)
|