bb668a3ac3
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>
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 = 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())
|