Files
kernbench2/scripts/paper/measure_gqa_decode_placement_comm.py
T
mukesh 8102ddbe30 paper(gqa): rename long-ctx artifacts to gqa_long_ctx_6cases_* and add 4 decode 6-case charts
Filename cleanup so every long-ctx GQA artifact has a consistent
"gqa_long_ctx_6cases_*" prefix (or "gqa_decode_long_ctx_6cases_*"
for decode-only charts). Old "4cases" / mixed names retired.

Renames (long_ctx + figures, content unchanged):
  gqa_hbm_budget.png                  -> gqa_long_ctx_6cases_hbm_budget.png
  gqa_4cases_summary.png              -> gqa_long_ctx_6cases_summary.png
  gqa_4cases_memory_comm_analytical   -> gqa_long_ctx_6cases_memory_comm_analytical.png
  gqa_4cases_memory_comm_paired       -> gqa_long_ctx_6cases_memory_comm_paired.png
  gqa_kv_sharding_6cases_diagram      -> gqa_long_ctx_6cases_kv_sharding_diagram.png
  gqa_kv_sharding_6cases_table        -> gqa_long_ctx_6cases_kv_sharding_table.png
  gqa_3cases_measured_comm.json       -> gqa_long_ctx_6cases_measured_comm.json
  gqa_decode_long_ctx_4cases_*.png    -> gqa_decode_long_ctx_6cases_*.png
                                         (figures dir; long_ctx never had old)

New 4-chart 6-case set in long_ctx output dir (regenerated by
paper_plot_gqa_decode_long_ctx_4cases.py, which now reads all 6
sweep_decode.json panels — Cases 1-6 with the same colour scheme
used elsewhere: red = overflow per-PE HBM, grey = neutral, blue
= Pareto-best ★):

  gqa_decode_long_ctx_6cases_latency.png
  gqa_decode_long_ctx_6cases_memory.png
  gqa_decode_long_ctx_6cases_parallelism.png
  gqa_decode_long_ctx_6cases_traffic.png

Generator scripts updated to write the new filenames + handle the
two new d_head-TP variants (Cases 4, 5) in their per-PE memory and
active-PE-count helpers. Figure widths bumped 10 -> 12 in to fit 6
multi-line case labels.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-06-18 11:49:03 -07:00

363 lines
15 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.
"""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 = 8 * 1024 # per-run simulator S_kv (small — 1/128th 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())