analytical-viz: interactive Streamlit tool for SIP transformer analysis
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>
This commit is contained in:
@@ -0,0 +1,126 @@
|
||||
"""Model preset library for the analytical visualization tool.
|
||||
|
||||
Comprehensive coverage of popular MHA, GQA, MQA, and MLA models across
|
||||
sizes. MoE models are approximated as dense with (ffn_dim = activated
|
||||
experts × per_expert_ffn) for compute-side estimates; the note field
|
||||
flags this. MLA models are approximated as GQA with matching H_kv until
|
||||
we add proper MLA support.
|
||||
|
||||
Source of numeric params: model cards / config.json from HuggingFace.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
from .model_config import ModelConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class Preset:
|
||||
label: str
|
||||
model: ModelConfig
|
||||
family: str = "" # e.g. "Llama 3", "Qwen 3"
|
||||
attn_type: str = "" # "MHA" | "GQA" | "MQA" | "MLA"
|
||||
note: str = ""
|
||||
|
||||
|
||||
def _mk(name: str, hidden: int, ffn: int, hq: int, hkv: int,
|
||||
dh: int, layers: int, *, family: str = "", attn: str = "GQA",
|
||||
note: str = "") -> Preset:
|
||||
return Preset(
|
||||
label=name,
|
||||
model=ModelConfig(
|
||||
name=name, hidden=hidden, ffn_dim=ffn,
|
||||
h_q=hq, h_kv=hkv, d_head=dh, layers=layers,
|
||||
),
|
||||
family=family, attn_type=attn, note=note,
|
||||
)
|
||||
|
||||
|
||||
PRESETS: dict[str, Preset] = {
|
||||
# ── MHA (H_q == H_kv) ─────────────────────────────────────────
|
||||
"GPT-2 (117M)": _mk("GPT-2 117M", 768, 3072, 12, 12, 64, 12, family="GPT-2", attn="MHA"),
|
||||
"GPT-2 XL (1.5B)": _mk("GPT-2 XL 1.5B", 1600, 6400, 25, 25, 64, 48, family="GPT-2", attn="MHA"),
|
||||
"Llama 2 7B": _mk("Llama 2 7B", 4096, 11008, 32, 32, 128, 32, family="Llama 2", attn="MHA"),
|
||||
"Llama 2 13B": _mk("Llama 2 13B", 5120, 13824, 40, 40, 128, 40, family="Llama 2", attn="MHA"),
|
||||
"Yi 6B": _mk("Yi 6B", 4096, 11008, 32, 4, 128, 32, family="Yi", attn="GQA"),
|
||||
"Yi 9B": _mk("Yi 9B", 4096, 11008, 32, 4, 128, 48, family="Yi", attn="GQA"),
|
||||
"OPT 30B": _mk("OPT 30B", 7168, 28672, 56, 56, 128, 48, family="OPT", attn="MHA"),
|
||||
|
||||
# ── GQA (H_q > H_kv > 1) ──────────────────────────────────────
|
||||
"Qwen 2 0.5B": _mk("Qwen 2 0.5B", 896, 4864, 14, 2, 64, 24, family="Qwen 2"),
|
||||
"Qwen 2 1.5B": _mk("Qwen 2 1.5B", 1536, 8960, 12, 2, 128, 28, family="Qwen 2"),
|
||||
"Qwen 2 7B": _mk("Qwen 2 7B", 3584, 18944, 28, 4, 128, 28, family="Qwen 2"),
|
||||
"Qwen 2 72B": _mk("Qwen 2 72B", 8192, 29568, 64, 8, 128, 80, family="Qwen 2"),
|
||||
|
||||
"Qwen 3 0.6B": _mk("Qwen 3 0.6B", 1024, 3072, 16, 8, 128, 28, family="Qwen 3"),
|
||||
"Qwen 3 1.7B": _mk("Qwen 3 1.7B", 2048, 6144, 16, 8, 128, 28, family="Qwen 3"),
|
||||
"Qwen 3 4B": _mk("Qwen 3 4B", 2560, 9728, 32, 8, 128, 36, family="Qwen 3"),
|
||||
"Qwen 3 8B": _mk("Qwen 3 8B", 4096, 12288, 32, 8, 128, 36, family="Qwen 3"),
|
||||
"Qwen 3 14B": _mk("Qwen 3 14B", 5120, 17408, 40, 8, 128, 40, family="Qwen 3"),
|
||||
"Qwen 3 32B": _mk("Qwen 3 32B", 5120, 25600, 64, 8, 128, 64, family="Qwen 3"),
|
||||
|
||||
"Llama 3 8B": _mk("Llama 3 8B", 4096, 14336, 32, 8, 128, 32, family="Llama 3"),
|
||||
"Llama 3 70B": _mk("Llama 3 70B", 8192, 28672, 64, 8, 128, 80, family="Llama 3"),
|
||||
"Llama 3.1 8B": _mk("Llama 3.1 8B", 4096, 14336, 32, 8, 128, 32, family="Llama 3.1"),
|
||||
"Llama 3.1 70B": _mk("Llama 3.1 70B", 8192, 28672, 64, 8, 128, 80, family="Llama 3.1"),
|
||||
"Llama 3.1 405B": _mk("Llama 3.1 405B", 16384,53248, 128,8, 128, 126,family="Llama 3.1"),
|
||||
"Llama 3.2 1B": _mk("Llama 3.2 1B", 2048, 8192, 32, 8, 64, 16, family="Llama 3.2"),
|
||||
"Llama 3.2 3B": _mk("Llama 3.2 3B", 3072, 8192, 24, 8, 128, 28, family="Llama 3.2"),
|
||||
|
||||
"Mistral 7B": _mk("Mistral 7B v0.3", 4096, 14336, 32, 8, 128, 32, family="Mistral"),
|
||||
"Mistral Nemo 12B": _mk("Mistral Nemo 12B", 5120, 14336, 32, 8, 128, 40, family="Mistral"),
|
||||
"Mistral Small 22B": _mk("Mistral Small 22B",6144, 16384, 48, 8, 128, 56, family="Mistral"),
|
||||
"Mistral Large 123B": _mk("Mistral Large 123B",12288,28672,96, 8, 128, 88, family="Mistral"),
|
||||
|
||||
"Gemma 2 2B": _mk("Gemma 2 2B", 2304, 9216, 8, 4, 256, 26, family="Gemma 2"),
|
||||
"Gemma 2 9B": _mk("Gemma 2 9B", 3584, 14336, 16, 8, 256, 42, family="Gemma 2"),
|
||||
"Gemma 2 27B": _mk("Gemma 2 27B", 4608, 36864, 32, 16, 128, 46, family="Gemma 2"),
|
||||
|
||||
"Phi 3 mini (3.8B)": _mk("Phi 3 mini 3.8B", 3072, 8192, 32, 32, 96, 32, family="Phi 3", attn="MHA"),
|
||||
"Phi 3 small (7B)": _mk("Phi 3 small 7B", 4096, 14336, 32, 8, 128, 32, family="Phi 3"),
|
||||
"Phi 3 medium (14B)": _mk("Phi 3 medium 14B", 5120, 17920, 40, 10, 128, 40, family="Phi 3"),
|
||||
|
||||
"Yi 34B": _mk("Yi 34B", 7168, 20480, 56, 8, 128, 60, family="Yi"),
|
||||
"Command R+ 104B": _mk("Command R+ 104B", 12288,33792, 96, 8, 128, 64, family="Cohere"),
|
||||
|
||||
# ── MoE (activated-only approximation) ────────────────────────
|
||||
"Mixtral 8x7B (MoE)": _mk("Mixtral 8x7B", 4096, 14336*2, 32, 8, 128, 32,
|
||||
family="Mistral MoE",
|
||||
note="MoE: 8 experts x 2 activated. FFN ~ 2 x per_expert."),
|
||||
"Mixtral 8x22B (MoE)":_mk("Mixtral 8x22B", 6144, 16384*2, 48, 8, 128, 56,
|
||||
family="Mistral MoE",
|
||||
note="MoE: 8 experts x 2 activated."),
|
||||
"Qwen 3 30B (MoE)": _mk("Qwen 3 30B (MoE)", 2048, 768*8, 32, 4, 128, 48,
|
||||
family="Qwen 3 MoE",
|
||||
note="128 experts x 8 activated per token."),
|
||||
"Qwen 3 235B (MoE)": _mk("Qwen 3 235B (MoE)",4096, 1536*8, 64, 4, 128, 94,
|
||||
family="Qwen 3 MoE",
|
||||
note="128 experts x 8 activated. Dense-approx overstates weight."),
|
||||
"DeepSeek V2 (dense-approx)":
|
||||
_mk("DeepSeek V2", 5120, 12288, 128, 128, 128, 60,
|
||||
family="DeepSeek", attn="MLA",
|
||||
note="MLA compresses KV to ~576 B/token; d_c=512 latent. Dense-approx."),
|
||||
"DeepSeek V3 (dense-approx)":
|
||||
_mk("DeepSeek V3", 7168, 16384, 128, 128, 128, 61,
|
||||
family="DeepSeek", attn="MLA",
|
||||
note="MoE 256 x 8 activated + MLA. Both approximations."),
|
||||
"Grok-1 314B (MoE)": _mk("Grok-1 314B", 6144, 32768*2, 48, 8, 128, 64,
|
||||
family="xAI",
|
||||
note="MoE 8 experts x 2 activated."),
|
||||
|
||||
# ── MQA (H_kv == 1) ───────────────────────────────────────────
|
||||
"Falcon 7B (MQA)": _mk("Falcon 7B", 4544, 18176, 71, 1, 64, 32, family="Falcon", attn="MQA"),
|
||||
"Falcon 40B (MQA)": _mk("Falcon 40B", 8192, 32768, 128,1, 64, 60, family="Falcon", attn="MQA"),
|
||||
|
||||
# ── Custom slot ──────────────────────────────────────────────
|
||||
"Custom": Preset(label="Custom", model=ModelConfig(name="Custom")),
|
||||
}
|
||||
|
||||
|
||||
def preset_names_by_family() -> dict[str, list[str]]:
|
||||
"""Return preset names grouped by family (for a nested dropdown)."""
|
||||
grouped: dict[str, list[str]] = {}
|
||||
for name, p in PRESETS.items():
|
||||
grouped.setdefault(p.family or "Other", []).append(name)
|
||||
return grouped
|
||||
Reference in New Issue
Block a user