Files
kernbench2/tests/attention/test_attention_8kv_groups_diag.py
T
mukesh 39fc2e953f attention: add 8-KV-group diag harness (CCL pattern + cube_start sub-meshes)
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) <noreply@anthropic.com>
2026-06-04 12:40:04 -07:00

287 lines
12 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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}"
)