Files
kernbench2/scripts/paper/measure_gqa_decode_placement_comm.py
mukesh bb668a3ac3 paper(gqa): derive (m,l,O) AR per-PE cost from kernel topology
Replaces the unsourced 32 KB/layer placeholder in the analytical chart
with T*(N-1)/N where T = h_q * S_q * (d_head*2 + 8) ~ 2.1 KB
(per-PE (m, l, O) payload) and N is the participant count of the
reduce-only chain per stage. Cases 2/3/6 analytical bars now match
the simulator-measured numbers within ~10% (was 6-30x over).

Also:
- Bumps the measurement S_kv from 8 K to 64 K to verify Cases 1, 2, 3, 6
  are genuinely S_kv-independent and Cases 4, 5 scale linearly.
- In-bar descriptor on the paired chart is now wrapped to fit the bar
  width, black, no background box, and centered between the analytical
  and measured bars so the label applies to both.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-06-18 15:38:20 -07:00

363 lines
15 KiB
Python
Raw Permalink 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.
"""Measured comm overlay for all 6 GQA decode KV placements.
Runs the simulator for Cases 1-6 (the same 6 placements the chart in
paper_plot_gqa_4cases_summary.py covers analytically) at S_kv = 64 K,
sums actual IPCQ-copy bytes from the engine op_log, projects to
per-token (x80 layers), scales the partial-score-AR component of
Cases 4/5 from S_kv = 64 K -> S_kv = 1 M (linear in S_kv; other
cases are S_kv-independent), adds the constant Wo + FFN AR
(1.25 MB / token), and writes the result to JSON for
paper_plot_gqa_4cases_summary.py to overlay on the analytical bars.
Single layer of decode attention only — the projection × 80 takes
the per-layer measurement to a per-token total.
Usage:
python scripts/paper/measure_gqa_decode_placement_comm.py
"""
from __future__ import annotations
import json
import os
import sys
from pathlib import Path
from kernbench.benches.gqa_helpers.long_ctx._gqa_attention_decode_long_ctx_cube_repl_pe_sp import (
gqa_attention_decode_long_ctx_cube_repl_pe_sp_kernel as _case3_kernel,
)
from kernbench.benches.gqa_helpers.long_ctx._gqa_attention_decode_long_ctx_cube_repl_pe_tp import (
gqa_attention_decode_long_ctx_cube_repl_pe_tp_kernel as _case1_kernel,
)
from kernbench.benches.gqa_helpers.long_ctx._gqa_attention_decode_long_ctx_cube_sp_pe_sp import (
gqa_attention_decode_long_ctx_cube_sp_pe_sp_kernel as _case6_kernel,
)
from kernbench.benches.gqa_helpers.long_ctx._gqa_attention_decode_long_ctx_cube_sp_pe_tp import (
gqa_attention_decode_long_ctx_cube_sp_pe_tp_kernel as _case2_kernel,
)
from kernbench.benches.gqa_helpers.long_ctx._gqa_attention_decode_long_ctx_cube_sp_pe_tp_dhead import (
gqa_attention_decode_long_ctx_cube_sp_pe_tp_dhead_kernel as _case4_kernel,
)
from kernbench.benches.gqa_helpers.long_ctx._gqa_attention_decode_long_ctx_cube_tp_dhead_pe_sp import (
gqa_attention_decode_long_ctx_cube_tp_dhead_pe_sp_kernel as _case5_kernel,
)
from kernbench.benches.gqa_helpers.shared._gqa_panel_helpers import _ccl_cfg
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
_C = 8
_P = 8
_N_LAYERS = 80
_S_KV_MEAS = 64 * 1024 # per-run simulator S_kv (1/16th of headline)
_S_KV_HEADLINE = 1 << 20 # 1 Mi tokens, the chart's headline S_kv
# Per-cube S_kv share for the d_head-TP partial-score AR cost.
# Case 4 (Cube-SP × PE-TP_dhead): per-cube = S_kv / C
# Case 5 (Cube-TP_dhead × PE-SP): per-cube = S_kv (KV replicated across
# cubes for the cube-axis d_head-TP), so partial-score AR scales
# with the full S_kv.
# Headline / measured per-cube ratios give the partial-score-AR scale-up
# from S_kv = 64 K to S_kv = 1 M. (m,,O) AR is S_kv-independent.
_PARTIAL_SCORE_SCALE = _S_KV_HEADLINE / _S_KV_MEAS # = 16
# Per-token Wo + FFN AR (constant across all cases, comes from the
# attn-output and FFN-down all-reduces NOT measured by the attention-
# only kernel run here).
_WO_PER_LAYER_BYTES = 8 * 1024
_FFN_PER_LAYER_BYTES = 8 * 1024
_WO_FFN_PER_TOKEN_BYTES = (
(_WO_PER_LAYER_BYTES + _FFN_PER_LAYER_BYTES) * _N_LAYERS
) # 1.25 MB
# Total PE count in one KV-head group — average per-PE comm = total / N.
_NUM_PES = _C * _P
# Total partial-score-AR slice produced by the attention compute when
# d_head is sharded. Used to split measured IPCQ traffic into the
# S_kv-scaling component (partial scores) vs the S_kv-independent
# component ((m,,O) merge). Same as the analytical formula in
# paper_plot_gqa_4cases_summary.py: h_q · S_q · per_cube_S_kv · 2 bytes.
_H_Q = 8
_S_Q = 1
_BYTES_PER_ELEM = 2
_PARAMS = dict(C=_C, P=_P, T_q=_S_Q, S_kv=_S_KV_MEAS,
d_head=128, h_q=_H_Q, h_kv=1)
def _bench_fn_case1(ctx):
"""Case 1: Cube-Repl x PE-repl (PE-TP doesn't shard KV)."""
p = _PARAMS
configure_sfr_intercube_multisip(ctx.engine, ctx.spec, _ccl_cfg())
dp = DPPolicy(cube="replicate", pe="replicate",
num_cubes=p["C"], num_pes=p["P"])
q = ctx.zeros((p["T_q"], p["h_q"] * p["d_head"]),
dtype="f16", dp=dp, name="q_c1")
k = ctx.zeros((p["S_kv"], p["h_kv"] * p["d_head"]),
dtype="f16", dp=dp, name="k_c1")
v = ctx.zeros((p["S_kv"], p["h_kv"] * p["d_head"]),
dtype="f16", dp=dp, name="v_c1")
o = ctx.empty((p["T_q"], p["h_q"] * p["d_head"]),
dtype="f16", dp=dp, name="o_c1")
ctx.launch("case1_repl_repl", _case1_kernel,
q, k, v, o,
p["T_q"], p["S_kv"], p["h_q"], p["h_kv"],
p["d_head"], p["C"], p["P"],
_auto_dim_remap=False)
def _bench_fn_case2(ctx):
"""Case 2: Cube-SP x PE-repl (PE-TP doesn't shard KV)."""
p = _PARAMS
configure_sfr_intercube_multisip(ctx.engine, ctx.spec, _ccl_cfg())
dp_full = DPPolicy(cube="replicate", pe="replicate",
num_cubes=p["C"], num_pes=p["P"])
dp_kv = DPPolicy(cube="row_wise", pe="replicate",
num_cubes=p["C"], num_pes=p["P"])
q = ctx.zeros((p["T_q"], p["h_q"] * p["d_head"]),
dtype="f16", dp=dp_full, name="q_c2")
k = ctx.zeros((p["S_kv"], p["h_kv"] * p["d_head"]),
dtype="f16", dp=dp_kv, name="k_c2")
v = ctx.zeros((p["S_kv"], p["h_kv"] * p["d_head"]),
dtype="f16", dp=dp_kv, name="v_c2")
o = ctx.empty((p["T_q"], p["h_q"] * p["d_head"]),
dtype="f16", dp=dp_full, name="o_c2")
ctx.launch("case2_sp_repl", _case2_kernel,
q, k, v, o,
p["T_q"], p["S_kv"], p["h_q"], p["h_kv"],
p["d_head"], p["C"], p["P"],
_auto_dim_remap=False)
def _bench_fn_case3(ctx):
"""Case 3: Cube-Repl x PE-SP."""
p = _PARAMS
configure_sfr_intercube_multisip(ctx.engine, ctx.spec, _ccl_cfg())
dp_full = DPPolicy(cube="replicate", pe="replicate",
num_cubes=p["C"], num_pes=p["P"])
dp_kv = DPPolicy(cube="replicate", pe="row_wise",
num_cubes=p["C"], num_pes=p["P"])
q = ctx.zeros((p["T_q"], p["h_q"] * p["d_head"]),
dtype="f16", dp=dp_full, name="q_c3")
k = ctx.zeros((p["S_kv"], p["h_kv"] * p["d_head"]),
dtype="f16", dp=dp_kv, name="k_c3")
v = ctx.zeros((p["S_kv"], p["h_kv"] * p["d_head"]),
dtype="f16", dp=dp_kv, name="v_c3")
o = ctx.empty((p["T_q"], p["h_q"] * p["d_head"]),
dtype="f16", dp=dp_full, name="o_c3")
ctx.launch("case3_repl_sp", _case3_kernel,
q, k, v, o,
p["T_q"], p["S_kv"], p["h_q"], p["h_kv"],
p["d_head"], p["C"], p["P"],
_auto_dim_remap=False)
def _bench_fn_case4(ctx):
p = _PARAMS
configure_sfr_intercube_multisip(ctx.engine, ctx.spec, _ccl_cfg())
dp_full = DPPolicy(cube="replicate", pe="column_wise",
num_cubes=p["C"], num_pes=p["P"])
dp_kv = DPPolicy(cube="row_wise", pe="column_wise",
num_cubes=p["C"], num_pes=p["P"])
q = ctx.zeros((p["T_q"], p["h_q"] * p["d_head"]),
dtype="f16", dp=dp_full, name="q_c4")
k = ctx.zeros((p["S_kv"], p["h_kv"] * p["d_head"]),
dtype="f16", dp=dp_kv, name="k_c4")
v = ctx.zeros((p["S_kv"], p["h_kv"] * p["d_head"]),
dtype="f16", dp=dp_kv, name="v_c4")
o = ctx.empty((p["T_q"], p["h_q"] * p["d_head"]),
dtype="f16", dp=dp_full, name="o_c4")
ctx.launch("case4_dhead_tp", _case4_kernel,
q, k, v, o,
p["T_q"], p["S_kv"], p["h_q"], p["h_kv"],
p["d_head"], p["C"], p["P"],
_auto_dim_remap=False)
def _bench_fn_case5(ctx):
p = _PARAMS
configure_sfr_intercube_multisip(ctx.engine, ctx.spec, _ccl_cfg())
dp_q = DPPolicy(cube="column_wise", pe="replicate",
num_cubes=p["C"], num_pes=p["P"])
dp_kv = DPPolicy(cube="column_wise", pe="row_wise",
num_cubes=p["C"], num_pes=p["P"])
q = ctx.zeros((p["T_q"], p["h_q"] * p["d_head"]),
dtype="f16", dp=dp_q, name="q_c5")
k = ctx.zeros((p["S_kv"], p["h_kv"] * p["d_head"]),
dtype="f16", dp=dp_kv, name="k_c5")
v = ctx.zeros((p["S_kv"], p["h_kv"] * p["d_head"]),
dtype="f16", dp=dp_kv, name="v_c5")
o = ctx.empty((p["T_q"], p["h_q"] * p["d_head"]),
dtype="f16", dp=dp_q, name="o_c5")
ctx.launch("case5_dhead_tp_inter", _case5_kernel,
q, k, v, o,
p["T_q"], p["S_kv"], p["h_q"], p["h_kv"],
p["d_head"], p["C"], p["P"],
_auto_dim_remap=False)
def _bench_fn_case6(ctx):
p = _PARAMS
configure_sfr_intercube_multisip(ctx.engine, ctx.spec, _ccl_cfg())
dp_full = DPPolicy(cube="replicate", pe="replicate",
num_cubes=p["C"], num_pes=p["P"])
dp_kv = DPPolicy(cube="row_wise", pe="row_wise",
num_cubes=p["C"], num_pes=p["P"])
q = ctx.zeros((p["T_q"], p["h_q"] * p["d_head"]),
dtype="f16", dp=dp_full, name="q_c6")
k = ctx.zeros((p["S_kv"], p["h_kv"] * p["d_head"]),
dtype="f16", dp=dp_kv, name="k_c6")
v = ctx.zeros((p["S_kv"], p["h_kv"] * p["d_head"]),
dtype="f16", dp=dp_kv, name="v_c6")
o = ctx.empty((p["T_q"], p["h_q"] * p["d_head"]),
dtype="f16", dp=dp_full, name="o_c6")
ctx.launch("case6_sp_sp", _case6_kernel,
q, k, v, o,
p["T_q"], p["S_kv"], p["h_q"], p["h_kv"],
p["d_head"], p["C"], p["P"],
_auto_dim_remap=False)
def _sum_ipcq_bytes(op_log) -> int:
"""Sum nbytes across all ipcq_copy records."""
return sum(
r.params.get("nbytes", 0)
for r in op_log
if r.op_kind == "memory" and r.op_name == "ipcq_copy"
)
def _partial_score_slices(case: int) -> int:
"""Divisor that splits S_kv into partial-score tiles, per the
analytical model in paper_plot_gqa_4cases_summary.py:
partial_score_per_PE = h_q * S_q * (s_kv / slices) * bytes
Case 4 (Cube-SP x PE-TP-dhead): intra-cube AR over d_head shards
on PE axis -> partial tile per PE has per-cube S_kv = s_kv/C.
Case 5 (Cube-TP-dhead x PE-SP): inter-cube AR over d_head shards
on cube axis -> partial tile per PE has per-PE S_kv = s_kv/P.
Cases 1, 2, 3, 6: no partial-score AR (only (m,l,O) merge).
"""
if case == 4:
return _C
if case == 5:
return _P
return 0
def _split_attn_layer_bytes(case: int, total_ipcq_bytes: int,
s_kv: int) -> tuple[int, int]:
"""Split per-layer attention-time IPCQ bytes into:
(partial_score_component, mlo_component).
The partial-score component scales with s_kv (so it must be scaled
when projecting from the measure-time s_kv to the headline s_kv);
the (m,l,O) component is constant in s_kv.
Partial-score size is analytically known per-case (formula in
_partial_score_slices); the remainder is treated as (m,l,O) + any
other S_kv-independent overhead. Per-PE = total / NUM_PES.
"""
per_pe_total = total_ipcq_bytes // _NUM_PES
slices = _partial_score_slices(case)
if slices == 0:
# No partial-score AR for this case.
return 0, per_pe_total
partial_score_per_pe = (
_H_Q * _S_Q * (s_kv // slices) * _BYTES_PER_ELEM
)
partial_score_per_pe = min(partial_score_per_pe, per_pe_total)
mlo_per_pe = per_pe_total - partial_score_per_pe
return partial_score_per_pe, mlo_per_pe
_KERNELS = (
(1, "Case 1 (Cube-Repl x PE-repl)", _bench_fn_case1),
(2, "Case 2 (Cube-SP x PE-repl)", _bench_fn_case2),
(3, "Case 3 (Cube-Repl x PE-SP)", _bench_fn_case3),
(4, "Case 4 (Cube-SP x PE-TP d_head)", _bench_fn_case4),
(5, "Case 5 (Cube-TP d_head x PE-SP)", _bench_fn_case5),
(6, "Case 6 (Cube-SP x PE-SP) [*]", _bench_fn_case6),
)
def main() -> int:
topology = os.environ.get("GQA_1H_TOPOLOGY", "topology.yaml")
topo = resolve_topology(topology)
out: dict = {
"S_kv_measured": _S_KV_MEAS,
"S_kv_headline": _S_KV_HEADLINE,
"n_layers": _N_LAYERS,
"num_pes": _NUM_PES,
"wo_ffn_per_token_bytes": _WO_FFN_PER_TOKEN_BYTES,
"cases": {},
}
print(f"Measuring at S_kv={_S_KV_MEAS:,} ; scaling partial-score AR "
f"to S_kv={_S_KV_HEADLINE:,} (×{int(_PARTIAL_SCORE_SCALE)})")
print()
for case_id, label, bench_fn in _KERNELS:
try:
res = run_bench(
topology=topo, bench_fn=bench_fn,
device=resolve_device(None),
engine_factory=lambda t, d: GraphEngine(
getattr(t, "topology_obj", t), enable_data=True,
),
)
except Exception as e:
print(f" {label:<42} FAIL: {type(e).__name__}: {e}")
return 1
if not res.completion.ok:
print(f" {label:<42} ENGINE FAIL: {res.completion}")
return 1
total_ipcq = _sum_ipcq_bytes(res.engine.op_log)
partial_pe, mlo_pe = _split_attn_layer_bytes(
case_id, total_ipcq, _S_KV_MEAS,
)
# Per-token attention-time comm at S_kv = 1 M:
# (partial_score_per_layer × scale + mlo_per_layer) × 80 layers
scaled_partial_per_token = (
partial_pe * int(_PARTIAL_SCORE_SCALE) * _N_LAYERS
)
mlo_per_token = mlo_pe * _N_LAYERS
attn_per_token = scaled_partial_per_token + mlo_per_token
total_per_token = attn_per_token + _WO_FFN_PER_TOKEN_BYTES
out["cases"][str(case_id)] = {
"label": label,
"total_ipcq_bytes_one_layer": total_ipcq,
"per_pe_partial_score_bytes_one_layer": partial_pe,
"per_pe_mlo_bytes_one_layer": mlo_pe,
"per_pe_attn_bytes_per_token_at_1M": attn_per_token,
"per_pe_total_bytes_per_token_at_1M": total_per_token,
}
print(f" {label:<42} "
f"ipcq_total={total_ipcq:>10,} "
f"per_pe_attn(1L)={(partial_pe + mlo_pe):>9,} "
f"per_pe_total/tok@1M={total_per_token / (1<<20):>7.2f} MB")
out_path = (
Path(__file__).resolve().parents[2]
/ "src" / "kernbench" / "benches"
/ "1H_milestone_output" / "gqa" / "long_ctx"
/ "gqa_long_ctx_6cases_measured_comm.json"
)
out_path.parent.mkdir(parents=True, exist_ok=True)
out_path.write_text(json.dumps(out, indent=2))
print()
print(f"wrote {out_path}")
return 0
if __name__ == "__main__":
sys.exit(main())