5a76ed4f6a
Make the file + function naming symmetric:
_gqa_decode.py -> _gqa_decode_long.py
_gqa_prefill.py -> _gqa_prefill_long.py
gqa_decode_kernel -> gqa_decode_long_kernel
gqa_prefill_kernel -> gqa_prefill_long_kernel
Mirrors the existing _gqa_{decode,prefill}_short.py naming. Updates the
two imports + two call sites in milestone_gqa_headline.py and the 9
attention tests that import the kernels.
Tests: 72/72 focused regression green (tests/attention/ + Phase E + TL
discipline). milestone-gqa-headline bench passes its 7 panel/schema
assertions.
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
158 lines
5.9 KiB
Python
158 lines
5.9 KiB
Python
"""Phase 1 spec test for Phase C: long-context regression for the GQA
|
||
kernels (ADR-0060 §B "long/short context split" + ADR-0063 §A.2).
|
||
|
||
The headline panel today caps prefill at S_kv=16 / decode at S_kv≤128 —
|
||
NOT because the algorithm fails, but because the bump allocator never
|
||
recycles. ADR-0063 §A.2 (test req 3) requires a sweep at an S that
|
||
would overflow 1 MiB without recycling to complete after scope
|
||
discipline lands.
|
||
|
||
Scratch-budget estimate at C=4, P=8 (current 1 MiB pool):
|
||
decode: per-rank S_local = S_kv / 32; intermediates ≲ 500 KB at
|
||
S_kv=32K (fits today — used as regression guard).
|
||
prefill: per-rank S_local = S_kv / 4; ~16·S_local bytes per ring
|
||
step × 4 steps ≈ 64·S_local bytes. Overflows at
|
||
~64 K tokens (64 × 16K > 1 MiB).
|
||
|
||
Phase 1 (this commit): tests only — production code lands in Phase 2.
|
||
The prefill test fails today with a RuntimeError("TLContext scratch
|
||
overflow"). After Phase 2 (scratch_scope + copy_to discipline) it
|
||
completes.
|
||
"""
|
||
from __future__ import annotations
|
||
|
||
from pathlib import Path
|
||
|
||
from kernbench.benches._gqa_decode_long import gqa_decode_long_kernel # noqa: F401
|
||
from kernbench.benches._gqa_prefill_long import gqa_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_multisip,
|
||
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)
|
||
|
||
|
||
# ── Decode at moderate-long context (regression guard) ───────────────
|
||
|
||
|
||
def test_decode_long_context_32k_completes():
|
||
"""Decode at S_kv=32K (C=1, P=8) — per-rank S_local=4K. Should
|
||
complete with current scratch usage (~few KB intermediates) and
|
||
must continue to complete after the rewrite.
|
||
|
||
This is a regression guard: scratch discipline shouldn't break
|
||
decode's existing long-context capability.
|
||
"""
|
||
topo = resolve_topology(str(TOPOLOGY_DEFAULT))
|
||
|
||
def _bench_fn(ctx):
|
||
configure_sfr_intercube_multisip(ctx.engine, ctx.spec, _ccl_cfg())
|
||
P = 8
|
||
S_kv = 32_768
|
||
dp_full = DPPolicy(cube="replicate", pe="replicate",
|
||
num_cubes=1, num_pes=P)
|
||
dp_kv = DPPolicy(cube="replicate", pe="row_wise",
|
||
num_cubes=1, num_pes=P)
|
||
ctx.zeros((1, 8 * D_HEAD), dtype=DTYPE, dp=dp_full, name="q_long_dec")
|
||
k = ctx.zeros((S_kv, D_HEAD),
|
||
dtype=DTYPE, dp=dp_kv, name="k_long_dec")
|
||
v = ctx.zeros((S_kv, D_HEAD),
|
||
dtype=DTYPE, dp=dp_kv, name="v_long_dec")
|
||
o = ctx.empty((1, 8 * D_HEAD),
|
||
dtype=DTYPE, dp=dp_full, name="o_long_dec")
|
||
q = ctx.zeros((1, 8 * D_HEAD),
|
||
dtype=DTYPE, dp=dp_full, name="q_long_dec_2")
|
||
ctx.launch(
|
||
"gqa_decode_long_32k",
|
||
gqa_decode_long_kernel,
|
||
q, k, v, o,
|
||
1, S_kv, 8, 1, D_HEAD,
|
||
1, P,
|
||
_auto_dim_remap=False,
|
||
)
|
||
|
||
result = run_bench(
|
||
topology=topo, bench_fn=_bench_fn,
|
||
device=resolve_device(None), engine_factory=_engine_factory,
|
||
)
|
||
assert result.completion.ok, (
|
||
f"decode at S_kv=32K must complete; got {result.completion}"
|
||
)
|
||
|
||
|
||
# ── Prefill at long context overflows without scope discipline ───────
|
||
|
||
|
||
def test_prefill_long_context_completes_after_scope_discipline():
|
||
"""ADR-0063 §A.2 Test Req 3: a sweep at an ``S`` that would overflow
|
||
1 MiB without recycling must complete after scope discipline lands.
|
||
|
||
With C=4 and S_kv chosen so per-rank S_local·8·4 > 1 MiB
|
||
(~16K tokens per rank ⇒ S_kv ≥ 64K), the current kernel overflows
|
||
the 1 MiB scratch pool. After Phase 2 (scratch_scope wraps each
|
||
ring step's intermediates; copy_to persists running state), it
|
||
completes.
|
||
|
||
This is the headline test that proves the S-ceiling is gone for
|
||
prefill.
|
||
"""
|
||
topo = resolve_topology(str(TOPOLOGY_DEFAULT))
|
||
|
||
def _bench_fn(ctx):
|
||
configure_sfr_intercube_ring(
|
||
ctx.engine, ctx.spec, _ccl_cfg(), ring_size=4,
|
||
)
|
||
T_q = 4
|
||
S_kv = 65_536 # 64 K — per-rank 16 K, ~2 MB scratch w/o scope
|
||
C = 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="q_long_pre")
|
||
k = ctx.zeros((S_kv, D_HEAD),
|
||
dtype=DTYPE, dp=dp_kv, name="k_long_pre")
|
||
v = ctx.zeros((S_kv, D_HEAD),
|
||
dtype=DTYPE, dp=dp_kv, name="v_long_pre")
|
||
o = ctx.empty((T_q * C, D_HEAD),
|
||
dtype=DTYPE, dp=dp_o, name="o_long_pre")
|
||
ctx.launch(
|
||
"gqa_prefill_long_64k",
|
||
gqa_prefill_long_kernel,
|
||
q, k, v, o,
|
||
T_q, S_kv, D_HEAD, C,
|
||
_auto_dim_remap=False,
|
||
)
|
||
|
||
result = run_bench(
|
||
topology=topo, bench_fn=_bench_fn,
|
||
device=resolve_device(None), engine_factory=_engine_factory,
|
||
)
|
||
assert result.completion.ok, (
|
||
f"prefill at S_kv=64K must complete after scope discipline lands; "
|
||
f"got {result.completion}"
|
||
)
|