Files
kernbench2/docs/sweeps/plot_scripts/plot_short_context.py
T
ywkang 711a9a257f gqa short-context: exclude KV deploy from wall; correct HBM peak; batch scaffold
Metric fixes (test harnesses, deploy artifact + wrong constant):
- wall = max(t_end) - min(t_start): exclude the one-time KV-cache deploy
  from the measured decode/prefill step (it was 92-99% of wall, mapping-
  invariant, and masked the real per-mapping separation).
- PEAK_PE_HBM_BPS 128 -> 256 B/ns (8 channels/PE x 32 GB/s); the old value
  was half the modeled per-PE HBM BW, so hbm_bw_util read >1.0 once wall
  was corrected. All six short-context sweep CSVs regenerated/repatched.

Result: the four KV mappings now separate along a {64,32,16,8}-active-PE
ladder (decode 8-kv/1-kv = 4.8x at 8K, 7.4x at 64K), not "modest/tied" as
before; decode is bandwidth-bound at 46-76% of the 256 GB/s per-PE ceiling.

Report (S5.2 rewrite):
- Replace the tied-wall / 8x-per-PE-util claims (both deploy artifacts)
  with the corrected separation and a density trade-off (per-CUBE KV
  footprint, Fig 16 top panel switched to a wall-invariant metric).
- Add a projected latency-vs-batched-throughput analysis (marked
  projected, not measured): dense mappings win throughput, 1-kv wins
  latency; converges at long context.
- Regenerate Fig 15/16/17/18; fix plot script hardcoded ROOT path.

Batch experiment (Part 2, cube_base):
- Add backward-compatible cube_base=0 scalar to the decode kernel so a
  batched user placed at DPPolicy.cube_start addresses 0-based shards.
  Default preserves single-user behavior (32 decode tests pass unchanged).
- New batch harness (skipped): concurrent B>=2 launches hit a sim routing
  issue (sip0.cube0.pe8); single-user path verified. Concurrency fix next.

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
2026-07-22 16:25:04 -07:00

