"""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