Files
kernbench2/scripts/paper/paper_plot_gqa_4cases_summary.py
T
mukesh 84bb418e1e gqa(decode): add d_head-TP kernels + measurement runner + figure generators
Two new long-ctx decode attention kernels for the d_head-TP sharding
variants the 6-case chart predicts:

  · _gqa_attention_decode_long_ctx_cube_sp_pe_tp_dhead.py
        Cube-SP × PE-TP-d_head (Case 4 in chart). Per cube holds
        S_kv/C tokens of full d_head; per PE holds same tokens but
        only d_head/P dims. Partial Q·Kᵀ scores reduced intra-cube
        before softmax; outer (m,ℓ,O) merge two-phase (intra+inter).

  · _gqa_attention_decode_long_ctx_cube_tp_dhead_pe_sp.py
        Cube-TP-d_head × PE-SP (Case 5). Per cube holds full S_kv
        with only d_head/C dims; per PE holds S_kv/P of those dims.
        Partial scores reduced inter-cube (UCIe) before softmax.

Sweep dispatch (gqa_decode_long_ctx_4cases.py) extended with two
new panels so the milestone-1h-gqa sweep covers all 6 cases.

Smoke test scripts/verify_case4_dhead_tp.py runs Cases 4/5/6 at
S_kv=2K to validate the kernels load and execute end-to-end.

Plus the figure-generation toolchain that produced the committed
PNGs in the prior commit (dd3337f):

  · paper_plot_gqa_4cases_summary.py    - 3-panel summary +
        2-panel (analytical / paired-measured) chart generator.
        _plot_comm now takes mode="analytical" | "paired".
  · paper_plot_gqa_kv_sharding_diagram.py  - 6-case 2-D KV-tensor
        diagram + companion comparison-table PNG.
  · measure_gqa_decode_placement_comm.py   - runs all 6 kernels at
        S_kv=8K, sums actual IPCQ-copy bytes from engine.op_log,
        scales partial-score AR ×128 to S_kv=1M, writes
        gqa_3cases_measured_comm.json (committed).

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