232 lines
8.8 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.
"""Generate 6 comparison figures for the short-context GQA attention sweep.
Sweep matrix: 3 variants × 2 phases × 4 modes × 4 contexts (8K/16K/32K/64K).
Figures
Wall-clock latency:
1. wall_variant1_mode_compare.png — (1) without composite
2. wall_variant2_mode_compare.png — (2) with composite (GEMM-only)
3. wall_variant3_mode_compare.png — (3) with composite + softmax_merge
4. wall_a1_variant_compare.png — 1-kv-per-cube: variant comparison
Mode trade-off + composite ablation:
5. per_cube_tradeoff_mode_compare.png — HBM BW + IPCQ (variant-invariant)
6. gemm_util_a1_variant_ablation.png — 1-kv-per-cube GEMM util
Output: docs/report/1H-codesign-paper/figures/gqa_short_context/
"""
from __future__ import annotations
import csv
from pathlib import Path
import matplotlib.pyplot as plt
import numpy as np
ROOT = Path(__file__).resolve().parents[3]
SWEEPS = ROOT / "docs" / "sweeps"
OUT = (ROOT / "docs" / "report" / "1H-codesign-paper"
/ "figures" / "gqa_short_context")
OUT.mkdir(parents=True, exist_ok=True)
CONTEXTS = [8192, 16384, 32768, 65536]
CONTEXT_LABELS = ["8K", "16K", "32K", "64K"]
MODES = ["A1", "A2", "A4", "B"]
MODE_LABEL = {
"A1": "1-kv-per-cube",
"A2": "2-kv-per-cube",
"A4": "4-kv-per-cube",
"B": "8-kv-per-cube",
}
MODE_COLOR = {
MODE_LABEL["A1"]: "#1f77b4",
MODE_LABEL["A2"]: "#2ca02c",
MODE_LABEL["A4"]: "#ffd000",
MODE_LABEL["B"]: "#d62728",
}
VARIANTS = [
("baseline", "(1) without composite"),
("composite", "(2) with composite (GEMM-only)"),
("composite_fused", "(3) with composite + softmax_merge"),
]
VARIANT_COLOR = {
"baseline": "#1f77b4",
"composite": "#2ca02c",
"composite_fused": "#d62728",
}
CSV_MAP = {
("decode", "baseline"): SWEEPS / "short_context_decode_sweep.csv",
("decode", "composite"): SWEEPS / "short_context_decode_composite_sweep.csv",
("decode", "composite_fused"): SWEEPS / "short_context_decode_composite_fused_sweep.csv",
("prefill", "baseline"): SWEEPS / "short_context_prefill_sweep.csv",
("prefill", "composite"): SWEEPS / "short_context_prefill_composite_sweep.csv",
("prefill", "composite_fused"): SWEEPS / "short_context_prefill_composite_fused_sweep.csv",
}
def load(phase, variant):
"""Return {(mode, S_kv): row dict}."""
rows = {}
with CSV_MAP[(phase, variant)].open() as f:
for r in csv.DictReader(f):
rows[(r["mode"], int(r["S_kv"]))] = r
return rows
def grouped_bars(ax, *, x_labels, groups, group_color, ylabel, title,
y_log=False, value_fmt="{:.3g}"):
"""Grouped bars: x = x_labels, groups = list of (label, values)."""
n = len(groups)
nx = len(x_labels)
width = 0.8 / n
x = np.arange(nx)
for i, (label, values) in enumerate(groups):
offset = (i - (n - 1) / 2) * width
bars = ax.bar(x + offset, values, width, label=label,
color=group_color[label], edgecolor="black",
linewidth=0.4)
for bar, v in zip(bars, values):
if v <= 0 and y_log:
continue
ax.text(bar.get_x() + bar.get_width() / 2, bar.get_height(),
value_fmt.format(v), ha="center", va="bottom",
fontsize=7)
ax.set_xticks(x)
ax.set_xticklabels(x_labels)
ax.set_xlabel("S_kv")
ax.set_ylabel(ylabel)
ax.set_title(title)
if y_log:
ax.set_yscale("log")
ax.grid(True, axis="y", alpha=0.3, which="both")
ax.legend(fontsize=8, loc="upper left")
# ── 1-3. Wall mode comparison per variant ──────────────────────────
def fig_wall_mode_compare(variant_key, suptitle, fname):
fig, axes = plt.subplots(1, 2, figsize=(13, 4.8))
for ax, phase in zip(axes, ("prefill", "decode")):
rows = load(phase, variant_key)
groups = [
(MODE_LABEL[m], [float(rows[(m, s)]["wall_us"]) for s in CONTEXTS])
for m in MODES
]
grouped_bars(ax, x_labels=CONTEXT_LABELS, groups=groups,
group_color=MODE_COLOR,
ylabel="Wall-clock latency (μs)",
title=phase.capitalize(),
y_log=True, value_fmt="{:.0f}")
fig.suptitle(suptitle, fontsize=12, fontweight="bold")
fig.tight_layout()
out = OUT / fname
fig.savefig(out, dpi=140, bbox_inches="tight")
plt.close(fig)
print(f"{out}")
# ── 4. Wall 1-kv-per-cube variant comparison ───────────────────────
def fig_wall_a1_variant_compare(fname):
fig, axes = plt.subplots(1, 2, figsize=(13, 4.8))
for ax, phase in zip(axes, ("prefill", "decode")):
groups = []
for variant_key, _ in VARIANTS:
rows = load(phase, variant_key)
values = [float(rows[("A1", s)]["wall_us"]) for s in CONTEXTS]
groups.append((variant_key, values))
grouped_bars(ax, x_labels=CONTEXT_LABELS, groups=groups,
group_color=VARIANT_COLOR,
ylabel="Wall-clock latency (μs)",
title=phase.capitalize(),
y_log=True, value_fmt="{:.0f}")
ax.legend([v[1] for v in VARIANTS], fontsize=8, loc="upper left")
fig.suptitle("1-kv-per-cube: variant comparison",
fontsize=12, fontweight="bold")
fig.tight_layout()
out = OUT / fname
fig.savefig(out, dpi=140, bbox_inches="tight")
plt.close(fig)
print(f"{out}")
# ── 5. Per-cube trade-off (HBM BW + IPCQ, mode compare) ────────────
def fig_per_cube_tradeoff(fname):
fig, axes = plt.subplots(2, 2, figsize=(13, 9))
metrics = [
("kv_cache_per_cube_mb", "Per-cube KV footprint (MB)", True, "{:.0f}"),
("ipcq_kb", "IPCQ traffic (KB)", True, "{:.1f}"),
]
for row, (key, ylabel, log, fmt) in enumerate(metrics):
for col, phase in enumerate(("prefill", "decode")):
ax = axes[row][col]
rows = load(phase, "baseline")
groups = [
(MODE_LABEL[m], [float(rows[(m, s)][key]) for s in CONTEXTS])
for m in MODES
]
grouped_bars(ax, x_labels=CONTEXT_LABELS, groups=groups,
group_color=MODE_COLOR, ylabel=ylabel,
title=phase.capitalize(),
y_log=log, value_fmt=fmt)
fig.suptitle("Per-cube trade-off: HBM BW + IPCQ mode comparison",
fontsize=12, fontweight="bold")
fig.tight_layout()
out = OUT / fname
fig.savefig(out, dpi=140, bbox_inches="tight")
plt.close(fig)
print(f"{out}")
# ── 6. GEMM util ablation (1-kv-per-cube, variant compare) ─────────
def fig_gemm_util_a1_ablation(fname):
fig, axes = plt.subplots(1, 2, figsize=(13, 5.4))
for ax, phase in zip(axes, ("prefill", "decode")):
groups = []
for variant_key, _ in VARIANTS:
rows = load(phase, variant_key)
values = [float(rows[("A1", s)]["gemm_util"]) for s in CONTEXTS]
groups.append((variant_key, values))
grouped_bars(ax, x_labels=CONTEXT_LABELS, groups=groups,
group_color=VARIANT_COLOR,
ylabel="GEMM engine utilization (per PE)",
title=phase.capitalize(),
y_log=False, value_fmt="{:.4f}")
ax.legend([v[1] for v in VARIANTS], fontsize=8, loc="upper left")
fig.suptitle("1-kv-per-cube: GEMM utilization variant comparison",
fontsize=12, fontweight="bold")
fig.text(
0.5, -0.04,
"⚠ Higher gemm_util on composite/fused is largely supertile "
"padding overhead (M=G=8 padded to TILE_M=32), not pure fusion "
"gain. See ADR-0070 limitation #4.",
ha="center", fontsize=9, style="italic", color="#555")
fig.tight_layout()
out = OUT / fname
fig.savefig(out, dpi=140, bbox_inches="tight")
plt.close(fig)
print(f"{out}")
if __name__ == "__main__":
print(f"Output: {OUT}")
fig_wall_mode_compare(
"baseline",
"(1) Without composite: latency mode comparison",
"wall_variant1_mode_compare.png")
fig_wall_mode_compare(
"composite",
"(2) With composite (GEMM-only): latency mode comparison",
"wall_variant2_mode_compare.png")
fig_wall_mode_compare(
"composite_fused",
"(3) With composite + softmax_merge: latency mode comparison",
"wall_variant3_mode_compare.png")
fig_wall_a1_variant_compare("wall_a1_variant_compare.png")
fig_per_cube_tradeoff("per_cube_tradeoff_mode_compare.png")
fig_gemm_util_a1_ablation("gemm_util_a1_variant_ablation.png")
print("\nDone.")