"""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())