Files
kernbench2/scripts/paper/paper_plot_gqa_decode_long_ctx_4cases.py
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

249 lines
9.7 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.
"""Comparative figures for milestone-gqa-decode-long-ctx-4cases.
Reads sweep_decode.json (emitted by the milestone-1h-gqa bench) and
writes four PNGs into the same bench-output dir
(src/kernbench/benches/1H_milestone_output/gqa/long_ctx/):
gqa_decode_long_ctx_6cases_latency.png end-to-end latency per case
gqa_decode_long_ctx_6cases_traffic.png ipcq/dma op-count breakdown
gqa_decode_long_ctx_6cases_memory.png per-PE KV bytes per case
gqa_decode_long_ctx_6cases_parallelism.png per-PE S_local (compute work)
Filename still says "4cases" for backwards compat, but the script now
covers all SIX kv-sharding strategies from the analytical chart
(`gqa_4cases_summary.png`) — the original 4 plus the two new
d_head-TP variants:
Case 1 Cube-Repl × PE-repl (PE-TP doesn't shard KV)
Case 2 Cube-SP × PE-repl
Case 3 Cube-Repl × PE-SP
Case 4 Cube-SP × PE-TP-d_head ← NEW
Case 5 Cube-TP-d_head × PE-SP ← NEW
Case 6 ★ Cube-SP × PE-SP (Pareto-best)
Run (after the bench):
GQA_DECODE_LONG_CTX_4CASES_RUN=1 python -m kernbench.cli.main run \\
--bench milestone-gqa-decode-long-ctx-4cases --topology topology.yaml
python scripts/paper/paper_plot_gqa_decode_long_ctx_4cases.py
"""
from __future__ import annotations
import json
from pathlib import Path
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt # noqa: E402
_REPO_ROOT = Path(__file__).resolve().parents[2]
# Sweep JSON + PNGs live together under the bench output dir.
_FIG_DIR = (
_REPO_ROOT / "src" / "kernbench" / "benches"
/ "1H_milestone_output" / "gqa" / "long_ctx"
)
_SWEEP_JSON = _FIG_DIR / "sweep_decode.json"
# Panel name → (short label, case ordinal, accent flag) using the
# analytical chart's memory-descending ordering. PE-TP doesn't shard
# KV memory, so the cube_repl_pe_tp panel maps to Case 1 (no
# sharding, KV-wise) and cube_sp_pe_tp panel maps to Case 2.
_NORMAL, _OVERFLOW, _PARETO = "normal", "overflow", "pareto"
_CASE_INFO = {
# panel name label ord flag
"single_kv_group_decode_long_ctx_gqa_cube_repl_pe_tp": ("Case 1\nCube-Repl × PE-repl", 1, _OVERFLOW),
"single_kv_group_decode_long_ctx_gqa_cube_sp_pe_tp": ("Case 2\nCube-SP × PE-repl", 2, _OVERFLOW),
"single_kv_group_decode_long_ctx_gqa_cube_repl_pe_sp": ("Case 3\nCube-Repl × PE-SP", 3, _OVERFLOW),
"single_kv_group_decode_long_ctx_gqa_cube_sp_pe_tp_dhead": ("Case 4\nCube-SP × PE-TP-d_head", 4, _NORMAL),
"single_kv_group_decode_long_ctx_gqa_cube_tp_dhead_pe_sp": ("Case 5\nCube-TP-d_head × PE-SP", 5, _NORMAL),
"single_kv_group_decode_long_ctx_gqa_cube_sp_pe_sp": ("Case 6 ★\nCube-SP × PE-SP", 6, _PARETO),
}
# Bar fill colour per flag (used by every panel).
_FLAG_COLOR = {
_NORMAL: "#888888", # neutral grey
_OVERFLOW: "#c0504d", # red — fails the per-PE HBM budget
_PARETO: "#3b6ea5", # blue — Pareto-best
}
def _load() -> list[dict]:
return json.loads(_SWEEP_JSON.read_text())["rows"]
def _sorted_by_case(rows: list[dict]) -> list[dict]:
return sorted(rows, key=lambda r: _CASE_INFO[r["panel"]][1])
def _bar_colors(rows: list[dict]) -> list[str]:
return [_FLAG_COLOR[_CASE_INFO[r["panel"]][2]] for r in rows]
def _plot_latency(rows: list[dict]) -> Path:
rows = _sorted_by_case(rows)
labels = [_CASE_INFO[r["panel"]][0] for r in rows]
lat_us = [r["latency_ns"] / 1e3 for r in rows]
fig, ax = plt.subplots(figsize=(12.0, 4.8))
bars = ax.bar(labels, lat_us, color=_bar_colors(rows), width=0.6)
ax.set_ylabel("end-to-end latency (µs)")
ax.set_title(
"Long-context decode 6-cases — end-to-end latency per case\n"
"LLaMA-3.1-70B single-KV-head group (8 cubes × 8 PEs)"
)
ax.bar_label(bars, fmt="%.1f", padding=3, fontsize=9)
ax.grid(axis="y", ls=":", alpha=0.5)
ax.set_ylim(0, max(lat_us) * 1.15)
fig.tight_layout()
out = _FIG_DIR / "gqa_decode_long_ctx_6cases_latency.png"
fig.savefig(out, dpi=150)
plt.close(fig)
return out
def _plot_traffic(rows: list[dict]) -> Path:
rows = _sorted_by_case(rows)
labels = [_CASE_INFO[r["panel"]][0] for r in rows]
x = list(range(len(rows)))
keys = ["ipcq_copy_count", "dma_read_count", "dma_write_count"]
disp = ["IPCQ copy", "DMA read", "DMA write"]
colors = ["#c0504d", "#9bbb59", "#8064a2"]
w = 0.25
fig, ax = plt.subplots(figsize=(11.0, 4.5))
for i, (k, d, c) in enumerate(zip(keys, disp, colors)):
vals = [r["op_log_summary"][k] for r in rows]
ax.bar([xi + (i - 1) * w for xi in x], vals, width=w, label=d, color=c)
ax.set_xticks(list(x))
ax.set_xticklabels(labels, fontsize=9)
ax.set_ylabel("op count")
ax.set_title("Long-context decode 6-cases — op-count breakdown per case")
ax.legend(fontsize=9)
ax.grid(axis="y", ls=":", alpha=0.5)
fig.tight_layout()
out = _FIG_DIR / "gqa_decode_long_ctx_6cases_traffic.png"
fig.savefig(out, dpi=150)
plt.close(fig)
return out
def _s_local_per_pe(panel: str, *, S_kv: int, C: int, P: int) -> int:
"""S_local (token count) each PE attends over locally.
cube_repl_pe_tp (Case 1): S_kv (no sharding, KV-wise)
cube_sp_pe_tp (Case 2): S_kv / C (cube splits S_kv, PEs replicate)
cube_repl_pe_sp (Case 3): S_kv / P
cube_sp_pe_tp_dhead (Case 4): S_kv / C (cube splits S_kv, PE splits d_head)
cube_tp_dhead_pe_sp (Case 5): S_kv / P (cube splits d_head, PE splits S_kv)
cube_sp_pe_sp (Case 6 ★): S_kv / (C·P)
"""
cube_splits_s = "cube_sp" in panel
pe_splits_s = "pe_sp" in panel
S_per_cube = S_kv // C if cube_splits_s else S_kv
return S_per_cube // P if pe_splits_s else S_per_cube
def _d_head_per_pe(panel: str, *, d_head: int, C: int, P: int) -> int:
"""d_head dims each PE owns (Cases 4 and 5 shard d_head)."""
if "cube_tp_dhead" in panel: # Case 5: cube shards d_head
return d_head // C
if "pe_tp_dhead" in panel: # Case 4: PE shards d_head
return d_head // P
return d_head # Cases 1, 2, 3, 6: full d_head per PE
def _active_pe_count(panel: str, *, C: int, P: int) -> int:
"""Number of PEs doing non-idle attention work.
cube_repl_pe_tp (Case 1): 1 (PE-TP idle for B=1; only one PE works)
cube_sp_pe_tp (Case 2): C (PE 0 of each cube; 7 PEs idle per cube)
cube_repl_pe_sp (Case 3): C·P (all PEs busy, cube-side redundant)
cube_sp_pe_tp_dhead (Case 4): C·P (PE shards d_head — all 64 active)
cube_tp_dhead_pe_sp (Case 5): C·P (PE shards S_kv — all active)
cube_sp_pe_sp (Case 6 ★): C·P (all 64 PEs doing unique work)
"""
if "cube_repl" in panel and "pe_tp" in panel and "dhead" not in panel:
return 1
if "cube_sp" in panel and "pe_tp" in panel and "dhead" not in panel:
return C
return C * P
def _kv_bytes_per_pe(panel: str, *, S_kv: int, h_kv: int,
d_head: int, C: int, P: int) -> int:
"""KV bytes a single PE references (K + V, f16, 2 B/elem)."""
s_local = _s_local_per_pe(panel, S_kv=S_kv, C=C, P=P)
d_local = _d_head_per_pe(panel, d_head=d_head, C=C, P=P)
return 2 * s_local * h_kv * d_local * 2
def _plot_memory(rows: list[dict]) -> Path:
"""Per-PE KV bytes — Case 6 ★ wins (64-way split)."""
rows = _sorted_by_case(rows)
labels = [_CASE_INFO[r["panel"]][0] for r in rows]
mib_per_pe = [
_kv_bytes_per_pe(
r["panel"], S_kv=r["S_kv"], h_kv=r["h_kv"],
d_head=r["d_head"], C=r["C"], P=r["P"],
) / (1024 * 1024)
for r in rows
]
fig, ax = plt.subplots(figsize=(12.0, 4.8))
bars = ax.bar(labels, mib_per_pe, color=_bar_colors(rows), width=0.6)
ax.set_ylabel("KV bytes per PE (MiB, K + V, f16)")
ax.set_title(
"Long-context decode 6-cases — KV memory per PE\n"
"(one KV-head group; per-layer, per-token state)"
)
ax.bar_label(bars, fmt="%.3f", padding=3, fontsize=9)
ax.grid(axis="y", ls=":", alpha=0.5)
ax.set_ylim(0, max(mib_per_pe) * 1.15)
fig.tight_layout()
out = _FIG_DIR / "gqa_decode_long_ctx_6cases_memory.png"
fig.savefig(out, dpi=150)
plt.close(fig)
return out
def _plot_parallelism(rows: list[dict]) -> Path:
"""Total active PE-token compute load — exposes redundant-work cases."""
rows = _sorted_by_case(rows)
labels = [_CASE_INFO[r["panel"]][0] for r in rows]
total_work = [
_active_pe_count(r["panel"], C=r["C"], P=r["P"])
* _s_local_per_pe(r["panel"], S_kv=r["S_kv"], C=r["C"], P=r["P"])
for r in rows
]
fig, ax = plt.subplots(figsize=(12.0, 4.8))
bars = ax.bar(labels, total_work, color=_bar_colors(rows), width=0.6)
ax.set_ylabel("active-PE × S_local (PE-tokens; lower ⇒ less wasted work)")
ax.set_title(
"Long-context decode 6-cases — total compute load across active PEs\n"
"(Case 3 replicates KV across 8 cubes → 8× wasted PE-tokens; "
"Case 6 ★ is fully parallel without replication)"
)
ax.bar_label(bars, fmt="%d", padding=3, fontsize=9)
ax.grid(axis="y", ls=":", alpha=0.5)
ax.set_ylim(0, max(total_work) * 1.15)
fig.tight_layout()
out = _FIG_DIR / "gqa_decode_long_ctx_6cases_parallelism.png"
fig.savefig(out, dpi=150)
plt.close(fig)
return out
def main() -> None:
rows = _load()
_FIG_DIR.mkdir(parents=True, exist_ok=True)
p1 = _plot_latency(rows)
p2 = _plot_traffic(rows)
p3 = _plot_memory(rows)
p4 = _plot_parallelism(rows)
print(f"wrote {p1}")
print(f"wrote {p2}")
print(f"wrote {p3}")
print(f"wrote {p4}")
if __name__ == "__main__":
main()