Files
kernbench2/tests/analytical_visualization/chip_roofline.py
T
mukesh 221097ef08 analytical-viz: roofline tab — headline step-latency + cost-per-token plots, B knob
Two new plots at the top of the Chip Roofline tab (right after the
knob row), side by side:

1. **Latency per decode step** — one forward pass's total time as B
   grows. weight_fetch (flat in B), compute (linear), KV (linear),
   total. The SLO / TTFT view — bigger B = longer step.

2. **Cost per token** = step latency ÷ B. weight (shrinks 1/B),
   compute (flat), KV (flat), total. The efficiency / cost view —
   bigger B (up to ~2·B*) = cheaper per token.

Same underlying decomposition, two divisors. Puts the classic
throughput ↔ latency tradeoff in one glance.

Also added a **B (batch size) slider** next to the existing S_kv
slider. Both plots draw a vertical marker at the current B; the
regime KPI ("memory-bound / compute-bound / kv-bound") now uses the
slider B, so you can sweep B live and watch the label flip when you
cross B*.

step_latency_curve + StepLatencyPoint added to chip_roofline; 5 new
tests cover: weight_s batch-invariant, compute_s/kv_s linear in B,
step_total == per_token_total × B (definitional), and step ==
per_token at B=1.

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

307 lines
11 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)
# ── Per-step latency (undivided by B) ─────────────────────────────
@dataclass
class StepLatencyPoint:
batch: int
weight_s: float # N·b / W — flat in B (loaded once per step)
compute_s: float # 2·N·B / C — linear in B
kv_s: float # B · S_kv · kv_bpt / W — linear in B
total_s: float
def step_latency_curve(machine: MachineParams, model: ModelConfig,
batch_range: list[int],
s_kv: int) -> list[StepLatencyPoint]:
"""Total time of one decode step (one forward pass), across a
range of batch sizes. NOT divided by B — this is the SLO view.
step_weight = N·b / W (batch-invariant)
step_compute = 2·N·B / C (linear in B)
step_kv = B · S_kv · kv_bpt / W (linear in B)
Per-token cost = step_total / B — the two views trade off:
bigger B lowers cost/token but raises step latency.
"""
n_active = total_active_params(model)
b = model.bytes_per_elem
weight_bytes = n_active * b
kv_bpt = kv_bytes_per_token(model)
step_weight = weight_bytes / machine.bw_hbm
points: list[StepLatencyPoint] = []
for bs in batch_range:
B = max(1, bs)
step_compute = _FLOPS_PER_PARAM_PER_TOKEN * n_active * B / machine.peak_flops
step_kv = B * s_kv * kv_bpt / machine.bw_hbm
total = step_weight + step_compute + step_kv
points.append(StepLatencyPoint(
batch=bs,
weight_s=step_weight,
compute_s=step_compute,
kv_s=step_kv,
total_s=total,
))
return points