8102ddbe30
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>
363 lines
15 KiB
Python
363 lines
15 KiB
Python
"""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())
|