607 lines
24 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.
"""GQA 4-cases (+ 2 d_head-TP variants) — combined memory + comm summary.
Single PNG with two panels side-by-side:
(left) Per-PE KV memory at S_kv = 1 M tokens (across 80 layers).
(right) Per-PE communication per output token at S_kv = 1 M (decode).
================================================================
SYSTEM INPUTS (LLaMA-3.1-70B single-KV-head group)
================================================================
N_layers = 80 (transformer layers per token)
h_kv = 1 (per KV group)
h_q = 8 (query heads per KV group)
d_head = 128
d_model = 8192 (LLaMA-3.1-70B hidden dim — for Wo / FFN AR)
bytes = 2 (FP16)
C = 8 cubes per KV group P = 8 PEs per cube
HBM_per_PE = 6.0 GB
Attn weights = 1.76 GB → KV headroom = 4.24 GB / PE
B = 1 user S_q = 1 token (decode)
================================================================
(1) PER-PE KV MEMORY (left panel)
================================================================
KV bytes per token across all 80 layers, single KV group:
KV/tok = 2 (K+V) · h_kv · d_head · bytes · N_layers
= 2 · 1 · 128 · 2 · 80
= 40 KB / token
Per-PE share = KV/tok ÷ divisor, where divisor depends on the K/V
tensor placement (NOT on the compute-side label):
Case 1 Cube-SP × PE-replicate divisor = C = 8
Case 1' Cube-SP × PE-TP (on d_head) divisor = C·P = 64
Case 2 Cube-Repl × PE-replicate divisor = 1 = 1
Case 3 Cube-Repl × PE-SP divisor = P = 8
Case 3' Cube-TP × PE-SP divisor = C·P = 64
Case 4 Cube-SP × PE-SP ★ divisor = C·P = 64
Per-PE bytes @ S_kv = 1 M = KV/tok · 1 M / divisor:
Case 1 → 5.0 GB ✗ (exceeds 4.24 GB headroom)
Case 1' → 640 MB ✓
Case 2 → 40.0 GB ✗
Case 3 → 5.0 GB ✗
Case 3' → 640 MB ✓
Case 4 → 640 MB ✓
Max KV context per PE = 4.24 GB · 1024² / (KV/tok ÷ divisor):
Case 1, 3 : 889 K tokens
Case 1', 3', 4 : 7.11 M tokens ✓
Case 2 : 111 K tokens
================================================================
(2) PER-PE COMMUNICATION (right panel)
================================================================
(A) WEIGHT AllReduces (constant across all cases)
Wo AR ≈ 8 KB / layer / PE (Wq replicated → AR partial Y)
FFN AR ≈ 8 KB / layer / PE
× 80 layers = 1.25 MB / token / PE (the bottom blue stack on every bar)
(B) ATTENTION-TIME collective (the differentiator)
Cases that compute partial attention locally (PE-SP / PE-repl):
Q is replicated and each rank attends to its own complete KV slice
(full d_head) → only the small (m, , O) triple needs AllReducing.
Payload (h_q · S_q · d_head · 2) ≈ 4 KB per AR step, hierarchical ≈ 32 KB / layer.
CONSTANT in S_kv.
Case 1 : inter-cube AR on (m,,O) → ~32 KB / layer
Case 2 : nothing → 0
Case 3 : intra-cube AR on (m,,O) → ~32 KB / layer
Case 4 ★ : intra + inter-cube AR on (m,,O) → ~64 KB / layer
d_head-TP variants (1', 3'):
Each rank holds only d_head/divisor dims → Q·K^T produces only
partial sums → must AR PARTIAL SCORES before softmax. The score
tile per AR is (h_q · S_q · S_kv / slices_for_seq) · 2 bytes,
which SCALES with S_kv.
Case 1' : intra-cube partial-score AR
payload = h_q · S_q · (S_kv/C) · 2 = 2 MB / layer at 1 M
Case 3' : inter-cube partial-score AR (UCIe, slower than NoC)
payload = h_q · S_q · (S_kv/P) · 2 = 2 MB / layer at 1 M
Plus tiny (m,,O) AR for both 1' and 3' to combine partial attentions
across the remaining axis (~32 KB / layer).
(C) Total per output token per PE (sum over 80 layers, S_kv = 1 M):
Case 1 ≈ 4 MB (1.25 MB Wo+FFN + 2.5 MB inter-cube AR)
Case 1' ≈ 166 MB (1.25 MB + 160 MB partial-score AR + ~5 MB other)
Case 2 ≈ 1.2 MB (Wo+FFN only)
Case 3 ≈ 4 MB (1.25 MB + 2.5 MB intra-cube AR)
Case 3' ≈ 166 MB (1.25 MB + 160 MB partial-score AR + ~5 MB other)
Case 4 ★ ≈ 6 MB (1.25 MB + 5 MB 2-phase AR)
================================================================
KEY TAKEAWAYS
================================================================
- Memory winners (640 MB / PE @ 1M): Cases 1', 3', 4 (any 64-way sharding).
- Comm winner among those three: Case 4 (6 MB). Cases 1' and 3' both
pay ~160 MB because the d_head-sharded variant ARs a 2 MB partial-
score tile every layer — the score tensor scales with S_kv while
the (m, , O) triple does not.
- Cases 1 and 3 (8-way sharding) don't fit 1M context (5 GB vs 4.24 GB
headroom).
- Case 2 has the cheapest comm but the most memory (40 GB / PE).
- Case 4 ★ is the Pareto-best: fits 1M context AND lowest comm
among memory-feasible options.
Output PNG:
src/kernbench/benches/1H_milestone_output/gqa/long_ctx/
gqa_4cases_summary.png
"""
from __future__ import annotations
import json
from pathlib import Path
import matplotlib.patches as mpatches
import matplotlib.pyplot as plt
# ── System constants (LLaMA-3.1-70B single-KV-head group, decode) ───
_N_LAYERS = 80
_H_KV = 1
_H_Q = 8
_D_HEAD = 128
_D_MODEL = 8192
_BYTES_PER_ELEM = 2 # FP16
_C = 8 # cubes per KV group
_P = 8 # PEs per cube
_B = 1
_S_Q = 1 # decode
_HBM_PER_PE_GB = 6.0
_WEIGHTS_PER_PE_GB = 1.76
_HEADROOM_GB = _HBM_PER_PE_GB - _WEIGHTS_PER_PE_GB # 4.24 GB
_HEADLINE_S_KV = 1 << 20 # 1 Mi tokens
# Per-token KV bytes (single KV group, all 80 layers).
_KV_PER_TOK_BYTES = (
2 * _H_KV * _D_HEAD * _BYTES_PER_ELEM * _N_LAYERS
) # 40 KB
# Per-token Wo + FFN AR (constant across cases).
_WO_PER_LAYER_BYTES = 8 * 1024
_FFN_PER_LAYER_BYTES = 8 * 1024
# ── Per-PE attention weight breakdown (single KV-head group) ────────
#
# LLaMA-3.1-70B single-KV-head group dimensions:
# d_model = 8192
# h_q per group = 8 (8 query heads attend to 1 KV head per group)
# h_kv per group = 1
# d_head = 128
# FP16 (2 bytes)
#
# Per layer, per PE:
# Wq shape (d_model, h_q · d_head) REPLICATED across all 64 PEs of group
# = 8192 · 8·128 · 2 bytes = 16 MB / layer / PE
# Wk shape (d_model, h_kv · d_head) REPLICATED
# = 8192 · 1·128 · 2 bytes = 2 MB / layer / PE
# Wv shape (d_model, h_kv · d_head) REPLICATED
# = 8192 · 1·128 · 2 bytes = 2 MB / layer / PE
# Wo shape (h_q · d_head, d_model) ROW-SPLIT across C=8 cubes,
# replicated within cube
# = (1024/C) · 8192 · 2 bytes = 2 MB / layer / PE
# ────────────────────────────────────────────────────────────────
# total attn weights / layer / PE = 22 MB
# × 80 layers = 1.76 GB / PE (decimal GB)
#
# KV headroom per PE = HBM_PER_PE - weights = 6.0 - 1.76 = 4.24 GB.
# (FFN weights are accounted for in a separate budget, not in this 4.24 GB.)
# Slide-17 convention: per-layer values in binary MB (MiB), totals in
# "GB" formed by ×80 layers ÷ 1000 — giving the canonical 1.76 GB total
# and 4.24 GB headroom that the slide-17 chart reports.
_WQ_MB_PER_LAYER = (_D_MODEL * _H_Q * _D_HEAD * _BYTES_PER_ELEM) / (1024 ** 2) # 16
_WK_MB_PER_LAYER = (_D_MODEL * _H_KV * _D_HEAD * _BYTES_PER_ELEM) / (1024 ** 2) # 2
_WV_MB_PER_LAYER = (_D_MODEL * _H_KV * _D_HEAD * _BYTES_PER_ELEM) / (1024 ** 2) # 2
_WO_MB_PER_LAYER = ((_H_Q * _D_HEAD) // _C * _D_MODEL
* _BYTES_PER_ELEM) / (1024 ** 2) # 2
_WQ_GB = _WQ_MB_PER_LAYER * _N_LAYERS / 1000 # 1.28 GB
_WK_GB = _WK_MB_PER_LAYER * _N_LAYERS / 1000 # 0.16 GB
_WV_GB = _WV_MB_PER_LAYER * _N_LAYERS / 1000 # 0.16 GB
_WO_GB = _WO_MB_PER_LAYER * _N_LAYERS / 1000 # 0.16 GB
_WEIGHTS_GB = _WQ_GB + _WK_GB + _WV_GB + _WO_GB # 1.76 GB
# (m, , O) AR payload — used by Cases 1, 3, 4 (different axes per case).
_MLO_INTRA_BYTES_PER_LAYER = 32 * 1024 # ~32 KB intra-cube AR
_MLO_INTER_BYTES_PER_LAYER = 32 * 1024 # ~32 KB inter-cube AR
# Cases — renumbered in MEMORY-DESCENDING ORDER (left to right):
# Case 1 (40 GB) : no sharding
# Cases 2, 3 ( 5 GB) : single-axis sharding (cube OR PE only)
# Cases 4, 5, 6 (640 MB) : two-axis sharding (cube AND PE)
# — Case 6 ★ is the Pareto-best (lowest comm)
_CASES = (1, 2, 3, 4, 5, 6)
_DIVISOR = {
1: 1, # Cube-Repl × PE-replicate — no sharding
2: _C, # Cube-SP × PE-replicate — cube-axis only
3: _P, # Cube-Repl × PE-SP — PE-axis (S_kv) only
4: _C * _P, # Cube-SP × PE-TP (d_head) — 64-way (d_head intra)
5: _C * _P, # Cube-TP × PE-SP — 64-way (d_head inter)
6: _C * _P, # Cube-SP × PE-SP — 64-way (S_kv both axes) ★
}
_CASE_LABEL = {
1: "Case 1\nCube-Repl\nPE-repl",
2: "Case 2\nCube-SP\nPE-repl",
3: "Case 3\nCube-Repl\nPE-SP",
4: "Case 4\nCube-SP\nPE-TP",
5: "Case 5\nCube-TP\nPE-SP",
6: "Case 6 ★\nCube-SP\nPE-SP",
}
_CASE_COLOR = {
1: "#C0504D", # red — worst memory (no sharding)
2: "#E0834A", # orange — single-axis sharded (cube)
3: "#EBA854", # tan — single-axis sharded (PE)
4: "#A6C2E0", # light blue — d_head-TP 64-way (PE)
5: "#C7D8A0", # light green — d_head-TP 64-way (cube)
6: "#8064A2", # purple — Pareto winner ★ (S_kv 64-way)
}
_ATTN_DESC = {
1: "none",
2: "online-softmax\nmerge of\n(m,,O)\ninter-cube",
3: "online-softmax\nmerge of\n(m,,O)\nintra-cube",
4: "partial\nscores\n+ (m,,O)\nmerge\n(d_head-TP)",
5: "partial\nscores\n+ (m,,O)\nmerge\n(d_head-TP)",
6: "online-softmax\nmerge of\n(m,,O)\nintra + inter",
}
_WO_COLOR = "#9EC5E8"
_FFN_COLOR = "#4A78B8"
_ATTN_COLOR = "#E07A3F"
_OUT_DIR = (
Path(__file__).resolve().parents[2]
/ "src" / "kernbench" / "benches"
/ "1H_milestone_output" / "gqa" / "long_ctx"
)
_MEASURED_JSON = _OUT_DIR / "gqa_3cases_measured_comm.json"
def _load_measured() -> dict[int, float] | None:
"""Load measured per-PE comm bytes (already scaled to S_kv=1M).
Returns {case_id: per_token_total_bytes} or None if JSON missing.
Produced by scripts/paper/measure_gqa_decode_placement_comm.py.
"""
if not _MEASURED_JSON.exists():
return None
data = json.loads(_MEASURED_JSON.read_text())
return {
int(cid): info["per_pe_total_bytes_per_token_at_1M"]
for cid, info in data["cases"].items()
}
# ── Formulae ────────────────────────────────────────────────────────
def kv_per_pe_bytes(case, s_kv: int) -> int:
"""Per-PE KV bytes at the given S_kv."""
return _KV_PER_TOK_BYTES * s_kv // _DIVISOR[case]
def max_s_kv(case) -> int:
"""Max KV context per PE given the 4.24 GB headroom."""
headroom_bytes = int(_HEADROOM_GB * (1 << 30))
bytes_per_tok = _KV_PER_TOK_BYTES // _DIVISOR[case]
return headroom_bytes // bytes_per_tok
def _partial_score_bytes_per_layer(s_kv: int, slices_for_seq: int) -> int:
"""AR payload for partial scores in d_head-TP cases."""
return _H_Q * _S_Q * (s_kv // slices_for_seq) * _BYTES_PER_ELEM
def attn_comm_per_layer_bytes(case: int, s_kv: int) -> int:
"""Per-layer per-PE attention-time comm bytes (decode, B=1, S_q=1).
Case numbering follows the memory-descending order defined in
_CASES (1=no sharding, 6=Pareto winner).
"""
if case == 1: # Cube-Repl × PE-repl — no sharding
return 0
if case == 2: # Cube-SP × PE-repl — inter-cube AR on (m,,O)
return _MLO_INTER_BYTES_PER_LAYER
if case == 3: # Cube-Repl × PE-SP — intra-cube AR on (m,,O)
return _MLO_INTRA_BYTES_PER_LAYER
if case == 4: # Cube-SP × PE-TP(d_head) — partial-score AR
return (_partial_score_bytes_per_layer(s_kv, _C)
+ _MLO_INTRA_BYTES_PER_LAYER
+ _MLO_INTER_BYTES_PER_LAYER)
if case == 5: # Cube-TP(d_head) × PE-SP — partial-score AR
return (_partial_score_bytes_per_layer(s_kv, _P)
+ _MLO_INTRA_BYTES_PER_LAYER
+ _MLO_INTER_BYTES_PER_LAYER)
if case == 6: # Cube-SP × PE-SP — 2-phase (m,,O) AR ★
return _MLO_INTRA_BYTES_PER_LAYER + _MLO_INTER_BYTES_PER_LAYER
raise ValueError(f"unknown case {case}")
def per_token_bytes(case, s_kv: int) -> tuple[int, int, int]:
"""(Wo AR, FFN AR, Attn) bytes per output token per PE — × 80 layers."""
wo = _WO_PER_LAYER_BYTES * _N_LAYERS
ffn = _FFN_PER_LAYER_BYTES * _N_LAYERS
attn = attn_comm_per_layer_bytes(case, s_kv) * _N_LAYERS
return wo, ffn, attn
# ── Formatters ──────────────────────────────────────────────────────
def _fmt_bytes(b: float) -> str:
if b >= (1 << 30):
return f"{b / (1 << 30):.2f} GB"
if b >= (1 << 20):
return f"{b / (1 << 20):.1f} MB"
if b >= (1 << 10):
return f"{b / (1 << 10):.0f} KB"
return f"{b:.0f} B"
def _fmt_tokens(n: int) -> str:
if n >= 1_000_000:
return f"{n / 1_000_000:.2f} M"
if n >= 1_000:
return f"{n / 1_000:.0f} K"
return f"{n}"
# ── Panels ──────────────────────────────────────────────────────────
def _plot_budget(ax) -> None:
"""HBM budget per PE — stacked weights + KV headroom + HBM ceiling.
Wq is the dominant weight slice (~1.28 GB). Wk, Wv, Wo are each
small (~0.16 GB) so their slice labels would overlap on the bar —
they're shown in the legend only, and only Wq + KV-headroom get
on-bar annotations.
"""
components = [
("Wq (REPL)", _WQ_GB, "#7B9CC4"),
("Wk (REPL)", _WK_GB, "#A0BBD8"),
("Wv (REPL)", _WV_GB, "#C5D6E8"),
("Wo (cube-split)", _WO_GB, "#E2EAF3"),
("KV cache headroom",
_HBM_PER_PE_GB - _WEIGHTS_GB, "#9BBB59"),
]
bottom = 0.0
for label, val, color in components:
ax.bar(0, val, bottom=bottom, color=color, edgecolor="black",
width=0.7, label=f"{label} · {val:.2f} GB")
# Only annotate slices thick enough to fit text without overlap.
if val >= 0.50:
ax.text(0, bottom + val / 2, f"{label}\n{val:.2f} GB",
ha="center", va="center", fontsize=9, weight="bold")
bottom += val
ax.axhline(_HBM_PER_PE_GB, color="red", ls="--", lw=1.4,
label=f"HBM = {_HBM_PER_PE_GB} GB")
ax.set_xticks([0])
ax.set_xticklabels(["per-PE HBM"], fontsize=10)
ax.set_ylabel("GB per PE")
ax.set_ylim(0, _HBM_PER_PE_GB * 1.10)
ax.set_title(
f"Per-PE HBM budget\n"
f"weights {_WEIGHTS_GB:.2f} GB + KV = {_HBM_PER_PE_GB} GB",
fontsize=10,
)
ax.grid(axis="y", ls=":", alpha=0.5)
ax.legend(loc="upper right", fontsize=7.5, framealpha=0.92)
def _plot_memory(ax) -> None:
vals_gb = [kv_per_pe_bytes(c, _HEADLINE_S_KV) / (1 << 30) for c in _CASES]
labels = [_CASE_LABEL[c] for c in _CASES]
colors = [_CASE_COLOR[c] for c in _CASES]
x = list(range(len(_CASES)))
bars = ax.bar(x, vals_gb, color=colors, width=0.65)
ax.set_xticks(x)
ax.set_xticklabels(labels, fontsize=9)
ax.set_ylabel("KV bytes per PE (GB, log)")
ax.set_yscale("log")
ax.set_title(
f"Per-PE KV memory at S_kv = {_HEADLINE_S_KV:,} tokens "
f"(across {_N_LAYERS} layers, FP16)",
fontsize=11,
)
ax.grid(axis="y", ls=":", alpha=0.5, which="both")
ax.axhline(_HEADROOM_GB, color="red", ls="--", lw=1.4,
label=f"KV headroom = {_HEADROOM_GB} GB / PE")
for bar, v_gb in zip(bars, vals_gb):
v_bytes = v_gb * (1 << 30)
fits = v_gb <= _HEADROOM_GB
ax.text(bar.get_x() + bar.get_width() / 2,
v_gb * 1.10,
_fmt_bytes(v_bytes) + ("" if fits else ""),
ha="center", va="bottom", fontsize=9,
color="green" if fits else "red", weight="bold")
ax.legend(loc="upper right", fontsize=9)
def _plot_comm(ax, *, mode: str = "analytical") -> None:
"""Per-PE comm panel.
mode = "analytical": single solid bars from per_token_bytes formula.
mode = "paired" : analytical (solid) + simulator-measured
(hatched) side-by-side per case, when the
measurement JSON is available.
"""
labels = [_CASE_LABEL[c] for c in _CASES]
wo_mb_list: list[float] = []
ffn_mb_list: list[float] = []
attn_mb: list[float] = []
for c in _CASES:
wo, ffn, attn = per_token_bytes(c, _HEADLINE_S_KV)
wo_mb_list.append(wo / (1 << 20))
ffn_mb_list.append(ffn / (1 << 20))
attn_mb.append(attn / (1 << 20))
measured = _load_measured() if mode == "paired" else None
paired = measured is not None
source_tag = "analytical (solid) vs simulator-measured (hatched)" \
if paired else "analytical"
n_cases = len(_CASES)
x = list(range(n_cases))
bar_w = 0.36 if paired else 0.65
x_ana = [xi - bar_w / 2 for xi in x] if paired else x
x_meas = [xi + bar_w / 2 for xi in x] if paired else None
# Analytical bars (solid).
ax.bar(x_ana, wo_mb_list, width=bar_w,
color=_WO_COLOR, edgecolor="black")
ax.bar(x_ana, ffn_mb_list, width=bar_w, bottom=wo_mb_list,
color=_FFN_COLOR, edgecolor="black")
bottoms_attn = [w + f for w, f in zip(wo_mb_list, ffn_mb_list)]
ax.bar(x_ana, attn_mb, width=bar_w, bottom=bottoms_attn,
color=_ATTN_COLOR, edgecolor="black")
# Measured bars (hatched) — same Wo+FFN base, attn from op_log.
meas_attn_mb: list[float] = []
if paired:
for i, c in enumerate(_CASES):
meas_total = measured.get(c, 0) / (1 << 20)
meas_attn_mb.append(
max(meas_total - wo_mb_list[i] - ffn_mb_list[i], 0.0))
ax.bar(x_meas, wo_mb_list, width=bar_w,
color=_WO_COLOR, edgecolor="black",
hatch="///", alpha=0.85)
ax.bar(x_meas, ffn_mb_list, width=bar_w, bottom=wo_mb_list,
color=_FFN_COLOR, edgecolor="black",
hatch="///", alpha=0.85)
ax.bar(x_meas, meas_attn_mb, width=bar_w, bottom=bottoms_attn,
color=_ATTN_COLOR, edgecolor="black",
hatch="///", alpha=0.85)
ax.set_xticks(x)
ax.set_xticklabels(labels, fontsize=9)
ax.set_ylabel("Comm per token per PE (MB, log)")
ax.set_yscale("log")
title = (
f"Per-PE communication at S_kv = {_HEADLINE_S_KV:,} tokens "
f"(decode S_q=1, B=1; {_N_LAYERS} layers) — {source_tag}"
)
if paired:
title += "\n(simulator measured at S_kv = 8K; " \
"partial-score AR scaled ×128 to S_kv = 1M)"
ax.set_title(title, fontsize=10)
ax.grid(axis="y", ls=":", alpha=0.5, which="both")
totals_ana = [wo_mb_list[i] + ffn_mb_list[i] + attn_mb[i]
for i in range(n_cases)]
ymax = max(totals_ana)
if paired:
ymax = max(ymax, max(
(measured.get(c, 0) / (1 << 20)) for c in _CASES))
ax.set_ylim(top=ymax * 22)
for i, c in enumerate(_CASES):
ana_bytes = totals_ana[i] * (1 << 20)
if paired:
meas_bytes = measured.get(c, 0)
top_y = max(totals_ana[i], meas_bytes / (1 << 20))
label = (f"ana: {_fmt_bytes(ana_bytes)}\n"
f"sim: {_fmt_bytes(meas_bytes)}")
else:
top_y = totals_ana[i]
label = _fmt_bytes(ana_bytes)
ax.text(x[i], top_y * 1.5, label,
ha="center", va="bottom",
fontsize=8 if paired else 9,
weight="bold", linespacing=1.05)
# In-bar attention-time AR description (online-softmax (m,,O)
# merge, partial-score AR, intra/inter-cube). Placed on the
# analytical bar; with paired mode it lives on the solid bar so
# the hatched measured bar stays uncluttered.
if attn_mb[i] > 0:
mid = bottoms_attn[i] + attn_mb[i] / 2
ax.text(x_ana[i], mid, _ATTN_DESC[c],
ha="center", va="center",
fontsize=6 if paired else 7,
color="black", weight="bold")
else:
ax.text(x_ana[i], totals_ana[i] * 0.4, _ATTN_DESC[c],
ha="center",
fontsize=6.5 if paired else 7.5,
color="grey", style="italic")
legend_handles = [
mpatches.Patch(facecolor=_WO_COLOR, edgecolor="black",
label=f"Wo AR (× {_N_LAYERS} layers)"),
mpatches.Patch(facecolor=_FFN_COLOR, edgecolor="black",
label=f"FFN AR (× {_N_LAYERS} layers)"),
mpatches.Patch(facecolor=_ATTN_COLOR, edgecolor="black",
label="Attn-time collective"),
]
if paired:
legend_handles.append(
mpatches.Patch(facecolor="white", edgecolor="black",
hatch="///", label="simulator-measured"))
ax.legend(handles=legend_handles, loc="upper right", fontsize=8,
framealpha=0.92)
def main() -> Path:
_OUT_DIR.mkdir(parents=True, exist_ok=True)
# (a) Per-PE HBM budget — standalone PNG.
fig_b, ax_b = plt.subplots(figsize=(4.0, 6.0))
_plot_budget(ax_b)
fig_b.tight_layout()
out_b = _OUT_DIR / "gqa_hbm_budget.png"
fig_b.savefig(out_b, dpi=150)
plt.close(fig_b)
print(f"wrote {out_b}")
# (b) Combined 3-panel summary — HBM budget + KV memory + comm.
fig = plt.figure(figsize=(22.0, 6.5))
gs = fig.add_gridspec(1, 3, width_ratios=[0.7, 1.6, 1.6], wspace=0.22)
ax_b2 = fig.add_subplot(gs[0, 0])
ax_m = fig.add_subplot(gs[0, 1])
ax_c = fig.add_subplot(gs[0, 2])
_plot_budget(ax_b2)
_plot_memory(ax_m)
_plot_comm(ax_c)
fig.tight_layout()
out = _OUT_DIR / "gqa_4cases_summary.png"
fig.savefig(out, dpi=150)
plt.close(fig)
print(f"wrote {out}")
# (c) 2-panel companion (analytical only).
fig2 = plt.figure(figsize=(18.0, 6.5))
gs2 = fig2.add_gridspec(1, 2, width_ratios=[1.0, 1.0], wspace=0.18)
ax_m2 = fig2.add_subplot(gs2[0, 0])
ax_c2 = fig2.add_subplot(gs2[0, 1])
_plot_memory(ax_m2)
_plot_comm(ax_c2, mode="analytical")
fig2.tight_layout()
out2 = _OUT_DIR / "gqa_4cases_memory_comm_analytical.png"
fig2.savefig(out2, dpi=150)
plt.close(fig2)
print(f"wrote {out2}")
# (d) 2-panel companion — analytical vs simulator-measured paired.
fig3 = plt.figure(figsize=(19.0, 6.5))
gs3 = fig3.add_gridspec(1, 2, width_ratios=[1.0, 1.0], wspace=0.18)
ax_m3 = fig3.add_subplot(gs3[0, 0])
ax_c3 = fig3.add_subplot(gs3[0, 1])
_plot_memory(ax_m3)
_plot_comm(ax_c3, mode="paired")
fig3.tight_layout()
out3 = _OUT_DIR / "gqa_4cases_memory_comm_paired.png"
fig3.savefig(out3, dpi=150)
plt.close(fig3)
print(f"wrote {out3}")
# Paper-ready table to stdout.
print()
print(f" {'Case':<27} {'KV/tok·PE':>12} {'KV @ 1M':>11} "
f"{'Max S_kv':>11} {'Comm @ 1M':>12}")
print(" " + "-" * 80)
for c in _CASES:
kv_per_tok_pe = _KV_PER_TOK_BYTES / _DIVISOR[c]
mem_1m = kv_per_pe_bytes(c, _HEADLINE_S_KV)
max_s = max_s_kv(c)
comm_1m = sum(per_token_bytes(c, _HEADLINE_S_KV))
label = _CASE_LABEL[c].replace(chr(10), " ")
print(f" {label:<27} "
f"{kv_per_tok_pe / 1024:>9.3f} KB "
f"{_fmt_bytes(mem_1m):>11} "
f"{_fmt_tokens(max_s):>11} "
f"{_fmt_bytes(comm_1m):>12}")
return out
if __name__ == "__main__":
main()