Files
kernbench2/tests/analytical_visualization/chip_roofline.py
T
mukesh c99a238826 analytical-viz: roofline tab — S_kv knob, regime formulas, good B/L, decomp plot, split table
Bundle of roofline-tab enhancements requested in-thread:

1. **S_kv slider at top of tab** — override the sidebar's S_kv locally
   so all four plots update as you drag it. Handy for exploring how
   the KV wall arrives without disturbing the rest of the app config.

2. **Good B / Good L KPI cards** — 3 metrics under the main KPI row:
     - Good B = 2·B* (Pope's rule of thumb)
     - Good L ceiling = L* (compute-friendly context ceiling)
     - Utilization @ current S_kv = 1 / (1 + S_kv/L*)

3. **Plot 4 — latency-decomposition per decode step** — five curves
   on one axis:
     - Compute (dashed, flat)
     - Weight fetch (triangles, shrinks 1/B)
     - KV fetch (squares, flat — batching doesn't help)
     - Memory total (weights + KV, purple)
     - Total (compute + memory, black bold)
   Makes "which term dominates at this B?" visible at a glance.

4. **Regime formulas table** — one row per cost term (t_com, weight
   fetch, KV fetch, bottleneck) × two columns (short-context vs
   long-context regime), plus a 'value now' column using the
   current S_kv slider.

5. **How to pick B and S_kv — one-paragraph guidance** with the
   numeric recommendations plugged in.

6. **Split formula table** — was cramming symbolic + substituted
   form into one cell with newlines; now has four proper columns:
   Symbol / Formula / With numbers / Meaning / Value.

New pure functions in chip_roofline.py: t_mem_short, t_mem_long,
t_com, good_batch, good_context, utilization_at. All roofline math
stays in the pure module; app.py just calls them and formats.

10 new tests bring the roofline test count to 27 (all 67 tests still
green): t_mem_long(L*) == t_com, doubling B halves t_mem_short,
good_batch scales with sparsity, utilization at 0/L*/2L* returns
1.0/0.5/0.333, monotonic decrease with context.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-07-29 13:27:17 -07:00

259 lines
9.4 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.
"""Chip-level roofline math: AI, B*, L*, per-token latency curves.
Surfaces the arithmetic-intensity story from LLM-serving practice:
- **AI** = C / W (peak FLOPs per byte of HBM bandwidth).
- **B\*** = C * b / (2 * W) * sparsity — the critical batch size at
which weight-fetch time equals compute time for one decode step.
Sparsity = N_total / N_active (MoE factor; 1 for dense).
- **L\*** = 2 * N_active / (AI * kv_bytes_per_token) — the balance
context length at which KV-read time equals compute time.
- **B_knee(S_kv)** = B* / (1 - S_kv/L*) — the batch size where the
cost curve bends (weight-fetch drops below the compute+KV floor).
Diverges at S_kv = L* and no knee exists past it.
Per-token decode-step latency, per PE (dense-approx, no comm):
t(B, S_kv) = N_active * b / (W * B) <-- weight fetch, 1/B
+ 2 * N_active / C <-- compute (peak), flat
+ S_kv * kv_bpt / W <-- KV read, flat
All numbers per PE / per one forward pass. **Peak roofline — no
utilization factor.** Comm cost and TP/CP sharding are intentionally
NOT in the roofline — this is the back-of-envelope chip-vs-model view
the transcript talks about, not the full latency model that
stage_latencies.py builds. With this convention weight_s == compute_s
exactly at B*.
"""
from __future__ import annotations
from dataclasses import dataclass
from .model_config import MachineParams, ModelConfig
# Per-token bytes of BF16 MAC arithmetic: one multiply + one add = 2 FLOPs.
_FLOPS_PER_PARAM_PER_TOKEN = 2
# ── Chip / model derived quantities ────────────────────────────────
def arithmetic_intensity(machine: MachineParams) -> float:
"""FLOPs per byte of HBM bandwidth. Peak roofline; utilization
is applied only in the compute-time formula, not here."""
return machine.peak_flops / machine.bw_hbm
def total_active_params(model: ModelConfig) -> int:
"""Full-model parameter count (attention + FFN, all layers).
Attention: 4 projections × hidden × H_q * d_head effective per layer.
(W_Q hidden×H_q*d_h, W_O H_q*d_h×hidden, W_K/W_V hidden×H_kv*d_h.)
FFN: 3 × hidden × ffn_dim per layer (gate, up, down).
"""
m = model
attn = (
m.hidden * m.h_q * m.d_head # W_Q
+ m.hidden * m.h_kv * m.d_head * 2 # W_K + W_V
+ m.h_q * m.d_head * m.hidden # W_O
)
ffn = 3 * m.hidden * m.ffn_dim
return (attn + ffn) * m.layers
def kv_bytes_per_token(model: ModelConfig) -> int:
"""Bytes of KV cache one new token adds across ALL layers, per one
sequence, un-sharded (K + V, H_kv heads * d_h * bytes)."""
m = model
return 2 * m.h_kv * m.d_head * m.bytes_per_elem * m.layers
def critical_batch(machine: MachineParams, model: ModelConfig,
sparsity: float = 1.0) -> float:
"""B* = C * b / (2 * W) * sparsity.
Sparsity = N_total / N_active (>= 1). Dense = 1. MoE 8-of-256 = 8.
"""
b = model.bytes_per_elem
ai = arithmetic_intensity(machine)
return ai * b / _FLOPS_PER_PARAM_PER_TOKEN * sparsity
def balance_context(machine: MachineParams, model: ModelConfig) -> float:
"""L* = 2 * N_active / (AI * kv_bpt).
Context length (in tokens) at which per-step KV read matches the
per-step compute cost. Beyond L*, the KV term is dominant and no
batch size gets you compute-bound.
"""
n_active = total_active_params(model)
ai = arithmetic_intensity(machine)
kv_bpt = kv_bytes_per_token(model)
return _FLOPS_PER_PARAM_PER_TOKEN * n_active / (ai * kv_bpt)
def knee_batch(machine: MachineParams, model: ModelConfig,
s_kv: int) -> float | None:
"""B_knee(S_kv) = B* / (1 - S_kv/L*).
Returns None when S_kv >= L* (no knee exists — the total-cost
curve never touches the compute floor).
"""
b_star = critical_batch(machine, model)
l_star = balance_context(machine, model)
r = s_kv / l_star
if r >= 1.0:
return None
return b_star / (1.0 - r)
# ── Per-token latency curves ───────────────────────────────────────
@dataclass
class RooflinePoint:
batch: int
weight_s: float # weight fetch time, 1/B
compute_s: float # compute time, flat
kv_s: float # KV read time, flat
total_s: float
def per_token_latency_curve(machine: MachineParams, model: ModelConfig,
batch_range: list[int],
s_kv: int) -> list[RooflinePoint]:
"""Per-token decode-step latency curve across a range of batch sizes.
Returns one point per batch. All times per PE, dense-approx,
utilization from machine.compute_util. Comm and TP/CP sharding are
excluded — this is the roofline model.
"""
n_active = total_active_params(model)
b = model.bytes_per_elem
weight_bytes = n_active * b
compute_flops = _FLOPS_PER_PARAM_PER_TOKEN * n_active
kv_read_bytes = s_kv * kv_bytes_per_token(model)
compute_s = compute_flops / machine.peak_flops
kv_s = kv_read_bytes / machine.bw_hbm
points: list[RooflinePoint] = []
for bs in batch_range:
weight_s = weight_bytes / machine.bw_hbm / max(1, bs)
points.append(RooflinePoint(
batch=bs,
weight_s=weight_s,
compute_s=compute_s,
kv_s=kv_s,
total_s=weight_s + compute_s + kv_s,
))
return points
def bound_regime(machine: MachineParams, model: ModelConfig,
batch: int, s_kv: int) -> str:
"""Which term dominates at the current (batch, S_kv) point.
Returns 'memory-bound' if weight_fetch is the largest term,
'kv-bound' if KV read is largest, 'compute-bound' if compute.
"""
pts = per_token_latency_curve(machine, model, [batch], s_kv)
p = pts[0]
parts = {"memory-bound": p.weight_s,
"kv-bound": p.kv_s,
"compute-bound": p.compute_s}
return max(parts, key=parts.get)
# ── Regime-dependent cost terms ────────────────────────────────────
def t_mem_short(machine: MachineParams, model: ModelConfig,
batch: int) -> float:
"""Per-token weight-fetch time (short-context regime term).
N_active · b / (W · B). Shrinks as B grows — this is what
batching amortizes.
"""
return (total_active_params(model) * model.bytes_per_elem
/ (machine.bw_hbm * max(1, batch)))
def t_mem_long(machine: MachineParams, model: ModelConfig,
s_kv: int) -> float:
"""Per-token KV-read time (long-context regime term).
S_kv · kv_bpt / W. Independent of B — each sequence reads its
own KV cache; batching doesn't help.
"""
return s_kv * kv_bytes_per_token(model) / machine.bw_hbm
def t_com(machine: MachineParams, model: ModelConfig) -> float:
"""Per-token compute time. Same in both regimes: 2·N/C, peak."""
return _FLOPS_PER_PARAM_PER_TOKEN * total_active_params(model) / machine.peak_flops
# ── "Good" batch / context recommendations ─────────────────────────
@dataclass
class BatchRecommendation:
target: float # Pope's rule: 2 × B*
b_star: float # B* itself
effective: float # what we recommend using
reason: str # short explanation
@dataclass
class ContextRecommendation:
l_star: float # balance context length
max_efficient: float # same as l_star (compute-friendly ceiling)
utilization_at: float # utilization at current s_kv
reason: str
def good_batch(machine: MachineParams, model: ModelConfig,
sparsity: float = 1.0) -> BatchRecommendation:
"""Recommended batch size: 2 × B* (Pope's rule of thumb).
Below B*: memory-bound, doubling B halves cost/token.
At 2×B*: 50% excess over compute floor — the sweet spot.
Beyond 3×B*: diminishing returns; latency keeps growing linearly.
"""
b_star = critical_batch(machine, model, sparsity)
target = 2 * b_star
return BatchRecommendation(
target=target, b_star=b_star, effective=target,
reason=(f"2·B* = 2 · {b_star:.0f} = {target:.0f}. "
"Below B*: memory-bound (doubling B halves cost/token). "
"Beyond 3·B*: diminishing returns."),
)
def good_context(machine: MachineParams, model: ModelConfig,
s_kv: int) -> ContextRecommendation:
"""Recommended max context: L* — the compute-friendly ceiling.
Below L*: KV read is cheap relative to compute → good utilization.
Above L*: KV bandwidth wall → utilization = 1/(1 + S_kv/L*).
"""
l_star = balance_context(machine, model)
util = utilization_at(s_kv, l_star)
return ContextRecommendation(
l_star=l_star, max_efficient=l_star, utilization_at=util,
reason=(f"L* = {l_star:,.0f} tokens. Below L*: compute-bound "
f"(good util). At {s_kv:,} tokens: peak utilization ≈ "
f"{util*100:.1f}% (1 / (1 + S_kv/L*))."),
)
def utilization_at(s_kv: int, l_star: float) -> float:
"""Peak compute utilization at context length s_kv, given L*.
util = compute / (compute + KV_read) = 1 / (1 + S_kv/L*).
At S_kv=0: 100%. At S_kv=L*: 50%. At 2·L*: 33.3%. At 5·L*: 16.7%.
"""
return 1.0 / (1.0 + s_kv / l_star)