"""Tests for prefill_long Ring KV at C=8 over a 2×4 snake sub-mesh. ADR-0060 §5.5 design target for the single-KV-group LLaMA-3.1-70B configuration: ``C = G = 8`` (one Q head per CUBE), Ring KV rotates the 8 KV slices around all 8 CUBEs. Increment 1 wired ``configure_sfr_intercube_ring(submesh_shape=(2, 4))`` to map the kernel's logical E/W to a snake/serpentine 1-hop path through the top 2×4 sub-mesh of the 4×4 SIP CUBE mesh. The kernel itself uses only logical ``dir="W"`` / ``dir="E"`` — the snake is transparent at kernel level. This file verifies the assembly works end-to-end at C=8: T1 kernel completes (no scratch overflow, no deadlock) T2 per-CUBE head-parallel output (8 dma_writes, one per CUBE) T3 tile-granular ring traffic at C=8 follows the same ``(C-1)·n_tiles·2·C`` formula as the existing tile-ring tests """ 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_c8_snake(*, T_q: int, S_kv: int): """Head-parallel prefill at C=8 with snake-mapped Ring KV over the top 2×4 sub-mesh of the 4×4 SIP CUBE mesh.""" topo = resolve_topology(str(TOPOLOGY_DEFAULT)) C = 8 def _bench_fn(ctx): configure_sfr_intercube_ring( ctx.engine, ctx.spec, _ccl_cfg(), submesh_shape=(2, 4), ) dp_q = DPPolicy(cube="replicate", pe="replicate", num_cubes=C, num_pes=1) dp_kv = DPPolicy(cube="row_wise", pe="replicate", num_cubes=C, num_pes=1) dp_o = DPPolicy(cube="row_wise", pe="replicate", num_cubes=C, num_pes=1) q = ctx.zeros((T_q, D_HEAD), dtype=DTYPE, dp=dp_q, name=f"q_c8_snake_t{T_q}_s{S_kv}") k = ctx.zeros((S_kv, D_HEAD), dtype=DTYPE, dp=dp_kv, name=f"k_c8_snake_t{T_q}_s{S_kv}") v = ctx.zeros((S_kv, D_HEAD), dtype=DTYPE, dp=dp_kv, name=f"v_c8_snake_t{T_q}_s{S_kv}") o = ctx.empty((T_q * C, D_HEAD), dtype=DTYPE, dp=dp_o, name=f"o_c8_snake_t{T_q}_s{S_kv}") ctx.launch( f"gqa_prefill_long_c8_snake_t{T_q}_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: kernel completes end-to-end at C=8 ─────────────────────────── def test_prefill_long_c8_snake_completes(): """ADR-0060 §5.5 design target: prefill Ring KV at C=G=8 over the 2×4 snake sub-mesh. With zero inputs, output is trivially zero; this test guards that the kernel reaches completion without deadlock, scratch overflow, or routing failure. Configuration: T_q=4, S_kv=8192 → S_local=1024, n_tiles=1 (TILE_S_KV=1024). Bounded scratch. """ result = _run_prefill_c8_snake(T_q=4, S_kv=8192) assert result.completion.ok, ( f"prefill at C=8 with snake ring must complete; " f"got {result.completion}" ) # ── T2: 8 dma_writes (head-parallel, one head per CUBE) ────────────── def test_prefill_long_c8_snake_distributed_output_count(): """Head-parallel: each CUBE owns one Q head and writes its own head's output (T_q, d_head) rows. No inter-CUBE reduce on the output side — so ``dma_write_count == C == 8``. This is the C=8 analogue of ``test_prefill_long_tile_ring_dma_write_count`` at C=4. """ result = _run_prefill_c8_snake(T_q=4, S_kv=8192) assert result.completion.ok n_writes = _count(result.engine.op_log, "dma_write") assert n_writes == 8, ( f"head-parallel C=8: expected 8 dma_writes (one per CUBE); " f"got {n_writes}" ) # ── T3: tile-granular ring ipcq_copy count scales as (C-1)·n_tiles·2·C def test_prefill_long_c8_snake_tile_ipcq_count(): """ADR-0060 §5.5.1 tile-granular ring: per-CUBE send count is ``2·n_tiles·(C-1)`` (K + V per tile per ring step). Aggregated across C CUBEs, total ipcq_copy = ``(C-1)·n_tiles·2·C``. Configuration: T_q=4, S_kv=8192, C=8 → S_local=1024, n_tiles=1 (TILE_S_KV=1024). Expected total = ``7·1·2·8 = 112``. Verifies that the snake-mapped 1-hop physical links carry the same logical ring traffic as the existing C=4 1D-row ring tests. """ C = 8 n_tiles = 1 # S_local=1024 / TILE_S_KV=1024 result = _run_prefill_c8_snake(T_q=4, S_kv=8192) assert result.completion.ok n_copy = _count(result.engine.op_log, "ipcq_copy") expected = (C - 1) * n_tiles * 2 * C assert n_copy == expected, ( f"tile-granular ring at C=8: expected {expected} ipcq_copy " f"((C-1)·n_tiles·2·C = {C - 1}·{n_tiles}·2·{C}); got {n_copy}" )