9fdde44922
New tests/analytical_visualization/ module - an interactive dashboard for exploring memory / latency tradeoffs of transformer inference on the SIP architecture. Highlights: - 30+ model presets (Qwen, Llama 2/3/3.1, Mistral, Gemma 2, Phi 3, DeepSeek MLA, Mixtral/Qwen 3 MoE, Grok-1, ...) - Placement toggles: TP and CP each on PE-level vs cube-level - CP ring variant: K/V ring vs Q+O/m/l ring (prefill); in decode the O/m/l all-reduce is folded into S8 (no separate C1 row) - SIP interconnect: ring / mesh2d / torus2d with matching link drawing - Per-stage latency table with compute + memory + comm formulas, auto-scaled ns/us/ms, colored by dominant bound - Ring attention loop indicator on the pipeline diagram (purple arc over S5-S8 with 'xN hops' badge) - Tensor sharding view with optional physical PE/cube annotations - Replication-waste + optimization-hints panel - Save & compare configurations (config1, config2, ...): summary table plus side-by-side per-stage attention and FFN latency, best-in-row highlighting - Symbol glossary with current values for every symbol used in formulas Not tied to production sim_engine or runtime API; purely analytical tooling for design-space exploration. Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
226 lines
8.2 KiB
Python
226 lines
8.2 KiB
Python
"""Per-PE memory footprint: weights + KV cache + activations + slack.
|
||
|
||
All formulas are per-PE, per-PP-stage. If PP > 1, each PE holds only
|
||
its stage's layers (layers/PP).
|
||
"""
|
||
from __future__ import annotations
|
||
|
||
from dataclasses import dataclass
|
||
|
||
from .model_config import FullConfig
|
||
|
||
|
||
@dataclass
|
||
class MemoryBreakdown:
|
||
weights_bytes: int
|
||
kv_cache_bytes: int
|
||
transient_bytes: int
|
||
budget_bytes: int
|
||
|
||
@property
|
||
def used_bytes(self) -> int:
|
||
return self.weights_bytes + self.kv_cache_bytes + self.transient_bytes
|
||
|
||
@property
|
||
def slack_bytes(self) -> int:
|
||
return max(0, self.budget_bytes - self.used_bytes)
|
||
|
||
@property
|
||
def over_budget(self) -> bool:
|
||
return self.used_bytes > self.budget_bytes
|
||
|
||
|
||
def per_pe_weight_bytes(cfg: FullConfig) -> int:
|
||
"""Attention + FFN weights per PE, bf16. Divided by TP, PP, and EP.
|
||
|
||
EP divides the FFN (experts) across ranks; attention weights are
|
||
unaffected. When kv_shard_mode='replicate' and TP > H_kv, each KV
|
||
head is replicated (per-PE W_K/W_V size doesn't drop below one head).
|
||
"""
|
||
m = cfg.model
|
||
tp = cfg.topo.tp
|
||
pp = cfg.topo.pp
|
||
ep = max(1, cfg.topo.ep)
|
||
|
||
hq_per_pe = cfg.h_q_per_pe
|
||
if cfg.topo.kv_shard_mode == "replicate":
|
||
# Each KV head held fully; replicated across ranks if TP > H_kv.
|
||
hkv_per_pe_bytes = max(1.0, m.h_kv / tp)
|
||
else: # "split"
|
||
# Head-dim split: fractional head allowed (< 1 head bytes when TP > H_kv).
|
||
hkv_per_pe_bytes = m.h_kv / tp
|
||
|
||
per_layer_attn = (
|
||
m.hidden * hq_per_pe * m.d_head + # W_Q
|
||
m.hidden * hkv_per_pe_bytes * m.d_head + # W_K
|
||
m.hidden * hkv_per_pe_bytes * m.d_head + # W_V
|
||
hq_per_pe * m.d_head * m.hidden # W_O
|
||
)
|
||
# FFN divisor: TP (default) or TP*CP or TP*CP*DP if user opts in.
|
||
# EP further divides (MoE experts).
|
||
ffn_div = cfg.ffn_shard_divisor * ep
|
||
per_layer_ffn = 3 * m.hidden * (m.ffn_dim // ffn_div)
|
||
per_layer = (per_layer_attn + per_layer_ffn) * m.bytes_per_elem
|
||
|
||
layers_per_stage = (m.layers + pp - 1) // pp
|
||
return int(per_layer * layers_per_stage)
|
||
|
||
|
||
def per_pe_kv_cache_bytes(cfg: FullConfig) -> int:
|
||
"""K + V per PE, across all layers this stage holds.
|
||
|
||
- CP shards the sequence dim -> S_local = S_kv/CP tokens per PE.
|
||
- TP splits KV heads across ranks. If kv_shard_mode='replicate' and
|
||
TP > H_kv, each KV head is duplicated across TP/H_kv ranks so
|
||
per-PE storage doesn't shrink below 1 head.
|
||
- PP: only layers_per_stage layers stored per PE.
|
||
"""
|
||
m = cfg.model
|
||
pp = cfg.topo.pp
|
||
tp = cfg.topo.tp
|
||
if cfg.topo.kv_shard_mode == "replicate":
|
||
hkv_per_pe_bytes = max(1.0, m.h_kv / tp)
|
||
else:
|
||
hkv_per_pe_bytes = m.h_kv / tp
|
||
layers_per_stage = (m.layers + pp - 1) // pp
|
||
per_layer = 2 * cfg.topo.s_local * hkv_per_pe_bytes * m.d_head * m.bytes_per_elem
|
||
return int(per_layer * layers_per_stage)
|
||
|
||
|
||
def per_pe_transient_bytes(cfg: FullConfig) -> int:
|
||
"""Rough peak transient (activations, GEMM outputs) per PE."""
|
||
m = cfg.model
|
||
T_q = cfg.topo.T_q
|
||
hq_per_pe = cfg.h_q_per_pe
|
||
|
||
if cfg.topo.mode == "decode":
|
||
return 4 * (m.hidden + m.head_dim_total_q // cfg.topo.tp) * m.bytes_per_elem
|
||
else:
|
||
TILE = 1024
|
||
tile_score = hq_per_pe * T_q * TILE * m.bytes_per_elem
|
||
return int(2 * tile_score + T_q * m.hidden * m.bytes_per_elem)
|
||
|
||
|
||
def compute_memory(cfg: FullConfig) -> MemoryBreakdown:
|
||
return MemoryBreakdown(
|
||
weights_bytes=per_pe_weight_bytes(cfg),
|
||
kv_cache_bytes=per_pe_kv_cache_bytes(cfg),
|
||
transient_bytes=per_pe_transient_bytes(cfg),
|
||
budget_bytes=cfg.machine.pe_budget_bytes,
|
||
)
|
||
|
||
|
||
def _one_layer_row(name: str,
|
||
global_shape: tuple[int, int],
|
||
per_pe_shape: tuple[int, int],
|
||
bytes_per_elem: int) -> dict:
|
||
"""Per-tensor row for ONE layer (both global and per-PE shard)."""
|
||
p_global = global_shape[0] * global_shape[1]
|
||
p_per_pe = per_pe_shape[0] * per_pe_shape[1]
|
||
return {
|
||
"Tensor": name,
|
||
"Global shape": f"({global_shape[0]}, {global_shape[1]})",
|
||
"Params/layer": f"{p_global/1e6:.2f} M",
|
||
"Bytes/layer (global)": f"{p_global * bytes_per_elem / 1e6:.2f} MB",
|
||
"Per-PE shape": f"({per_pe_shape[0]}, {per_pe_shape[1]})",
|
||
"Bytes/layer (per PE)": f"{p_per_pe * bytes_per_elem / 1e6:.2f} MB",
|
||
"_p_global": p_global,
|
||
"_p_per_pe": p_per_pe,
|
||
}
|
||
|
||
|
||
def attention_weight_rows(cfg: FullConfig) -> list[dict]:
|
||
"""Per-tensor rows for attention weights (one layer)."""
|
||
m = cfg.model
|
||
hq_per_pe = cfg.h_q_per_pe
|
||
hkv_per_pe = max(1, m.h_kv // cfg.topo.tp)
|
||
return [
|
||
_one_layer_row("W_Q", (m.hidden, m.h_q * m.d_head),
|
||
(m.hidden, hq_per_pe * m.d_head),
|
||
m.bytes_per_elem),
|
||
_one_layer_row("W_K", (m.hidden, m.h_kv * m.d_head),
|
||
(m.hidden, hkv_per_pe * m.d_head),
|
||
m.bytes_per_elem),
|
||
_one_layer_row("W_V", (m.hidden, m.h_kv * m.d_head),
|
||
(m.hidden, hkv_per_pe * m.d_head),
|
||
m.bytes_per_elem),
|
||
_one_layer_row("W_O", (m.h_q * m.d_head, m.hidden),
|
||
(hq_per_pe * m.d_head, m.hidden),
|
||
m.bytes_per_elem),
|
||
]
|
||
|
||
|
||
def ffn_weight_rows(cfg: FullConfig) -> list[dict]:
|
||
"""Per-tensor rows for FFN weights (one layer, activated for MoE)."""
|
||
m = cfg.model
|
||
ffn_per_pe = m.ffn_dim // cfg.topo.tp
|
||
return [
|
||
_one_layer_row("W_gate", (m.hidden, m.ffn_dim),
|
||
(m.hidden, ffn_per_pe), m.bytes_per_elem),
|
||
_one_layer_row("W_up", (m.hidden, m.ffn_dim),
|
||
(m.hidden, ffn_per_pe), m.bytes_per_elem),
|
||
_one_layer_row("W_down", (m.ffn_dim, m.hidden),
|
||
(ffn_per_pe, m.hidden), m.bytes_per_elem),
|
||
]
|
||
|
||
|
||
def kv_cache_rows(cfg: FullConfig) -> list[dict]:
|
||
"""Per-tensor row for K and V cache (one layer)."""
|
||
m = cfg.model
|
||
hkv_per_pe = max(1, m.h_kv // cfg.topo.tp)
|
||
b = m.bytes_per_elem
|
||
global_k = cfg.topo.s_kv * m.h_kv * m.d_head
|
||
per_pe_k = cfg.topo.s_local * hkv_per_pe * m.d_head
|
||
return [
|
||
{
|
||
"Tensor": "K cache",
|
||
"Global shape": f"({cfg.topo.s_kv:,}, {m.h_kv * m.d_head})",
|
||
"Params/layer": f"{global_k/1e6:.2f} M",
|
||
"Bytes/layer (global)": f"{global_k * b / 1e6:.2f} MB",
|
||
"Per-PE shape": f"({cfg.topo.s_local:,}, {hkv_per_pe * m.d_head})",
|
||
"Bytes/layer (per PE)": f"{per_pe_k * b / 1e6:.2f} MB",
|
||
"_p_global": global_k,
|
||
"_p_per_pe": per_pe_k,
|
||
},
|
||
{
|
||
"Tensor": "V cache",
|
||
"Global shape": f"({cfg.topo.s_kv:,}, {m.h_kv * m.d_head})",
|
||
"Params/layer": f"{global_k/1e6:.2f} M",
|
||
"Bytes/layer (global)": f"{global_k * b / 1e6:.2f} MB",
|
||
"Per-PE shape": f"({cfg.topo.s_local:,}, {hkv_per_pe * m.d_head})",
|
||
"Bytes/layer (per PE)": f"{per_pe_k * b / 1e6:.2f} MB",
|
||
"_p_global": global_k,
|
||
"_p_per_pe": per_pe_k,
|
||
},
|
||
]
|
||
|
||
|
||
def sum_bytes_all_layers(rows: list[dict], cfg: FullConfig,
|
||
per_pe: bool = False) -> int:
|
||
"""Sum bytes across all rows × N layers (or layers/PP for per_pe)."""
|
||
b = cfg.model.bytes_per_elem
|
||
layers = cfg.model.layers
|
||
if per_pe:
|
||
layers = (cfg.model.layers + cfg.topo.pp - 1) // cfg.topo.pp
|
||
key = "_p_per_pe" if per_pe else "_p_global"
|
||
return sum(r[key] for r in rows) * b * layers
|
||
|
||
|
||
def total_weight_bytes_full_model(cfg: FullConfig) -> int:
|
||
"""Total bytes of ALL weights (attention + FFN) for the full unsharded model."""
|
||
m = cfg.model
|
||
attn_per_layer = (
|
||
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_per_layer = 3 * m.hidden * m.ffn_dim
|
||
return (attn_per_layer + ffn_per_layer) * m.bytes_per_elem * m.layers
|
||
|
||
|
||
def total_kv_bytes_full_model(cfg: FullConfig) -> int:
|
||
"""Total bytes of KV cache for full unsharded model at cfg.topo.s_kv."""
|
||
m = cfg.model
|
||
per_layer = 2 * cfg.topo.s_kv * m.h_kv * m.d_head * m.bytes_per_elem
|
||
return per_layer * m.layers
|