From 39fc2e953f8a4f1c7c8c1391fd1e6684bdf91260 Mon Sep 17 00:00:00 2001 From: Mukesh Garg Date: Thu, 4 Jun 2026 12:40:04 -0700 Subject: [PATCH] attention: add 8-KV-group diag harness (CCL pattern + cube_start sub-meshes) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit New ``test_attention_8kv_groups_diag.py`` probes the path from validation scale to the GQA Llama-70B 1-Q-head-per-cube headline (64 cubes = 8 KV-groups x 8 cubes/group on a 4-SIP topology) in three incremental steps: step_1_single_kv_group_at_full_breadth ONE 2x4 multi_user_decode launch (8 cubes) via the 2D mesh-mlo kernel. Verifies the per-KV-group 2D AllReduce works at full breadth. step_2_four_kv_groups_one_per_sip Four sequential 2x4 launches, one per SIP. Uses the CCL milestone's set_device pattern (milestone_1h_ccl.py:283-292): ONE run_bench with ``target_device="all"`` and ``ctx.ahbm.set_device(sip)`` between launches — not four separate run_bench calls with ``DeviceSelector("sip:N")`` (which would misuse DeviceSelector and hit a ``DPPolicy x target_device`` allocator mismatch). step_3_eight_kv_groups_two_per_sip Two 2x4 launches per SIP x 4 SIPs = 8 KV-groups (64 cubes total). Second launch per SIP uses ``cube_start=8`` to address cubes 8..15 — the disjoint sub-mesh that was unreachable before ``DPPolicy.cube_start`` landed. Demonstrates the headline-enabling use of cube_start end-to-end (DPPolicy + kernel kwarg). All three steps pass at validation dims (S_q=1, S_kv=16, h=1, d_head=64). Headline-scale dims, multi_user_prefill 2D, and the B=8 "Batch on batch" PE parallelism remain follow-up work. Co-Authored-By: Claude Opus 4.7 (1M context) --- .../test_attention_8kv_groups_diag.py | 286 ++++++++++++++++++ 1 file changed, 286 insertions(+) create mode 100644 tests/attention/test_attention_8kv_groups_diag.py diff --git a/tests/attention/test_attention_8kv_groups_diag.py b/tests/attention/test_attention_8kv_groups_diag.py new file mode 100644 index 0000000..46b78d8 --- /dev/null +++ b/tests/attention/test_attention_8kv_groups_diag.py @@ -0,0 +1,286 @@ +"""Diagnostic harness for the Llama-70B "1 Q-head per cube" target. + +Per the GQA Llama-70B sharding study at +``llm_paper_review/notes/GQA_MHA_sharding/scripts/_gen_llama70b_1M_4cases.py``, +the 1 Q-head/cube baseline uses 64 cubes (4 SIPs × 16 cubes/SIP) organized +into 8 KV-groups of 8 cubes each. Each KV-group occupies a ``2×4`` +sub-mesh within a SIP's ``4×4`` cube grid and runs the C2 2D row-then-col +AllReduce-mlo (ADR-0059 extension). This harness probes the gap between +validation and headline in three incrementally-larger steps: + + step_1_single_kv_group_at_full_breadth + ONE multi_user_decode launch on a 2×4 sub-mesh (8 cubes) of the + 4-SIP topology, via the 2D mesh-mlo kernel. Verifies the per-KV-group + 2D AllReduce works at full breadth. Smallest dim possible + (S_q=1, S_kv=16, h=1, d_head=64) to keep wall time bounded. + + step_2_four_kv_groups_one_per_sip + Four sequential multi_user_decode launches, each targeting a + different SIP. Verifies that per-SIP isolation works (each SIP holds + its own 2×4 KV-group; the SFR install only writes intra-SIP + E/W + N/S edges so the 4 groups don't see each other). + + step_3_eight_kv_groups_two_per_sip + The actual study target: 8 KV-groups, two per SIP (cubes 0..7 vs + cubes 8..15 within each SIP). Expected to FAIL with current infra — + DPPolicy doesn't take a cube offset and target_device is SIP-level, + so back-to-back launches both land on cubes 0..7 of their target SIP. + Captures what's needed to lift the 4-group cap to 8. + +Each step prints what it observed; the test asserts only the documented +expected outcomes so we can land it, watch CI, and iterate. Steps that +are *expected* to fail (step 3) are marked xfail with a precise reason. +""" +from __future__ import annotations + +import traceback +from pathlib import Path + +import pytest + +from kernbench.benches._attention_mesh_mlo_2d import attention_mesh_mlo_2d_kernel +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 DeviceSelector, resolve_device +from kernbench.sim_engine.engine import GraphEngine +from kernbench.topology.builder import resolve_topology + +TOPOLOGY_4SIP = ( + Path(__file__).resolve().parents[2] / "topologies" / "llama70b_4sip.yaml" +) +TOPOLOGY_DEFAULT = Path(__file__).resolve().parents[2] / "topology.yaml" + +S_Q_DECODE = 1 +S_KV_PER_RANK = 16 +H_Q = 1 +H_KV = 1 +D_HEAD = 64 +# 2×4 sub-mesh per KV-group (study: 8 cubes per KV-group at Q/cube=1). +MESH_ROWS = 2 +MESH_COLS = 4 +N_CUBES_PER_KV_GROUP = MESH_ROWS * MESH_COLS +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 _make_one_kv_group_bench(mesh_rows: int, mesh_cols: int): + """Return a bench_fn that runs ONE multi_user_decode kernel on a + ``mesh_rows × mesh_cols`` sub-mesh.""" + n_cubes = mesh_rows * mesh_cols + + def _bench_fn(ctx): + configure_sfr_intercube_multisip(ctx.engine, ctx.spec, _ccl_cfg()) + dp_full = DPPolicy(cube="replicate", pe="replicate", + num_cubes=n_cubes, num_pes=8) + dp_kv = DPPolicy(cube="row_wise", pe="replicate", + num_cubes=n_cubes, num_pes=8) + q = ctx.zeros((S_Q_DECODE, H_Q * D_HEAD), + dtype=DTYPE, dp=dp_full, name="q") + k = ctx.zeros((S_KV_PER_RANK * n_cubes, H_KV * D_HEAD), + dtype=DTYPE, dp=dp_kv, name="k") + v = ctx.zeros((S_KV_PER_RANK * n_cubes, H_KV * D_HEAD), + dtype=DTYPE, dp=dp_kv, name="v") + o = ctx.empty((S_Q_DECODE, H_Q * D_HEAD), + dtype=DTYPE, dp=dp_full, name="o") + ctx.launch( + f"single_kv_group_{mesh_rows}x{mesh_cols}", + attention_mesh_mlo_2d_kernel, + q, k, v, o, + S_Q_DECODE, S_KV_PER_RANK, H_Q, H_KV, D_HEAD, + mesh_rows, mesh_cols, + 1, # rank_axis=1 → cube-level ring + 0, # cube_start=0 — single sub-mesh launch + _auto_dim_remap=False, + ) + + return _bench_fn + + +def _run_one_kv_group(topology_path: Path, mesh_rows: int, mesh_cols: int, + target_device=None): + topo = resolve_topology(str(topology_path)) + captured: dict = {"engine": None} + + def factory(t, d): + eng = _engine_factory(t, d) + captured["engine"] = eng + return eng + + exc = None + result = None + try: + result = run_bench( + topology=topo, + bench_fn=_make_one_kv_group_bench(mesh_rows, mesh_cols), + device=target_device or resolve_device(None), + engine_factory=factory, + ) + except BaseException as e: # noqa: BLE001 + exc = e + return exc, result, captured["engine"] + + +# ── Step 1 — single KV-group at the study's full breadth ────────── + + +def test_step_1_single_kv_group_at_full_breadth(): + """One multi_user_decode launch on a 2×4 sub-mesh, 4-SIP topology. + + Uses the C2 2D row-then-col AllReduce-mlo kernel: stage 1 reduces + across cols (E/W) within each row, stage 2 reduces across rows (N/S). + Expected to PASS — N/S edges are wired by + ``configure_sfr_intercube_multisip`` and the 2D fan-out avoids the + row-boundary IpcqInvalidDirection that the 1D kernel hit at cube 4. + """ + if not TOPOLOGY_4SIP.exists(): + pytest.skip(f"4-SIP topology missing: {TOPOLOGY_4SIP}") + exc, result, engine = _run_one_kv_group( + TOPOLOGY_4SIP, mesh_rows=MESH_ROWS, mesh_cols=MESH_COLS, + ) + if exc is not None: + oplog_len = len(getattr(engine, "op_log", []) or []) if engine else 0 + print(f"\nstep_1 FAIL — op_log records before crash: {oplog_len}") + traceback.print_exception(type(exc), exc, exc.__traceback__) + raise AssertionError(f"step_1 failed: {exc}") from exc + assert result is not None and result.completion.ok, ( + f"step_1: completion not ok — {result.completion if result else None}" + ) + + +# ── Step 2 — 4 KV-groups, one per SIP, sequential launches ──────── + + +def _make_multi_sip_bench_fn(sip_groups: list[tuple[int, str, int]]): + """One bench_fn that does one 2×4 multi_user_decode launch per item. + + Each ``(sip, tag, cube_start)`` tuple becomes one launch: + - ``ctx.ahbm.set_device(sip)`` switches allocations to that SIP + (mirrors ``milestone_1h_ccl.py:283-292``). + - ``cube_start`` selects which 8-cube sub-mesh within the SIP: + ``0`` → cubes 0..7 (rows 0..1), ``8`` → cubes 8..15 (rows 2..3). + - ``tag`` disambiguates tensor names so launches in the same + run_bench don't collide on the allocator namespace. + """ + n_cubes = MESH_ROWS * MESH_COLS + + def _bench_fn(ctx): + configure_sfr_intercube_multisip(ctx.engine, ctx.spec, _ccl_cfg()) + for sip, tag, cube_start in sip_groups: + ctx.ahbm.set_device(sip) + dp_full = DPPolicy(cube="replicate", pe="replicate", + num_cubes=n_cubes, num_pes=8, + cube_start=cube_start) + dp_kv = DPPolicy(cube="row_wise", pe="replicate", + num_cubes=n_cubes, num_pes=8, + cube_start=cube_start) + q = ctx.zeros((S_Q_DECODE, H_Q * D_HEAD), + dtype=DTYPE, dp=dp_full, name=f"q_{tag}") + k = ctx.zeros((S_KV_PER_RANK * n_cubes, H_KV * D_HEAD), + dtype=DTYPE, dp=dp_kv, name=f"k_{tag}") + v = ctx.zeros((S_KV_PER_RANK * n_cubes, H_KV * D_HEAD), + dtype=DTYPE, dp=dp_kv, name=f"v_{tag}") + o = ctx.empty((S_Q_DECODE, H_Q * D_HEAD), + dtype=DTYPE, dp=dp_full, name=f"o_{tag}") + ctx.launch( + f"kv_group_{tag}", attention_mesh_mlo_2d_kernel, + q, k, v, o, + S_Q_DECODE, S_KV_PER_RANK, H_Q, H_KV, D_HEAD, + MESH_ROWS, MESH_COLS, + 1, # rank_axis=1 → cube-level ring + cube_start, # converts physical id → launch-local rank + _auto_dim_remap=False, + ) + + return _bench_fn + + +def _run_multi_sip(sip_groups: list[tuple[int, str, int]]): + """Run a single run_bench call covering all (sip, tag) groups.""" + topo = resolve_topology(str(TOPOLOGY_4SIP)) + captured: dict = {"engine": None} + + def factory(t, d): + eng = _engine_factory(t, d) + captured["engine"] = eng + return eng + + exc = None + result = None + try: + result = run_bench( + topology=topo, + bench_fn=_make_multi_sip_bench_fn(sip_groups), + device=resolve_device(None), # "all" SIPs in scope + engine_factory=factory, + ) + except BaseException as e: # noqa: BLE001 + exc = e + return exc, result, captured["engine"] + + +def test_step_2_four_kv_groups_one_per_sip(): + """Four multi_user_decode launches, one per SIP, in ONE run_bench call. + + Uses the CCL milestone pattern (``milestone_1h_ccl.py:283-292``): + ``target_device="all"`` scopes the runtime to every SIP; then + ``ctx.ahbm.set_device(sip)`` before each ``ctx.zeros``/``launch`` + switches which SIP the next allocation+launch lands on. This is the + canonical sequential per-SIP pattern in the codebase — four separate + ``run_bench`` calls with ``DeviceSelector("sip:N")`` is a misuse. + + Expected to PASS — the SFR install draws intra-SIP edges only, so the + 4 KV-groups can't see each other. + """ + if not TOPOLOGY_4SIP.exists(): + pytest.skip(f"4-SIP topology missing: {TOPOLOGY_4SIP}") + sip_groups = [(sip, f"sip{sip}", 0) for sip in range(4)] + exc, result, engine = _run_multi_sip(sip_groups) + if exc is not None: + oplog_len = len(getattr(engine, "op_log", []) or []) if engine else 0 + print(f"\nstep_2 FAIL — op_log records before crash: {oplog_len}") + traceback.print_exception(type(exc), exc, exc.__traceback__) + raise AssertionError(f"step_2 failed: {exc}") from exc + assert result is not None and result.completion.ok, ( + f"step_2: completion not ok — {result.completion if result else None}" + ) + + +# ── Step 3 — 8 KV-groups, two per SIP (study target) ────────────── + + +def test_step_3_eight_kv_groups_two_per_sip(): + """Two launches per SIP × 4 SIPs = 8 KV-groups total in one run_bench. + + The headline target: 64 cubes serving 8 KV-groups, two disjoint 2×4 + sub-meshes per SIP. ``cube_start=0`` puts the first KV-group on + cubes 0..7 (rows 0..1); ``cube_start=8`` puts the second on cubes + 8..15 (rows 2..3). This is the use case ``DPPolicy.cube_start`` was + added to enable. + + Expected to PASS — the SFR install draws intra-SIP edges only, so + the 8 KV-groups can't see each other; ``cube_start`` ensures the + two halves of each SIP land on disjoint cubes. + """ + if not TOPOLOGY_4SIP.exists(): + pytest.skip(f"4-SIP topology missing: {TOPOLOGY_4SIP}") + sip_groups = [ + (sip, f"sip{sip}_half{half}", half * N_CUBES_PER_KV_GROUP) + for sip in range(4) for half in (0, 1) + ] + exc, result, engine = _run_multi_sip(sip_groups) + if exc is not None: + raise AssertionError(f"step_3 failed: {exc}") from exc + assert result is not None and result.completion.ok, ( + f"step_3: completion not ok — {result.completion if result else None}" + )