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>
This commit is contained in:
@@ -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}"
|
||||
)
|
||||
Reference in New Issue
Block a user