"""Phase 1 spec test for P3c: tile-granular Ring KV in prefill_long (ADR-0060 §5.5 amendment + ADR-0063 §A.2). Today's prefill_long ring sends and receives full ``(d_head, S_local)`` KV slices per step. Step 0's local-attention intermediates also live outside any ``scratch_scope`` (they're loaded as full-slice ``Kc``, ``Vc`` and feed ``scores``, ``exp_scores`` as persistent allocations). At larger ``S_local`` both the step-0 leak and the ring step's in-scope intermediates grow linearly with ``S_local``, and at ``S_local = 32K`` (S_kv=128K, C=4, T_q=4) the peak overflows the 1 MiB pool. P3c converts the ring to **tile-granular**: a nested loop ``for k in range(C): for t in range(n_tiles): ...`` where each iteration sends/recvs one ``(d_head, TILE_S_KV)`` tile (and its V counterpart). The persistent state shrinks to ``(m, ℓ, O)`` only (~1 KB); per-tile in-scope scratch is bounded by ``TILE_S_KV`` regardless of ``S_local``. Ceiling lifted. Trade-off: the per-CUBE send count grows from ``2·(C-1)`` to ``2·n_tiles·(C-1)``. At ``n_tiles=1`` (small ``S_local``) the count is unchanged, so the existing ``test_prefill_ring_c_*`` tests at ``S_kv ∈ {16, 32}`` still pass. Phase 1 (this commit): tests only — production code lands in Phase 2. T1 fails today with a ``TLContext scratch overflow``; T2 fails today with the slice-granular ipcq_copy count; T5 passes today and is a regression guard for the per-CUBE output write count. """ from __future__ import annotations from pathlib import Path from kernbench.benches._gqa_attention_prefill_long import gqa_attention_prefill_long_kernel # noqa: F401 from kernbench.ccl.install import load_ccl_config, resolve_algorithm_config from kernbench.ccl.sfr_config import configure_sfr_intercube_ring 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) def _run_prefill_ring(*, T_q: int, S_kv: int, C: int): """Head-parallel prefill with Ring KV across C CUBEs.""" topo = resolve_topology(str(TOPOLOGY_DEFAULT)) def _bench_fn(ctx): configure_sfr_intercube_ring( ctx.engine, ctx.spec, _ccl_cfg(), ring_size=C, ) dp_q = DPPolicy(cube="replicate", pe="replicate", num_cubes=C, num_pes=1) dp_kv = DPPolicy(cube="row_wise" if C > 1 else "replicate", pe="replicate", num_cubes=C, num_pes=1) dp_o = DPPolicy(cube="row_wise" if C > 1 else "replicate", pe="replicate", num_cubes=C, num_pes=1) q = ctx.zeros((T_q, D_HEAD), dtype=DTYPE, dp=dp_q, name=f"q_tlr_t{T_q}_c{C}_s{S_kv}") k = ctx.zeros((S_kv, D_HEAD), dtype=DTYPE, dp=dp_kv, name=f"k_tlr_t{T_q}_c{C}_s{S_kv}") v = ctx.zeros((S_kv, D_HEAD), dtype=DTYPE, dp=dp_kv, name=f"v_tlr_t{T_q}_c{C}_s{S_kv}") o = ctx.empty((T_q * C, D_HEAD), dtype=DTYPE, dp=dp_o, name=f"o_tlr_t{T_q}_c{C}_s{S_kv}") ctx.launch( f"gqa_prefill_long_tile_ring_t{T_q}_c{C}_s{S_kv}", gqa_attention_prefill_long_kernel, q, k, v, o, T_q, S_kv, D_HEAD, C, _auto_dim_remap=False, ) return run_bench( topology=topo, bench_fn=_bench_fn, device=resolve_device(None), engine_factory=_engine_factory, ) # ── T1: 128K ceiling lift ──────────────────────────────────────────── def test_prefill_long_context_128k_completes(): """ADR-0063 §A.2 headline ceiling lift. At S_kv=128K with C=4, S_local=32K. Today: step-0 score-stack (~768 KB persistent) + ring scope (~768 KB) → 1.5 MB peak → TLContext scratch overflow. After P3c the persistent state is just ``(m, ℓ, O)`` (≈ 1 KB) and the per-tile in-scope scratch is bounded by TILE_S_KV. """ result = _run_prefill_ring(T_q=4, S_kv=131_072, C=4) assert result.completion.ok, ( f"prefill_long at S_kv=128K must complete after tile-granular " f"ring lands; got {result.completion}" ) # ── T2: tile-granular ipcq_copy count ──────────────────────────────── def test_prefill_long_tile_granular_ipcq_count(): """ADR-0060 §5.5 amendment: with tile-granular sends, the per-CUBE send count grows from ``2·(C-1)`` to ``2·n_tiles·(C-1)``. Aggregated across all C CUBEs the total ipcq_copy becomes ``2·n_tiles·(C-1)·C``. Config: T_q=4, S_kv=4096, C=2 → S_local=2048, n_tiles=2 (with TILE_S_KV=1024). Today: slice-granular total = ``(C-1)·2·C = 4``. After P3c: tile-granular total = ``(C-1)·n_tiles·2·C = 8``. """ C = 2 n_tiles = 2 # S_local=2048 / TILE_S_KV=1024 result = _run_prefill_ring(T_q=4, S_kv=4096, C=C) assert result.completion.ok, ( f"prefill_long multi-tile ring must complete; got {result.completion}" ) n_copy = _count(result.engine.op_log, "ipcq_copy") expected = (C - 1) * n_tiles * 2 * C assert n_copy == expected, ( f"tile-granular ring: expected {expected} ipcq_copy " f"((C-1)·n_tiles·2·C = {C - 1}·{n_tiles}·2·{C}); got {n_copy}" ) # ── T5: per-CUBE distributed output unchanged (regression guard) ───── def test_prefill_long_tile_ring_dma_write_count(): """ADR-0060 §5.5: per-CUBE distributed output must hold under the tile-granular ring rewrite. Each CUBE still writes its own head's output (no inter-CUBE reduce); dma_write_count == C. Today passes; must continue to pass after P3c. """ C = 4 result = _run_prefill_ring(T_q=4, S_kv=4096, C=C) assert result.completion.ok n_writes = _count(result.engine.op_log, "dma_write") assert n_writes == C, ( f"per-CUBE distributed output: expected {C} dma_writes (one per " f"CUBE); got {n_writes}" )