"""Tests for lrab-adapted center-root inter-CUBE reduce in decode_long. Verifies the ADR-0060 §4.2 prescribed Level-1 collective for the single-KV-group LLaMA-3.1-70B target (C=8 over a 2×4 sub-mesh, root at the geometric center CUBE). Activation contract: - ``sub_w == 0`` (default; omitted from launch args) → existing 1D inter-CUBE chain that converges at CUBE 0. Byte-for-byte unchanged. - ``sub_w >= 2`` with ``sub_h = C // sub_w >= 2`` → lrab-adapted center-root mesh reduce. Root cube is ``(sub_h//2)*sub_w + (sub_w//2)``. Degenerate (``sub_h < 2`` or ``sub_w < 2``) combinations are rejected at launch time per ADR-0060 §4.2 + CLAUDE.md "Simplicity First": those configurations belong to the 1D-chain code path, not lrab. Phase 1 (this commit): tests only — production code lands in Phase 2. T1, T2, T4 fail today (kernel signature does not accept ``sub_w``). T3 passes today as the backward-compat anchor for the existing 1D-chain path. """ from __future__ import annotations import re from pathlib import Path import pytest from kernbench.benches._gqa_attention_decode_long import gqa_attention_decode_long_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" _CUBE_RE = re.compile(r"\bcube(\d+)\b") 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 _dma_write_cubes(op_log) -> list[int]: """Return the list of CUBE ids that emitted a ``dma_write``. Decode's final-store path uses one ``tl.store`` from the root rank. Multiple HBM-channel-level write records may correspond to the same logical store; this just reports the source CUBE for each. """ cubes: list[int] = [] for r in op_log: if r.op_name != "dma_write": continue m = _CUBE_RE.search(r.component_id) if m is not None: cubes.append(int(m.group(1))) return cubes def _run_decode_long( *, C: int, P: int, S_kv: int, sub_w: int | None, ): """Drive a decode_long launch. ``sub_w=None`` means omit the arg entirely (exercises the kernel's current default behaviour).""" 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) T_q = 1 h_q = 8 h_kv = 1 ctx.zeros((T_q, h_q * D_HEAD), dtype=DTYPE, dp=dp_full, name=f"q_dec_2d_c{C}") k = ctx.zeros((S_kv, h_kv * D_HEAD), dtype=DTYPE, dp=dp_kv, name=f"k_dec_2d_c{C}") v = ctx.zeros((S_kv, h_kv * D_HEAD), dtype=DTYPE, dp=dp_kv, name=f"v_dec_2d_c{C}") o = ctx.empty((T_q, h_q * D_HEAD), dtype=DTYPE, dp=dp_full, name=f"o_dec_2d_c{C}") q = ctx.zeros((T_q, h_q * D_HEAD), dtype=DTYPE, dp=dp_full, name=f"q_dec_2d_c{C}_2") launch_args = [q, k, v, o, T_q, S_kv, h_q, h_kv, D_HEAD, C, P] if sub_w is not None: launch_args.append(sub_w) ctx.launch( f"gqa_decode_long_2d_c{C}_sw{sub_w}", gqa_attention_decode_long_kernel, *launch_args, _auto_dim_remap=False, ) return run_bench( topology=topo, bench_fn=_bench_fn, device=resolve_device(None), engine_factory=_engine_factory, ) # ── T1: C=8 lrab-adapted center-root completes ─────────────────────── def test_decode_2d_sub_w4_sub_h2_completes(): """ADR-0060 §4.2 prescribed center-root mesh reduce at the milestone target: C=8 over a 2×4 sub-mesh (sub_w=4, sub_h=2), P=8 PEs per CUBE. With zero inputs, output is trivially zero; the test guards that the kernel reaches completion without deadlock or scratch overflow. """ result = _run_decode_long(C=8, P=8, S_kv=2048, sub_w=4) assert result.completion.ok, ( f"decode at C=8, sub_w=4 (lrab-adapted) must complete; " f"got {result.completion}" ) # ── T2: root lives at the geometric center CUBE (cube 6) ───────────── def test_decode_2d_sub_w4_sub_h2_root_at_center_cube_6(): """For sub_w=4, sub_h=2: root_col=2, root_row=1, root_cube=6. Only the root CUBE issues the final ``tl.store`` — every other rank short-circuits. The dma_write records must come exclusively from CUBE 6. """ result = _run_decode_long(C=8, P=8, S_kv=2048, sub_w=4) assert result.completion.ok cubes = _dma_write_cubes(result.engine.op_log) assert cubes, "expected at least one dma_write for the final O store" distinct = set(cubes) assert distinct == {6}, ( f"final dma_write must come exclusively from CUBE 6 (sub_w=4, " f"sub_h=2 root); got cubes={sorted(distinct)}" ) # ── T3: backward-compat — sub_w omitted → existing 1D chain at C=4 ─── def test_decode_2d_backward_compat_sub_w0_default(): """When ``sub_w`` is omitted from the launch, the kernel must behave identically to today: 1D inter-CUBE chain along W, converging at CUBE 0. This serves as a regression anchor — Phase 2 must not change the existing C=4 panel behaviour. Passes today as well as after Phase 2. """ result = _run_decode_long(C=4, P=8, S_kv=2048, sub_w=None) assert result.completion.ok, ( f"decode at C=4 with default sub_w (1D chain) must complete; " f"got {result.completion}" ) cubes = _dma_write_cubes(result.engine.op_log) assert cubes, "expected at least one dma_write for the final O store" distinct = set(cubes) assert distinct == {0}, ( f"1D-chain root must be CUBE 0; got cubes={sorted(distinct)}" ) # ── T4: degenerate sub_w configurations are rejected ───────────────── def test_decode_2d_invalid_sub_w_rejected_when_sub_h_lt_2(): """ADR-0060 §4.2 + CLAUDE.md "Simplicity First": only sub-meshes with both sub_w >= 2 and sub_h >= 2 use lrab. C=4 with sub_w=4 would give sub_h=1 (degenerate; Phase 2 of lrab becomes a no-op). Callers must use the 1D-chain path (sub_w=0/omitted) for that case. The kernel must reject the degenerate combination with a clear error rather than silently producing a 1D-chain result at the wrong root location. """ with pytest.raises((ValueError, AssertionError), match=r"sub_[wh]"): _run_decode_long(C=4, P=8, S_kv=2048, sub_w=4)