|
Download README.md from AlgoDriveAI/AIMO_EP_25: direct link, hf CLI and curl.
- Browser
- Download file 49.2 kB
-
https://huggingface.co/AlgoDriveAI/AIMO_EP_25/resolve/main/README.md
- Command line
-
hf download hf://AlgoDriveAI/AIMO_EP_25/README.md
-
curl -L -H "Authorization: Bearer $HF_TOKEN" -o README.md https://huggingface.co/AlgoDriveAI/AIMO_EP_25/resolve/main/README.md
49.2 kB
| license: mit | |
| language: | |
| - en | |
| pipeline_tag: text-generation | |
| library_name: pytorch | |
| tags: | |
| - mathematics | |
| - math | |
| - text-generation | |
| - pretrained | |
| - base-model | |
| - custom-code | |
| - mla | |
| - pytorch | |
| - gradio | |
| # The Mathematician V2 β Beta | |
| This is the beta release of **The Mathematician V2**, a pretrained mathematical language model. | |
| This is a raw pretrained base model with no SFT, GRPO, or RL applied. | |
| The inference script below automatically downloads the checkpoint from: | |
| `AlgoDriveAI/AIMO_EP_25/checkpoints_optimized2/latest_checkpoint.pt` | |
| ## Run in Google Colab | |
| Copy and paste the following install command into a Colab cell and run it: | |
| ```bash | |
| !pip install -q torch transformers gradio huggingface_hub sentencepiece | |
| ``` | |
| Then copy and paste the **full inference code below** into the next Colab cell and run it. | |
| The checkpoint is downloaded automatically from Hugging Face. You do not need to manually upload the `.pt` file or edit a `/content/checkpoints/...` path. | |
| ## Full Inference Code | |
| ```python | |
| #!/usr/bin/env python3 | |
| # -*- coding: utf-8 -*- | |
| """ | |
| THE MATHEMATICIAN V2 β BASE (PRETRAINED) INFERENCE | |
| =============================================================================== | |
| Gradio + CLI inference for the *pretrained base* checkpoint. No GRPO, no SFT, | |
| no reward machinery β this is raw next-token completion from the pretrain run. | |
| ARCHITECTURE (older V2 MLA variant β matches the training script exactly): | |
| β’ q_proj : single full projection, split at runtime into [q_nope | q_rope] | |
| (NOT the canonical q_a_proj / q_a_norm / q_b_proj compression) | |
| β’ k_rope_proj : PER-HEAD (d_model β n_heads * qk_rope_head_dim) | |
| (NOT the shared single-head broadcast used in later versions) | |
| β’ k_up / v_up : separate up-projections off the latent | |
| (NOT a fused kv_b_proj) | |
| β’ QK-norm : applied AFTER concat + AFTER RoPE, over the FULL head_dim | |
| (NOT pre-RoPE on nope dims only) | |
| β’ SwiGLU : fused gate_up_proj β chunk(2) | |
| β’ lm_head weight-tied to embed | |
| β’ RoPE over qk_rope_head_dim only, rotate_half convention, base 10000 | |
| If you point this at a checkpoint from a *newer* architecture it will fail | |
| loudly on load_state_dict rather than silently mis-wire β that's intentional. | |
| FEATURES: | |
| β’ Incremental KV-cache decoding (prefill once, then 1 token/step) | |
| β’ Streaming Gradio output + Stop button | |
| β’ Architecture auto-detection from the checkpoint state dict (see | |
| AUTO_ADOPT_CHECKPOINT_ARCH) β no more guessing whether it's 2048d or 2304d | |
| β’ maj@k tab: batched self-consistency sampling with majority vote over | |
| extracted answers | |
| β’ CLI mode via RUN_CLI | |
| CONFIG: all hardcoded at the top of this file. No argparse. | |
| """ | |
| import os | |
| import re | |
| import gc | |
| import time | |
| import random | |
| import warnings | |
| from collections import Counter | |
| from dataclasses import dataclass | |
| from typing import Optional, Tuple, List, Dict, Any | |
| os.environ.setdefault("TOKENIZERS_PARALLELISM", "false") | |
| os.environ["HF_HUB_DISABLE_XET"] = "1" | |
| warnings.filterwarnings("ignore", message=r".*UnsupportedFieldAttributeWarning.*") | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| from transformers import AutoTokenizer | |
| from huggingface_hub import hf_hub_download | |
| # ============================================================================= | |
| # HARDCODED CONFIGURATION | |
| # ============================================================================= | |
| # ββ Checkpoint βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| # Download the beta checkpoint directly from Hugging Face. | |
| CHECKPOINT_REPO = "AlgoDriveAI/AIMO_EP_25" | |
| CHECKPOINT_SUBFOLDER = "checkpoints_optimized2" | |
| CHECKPOINT_FILE = "latest_checkpoint.pt" | |
| print(f"π₯ Downloading checkpoint from Hugging Face: {CHECKPOINT_REPO}/{CHECKPOINT_SUBFOLDER}/{CHECKPOINT_FILE}") | |
| CHECKPOINT_PATH = hf_hub_download( | |
| repo_id=CHECKPOINT_REPO, | |
| filename=f"{CHECKPOINT_SUBFOLDER}/{CHECKPOINT_FILE}", | |
| ) | |
| # The rest of the inference code expects a directory + filename. | |
| CHECKPOINT_DIR = os.path.dirname(CHECKPOINT_PATH) | |
| print(f" β Checkpoint ready: {CHECKPOINT_PATH}") | |
| # ββ Run mode βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| RUN_CLI = False # True β terminal REPL instead of Gradio | |
| GRADIO_SHARE = True | |
| GRADIO_PORT = 7860 | |
| # ββ Architecture (must match the pretrain run) βββββββββββββββββββββββββββββββ | |
| D_MODEL = 2048 | |
| N_LAYERS = 36 | |
| N_HEADS = 16 | |
| KV_LORA_RANK = 512 # MLA latent compression rank | |
| QK_ROPE_HEAD_DIM = 64 # decoupled RoPE dims per head | |
| FF_MULT = 3.5 # SwiGLU multiplier | |
| CONTEXT_LEN = 4096 | |
| QK_NORM = True | |
| ROPE_BASE = 10000.0 | |
| # Derived (recomputed after auto-detect) | |
| HEAD_DIM = D_MODEL // N_HEADS | |
| QK_NOPE_HEAD_DIM = HEAD_DIM - QK_ROPE_HEAD_DIM | |
| # If True, read the true architecture out of the checkpoint's state dict and | |
| # override the constants above before building the model. Every dim in this | |
| # architecture is recoverable from tensor shapes, so this is exact, not a guess. | |
| AUTO_ADOPT_CHECKPOINT_ARCH = True | |
| # ββ Tokenizer ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| VOCAB_NAME = "mistralai/Mistral-7B-v0.1" | |
| DOC_EOS_TOKEN = "<|endoftext|>" | |
| VOCAB_PAD_MULTIPLE = 64 | |
| # ββ Generation defaults ββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| DEFAULT_MAX_NEW_TOKENS = 512 | |
| DEFAULT_TEMPERATURE = 0.8 | |
| DEFAULT_TOP_K = 50 | |
| DEFAULT_TOP_P = 0.95 | |
| DEFAULT_REP_PENALTY = 1.05 | |
| DEFAULT_MAJ_K = 8 # samples for the maj@k tab | |
| # Prompt templates offered in the UI. The base model saw `[SOLUTION]` section | |
| # headers throughout the synthetic pretrain corpus, so that template is usually | |
| # the strongest completion trigger. "Raw" sends your text untouched. | |
| PROMPT_TEMPLATES = { | |
| "Raw (no template)": "{p}", | |
| "[SOLUTION] header": "{p}\n[SOLUTION]\n", | |
| "Solve step by step": "{p} Solve step by step:\n", | |
| } | |
| DEFAULT_TEMPLATE = "[SOLUTION] header" | |
| # ============================================================================= | |
| # RMSNorm β fused if available, custom fallback otherwise | |
| # ============================================================================= | |
| try: | |
| from torch.nn import RMSNorm | |
| _rmsnorm_source = "nn.RMSNorm (fused)" | |
| except ImportError: | |
| class RMSNorm(nn.Module): | |
| __constants__ = ["eps"] | |
| def __init__(self, normalized_shape, eps: float = 1e-6, **kwargs): | |
| super().__init__() | |
| if isinstance(normalized_shape, int): | |
| normalized_shape = (normalized_shape,) | |
| self.eps = eps | |
| self.weight = nn.Parameter(torch.ones(normalized_shape)) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| return self.weight * ( | |
| x.float() * torch.rsqrt(x.float().pow(2).mean(-1, keepdim=True) + self.eps) | |
| ).to(x.dtype) | |
| _rmsnorm_source = "custom fallback" | |
| # ============================================================================= | |
| # ROTARY EMBEDDINGS β precomputed, rotate_half convention | |
| # ============================================================================= | |
| class RotaryEmbedding(nn.Module): | |
| def __init__(self, dim: int, base: float = ROPE_BASE, max_seq_len: int = 8192): | |
| super().__init__() | |
| inv_freq = 1.0 / (base ** (torch.arange(0, dim, 2, dtype=torch.float32) / dim)) | |
| t = torch.arange(max_seq_len, dtype=torch.float32) | |
| freqs = torch.outer(t, inv_freq) | |
| self.register_buffer("cos_cached", freqs.cos().repeat(1, 2), persistent=False) | |
| self.register_buffer("sin_cached", freqs.sin().repeat(1, 2), persistent=False) | |
| def forward(self, seq_len: int, dtype: torch.dtype): | |
| return self.cos_cached[:seq_len].to(dtype), self.sin_cached[:seq_len].to(dtype) | |
| def _rotate_half(x: torch.Tensor) -> torch.Tensor: | |
| half = x.shape[-1] // 2 | |
| return torch.cat([-x[..., half:], x[..., :half]], dim=-1) | |
| def apply_rotary_emb(q, k, cos, sin): | |
| cos = cos.unsqueeze(0).unsqueeze(0) | |
| sin = sin.unsqueeze(0).unsqueeze(0) | |
| return (q * cos) + (_rotate_half(q) * sin), (k * cos) + (_rotate_half(k) * sin) | |
| # ============================================================================= | |
| # MLA β verbatim from the training script (per-head k_rope, split q_proj) | |
| # ============================================================================= | |
| class MLA(nn.Module): | |
| def __init__(self, d_model, n_heads, kv_lora_rank, qk_rope_head_dim, rope, | |
| qk_norm=False, attn_dropout=0.0): | |
| super().__init__() | |
| assert d_model % n_heads == 0 | |
| self.n_heads = n_heads | |
| self.head_dim = d_model // n_heads | |
| self.qk_rope_head_dim = qk_rope_head_dim | |
| self.qk_nope_head_dim = self.head_dim - qk_rope_head_dim | |
| self.kv_lora_rank = kv_lora_rank | |
| self.attn_drop = attn_dropout | |
| self.rope = rope | |
| self.q_proj = nn.Linear(d_model, n_heads * self.head_dim, bias=False) | |
| self.kv_down = nn.Linear(d_model, kv_lora_rank, bias=False) | |
| self.kv_norm = RMSNorm(kv_lora_rank) | |
| self.k_up = nn.Linear(kv_lora_rank, n_heads * self.qk_nope_head_dim, bias=False) | |
| self.v_up = nn.Linear(kv_lora_rank, n_heads * self.head_dim, bias=False) | |
| self.k_rope_proj = nn.Linear(d_model, n_heads * qk_rope_head_dim, bias=False) | |
| self.o_proj = nn.Linear(d_model, d_model, bias=False) | |
| self.o_proj._is_residual = True | |
| self.qk_norm = qk_norm | |
| if qk_norm: | |
| self.q_norm = RMSNorm(self.head_dim) | |
| self.k_norm = RMSNorm(self.head_dim) | |
| def forward(self, x, past_kv=None, use_cache=False): | |
| B, T, _ = x.shape | |
| q = self.q_proj(x).view(B, T, self.n_heads, self.head_dim) | |
| q_nope = q[..., :self.qk_nope_head_dim] | |
| q_rope = q[..., self.qk_nope_head_dim:] | |
| c = self.kv_norm(self.kv_down(x)) | |
| k_nope = self.k_up(c).view(B, T, self.n_heads, self.qk_nope_head_dim) | |
| v = self.v_up(c).view(B, T, self.n_heads, self.head_dim) | |
| k_rope = self.k_rope_proj(x).view(B, T, self.n_heads, self.qk_rope_head_dim) | |
| pos_offset = 0 | |
| if past_kv is not None: | |
| pos_offset = past_kv[0].shape[2] | |
| cos, sin = self.rope(pos_offset + T, x.dtype) | |
| cos = cos[pos_offset:pos_offset + T] | |
| sin = sin[pos_offset:pos_offset + T] | |
| q_rope = q_rope.transpose(1, 2) | |
| k_rope = k_rope.transpose(1, 2) | |
| q_rope, k_rope = apply_rotary_emb(q_rope, k_rope, cos, sin) | |
| q_rope = q_rope.transpose(1, 2) | |
| k_rope = k_rope.transpose(1, 2) | |
| q = torch.cat([q_nope, q_rope], dim=-1).transpose(1, 2).contiguous() | |
| k = torch.cat([k_nope, k_rope], dim=-1).transpose(1, 2).contiguous() | |
| v = v.transpose(1, 2).contiguous() | |
| # QK-norm AFTER assembly, over full head_dim β matches training. | |
| if self.qk_norm: | |
| q = self.q_norm(q) | |
| k = self.k_norm(k) | |
| present_kv = None | |
| if use_cache: | |
| if past_kv is not None: | |
| k = torch.cat([past_kv[0], k], dim=2) | |
| v = torch.cat([past_kv[1], v], dim=2) | |
| present_kv = (k, v) | |
| drop_p = self.attn_drop if self.training else 0.0 | |
| is_causal = (T > 1) and (past_kv is None) | |
| out = F.scaled_dot_product_attention(q, k, v, dropout_p=drop_p, is_causal=is_causal) | |
| return self.o_proj(out.transpose(1, 2).reshape(B, T, -1)), present_kv | |
| class SwiGLU(nn.Module): | |
| def __init__(self, d_model, hidden_mult=3.5): | |
| super().__init__() | |
| inner = int(hidden_mult * d_model) | |
| self.gate_up_proj = nn.Linear(d_model, 2 * inner, bias=False) | |
| self.down_proj = nn.Linear(inner, d_model, bias=False) | |
| self.down_proj._is_residual = True | |
| def forward(self, x): | |
| gate, up = self.gate_up_proj(x).chunk(2, dim=-1) | |
| return self.down_proj(F.silu(gate) * up) | |
| class Block(nn.Module): | |
| def __init__(self, d_model, n_heads, kv_lora_rank, qk_rope_head_dim, rope, | |
| ff_hidden_mult=3.5, qk_norm=False, attn_dropout=0.0, resid_dropout=0.0): | |
| super().__init__() | |
| self.ln_attn = RMSNorm(d_model) | |
| self.ln_ff = RMSNorm(d_model) | |
| self.attn = MLA(d_model, n_heads, kv_lora_rank, qk_rope_head_dim, | |
| rope, qk_norm=qk_norm, attn_dropout=attn_dropout) | |
| self.ff = SwiGLU(d_model, hidden_mult=ff_hidden_mult) | |
| self.resid_drop = nn.Dropout(resid_dropout) if resid_dropout > 0 else nn.Identity() | |
| def forward(self, x, past_kv=None, use_cache=False): | |
| attn_out, present_kv = self.attn(self.ln_attn(x), past_kv=past_kv, use_cache=use_cache) | |
| x = x + self.resid_drop(attn_out) | |
| x = x + self.resid_drop(self.ff(self.ln_ff(x))) | |
| return x, present_kv | |
| @dataclass | |
| class ModelConfig: | |
| vocab_size: int | |
| d_model: int | |
| n_layers: int | |
| n_heads: int | |
| kv_lora_rank: int | |
| qk_rope_head_dim: int | |
| ff_hidden_mult: float | |
| qk_norm: bool | |
| max_seq_len: int = 8192 | |
| attn_dropout: float = 0.0 | |
| resid_dropout: float = 0.0 | |
| emb_dropout: float = 0.0 | |
| @property | |
| def head_dim(self): return self.d_model // self.n_heads | |
| @property | |
| def qk_nope_head_dim(self): return self.head_dim - self.qk_rope_head_dim | |
| class DenseLLM(nn.Module): | |
| def __init__(self, cfg: ModelConfig): | |
| super().__init__() | |
| self.cfg = cfg | |
| self.embed = nn.Embedding(cfg.vocab_size, cfg.d_model) | |
| self.emb_drop = nn.Identity() | |
| self.rope = RotaryEmbedding(cfg.qk_rope_head_dim, max_seq_len=cfg.max_seq_len) | |
| self.blocks = nn.ModuleList([ | |
| Block(cfg.d_model, cfg.n_heads, cfg.kv_lora_rank, cfg.qk_rope_head_dim, | |
| self.rope, ff_hidden_mult=cfg.ff_hidden_mult, qk_norm=cfg.qk_norm, | |
| attn_dropout=cfg.attn_dropout, resid_dropout=cfg.resid_dropout) | |
| for _ in range(cfg.n_layers) | |
| ]) | |
| self.ln_f = RMSNorm(cfg.d_model) | |
| self.lm_head = nn.Linear(cfg.d_model, cfg.vocab_size, bias=False) | |
| self.lm_head.weight = self.embed.weight | |
| def forward(self, idx, past_kvs=None, use_cache=False): | |
| x = self.emb_drop(self.embed(idx)) | |
| present_kvs = [] if use_cache else None | |
| for i, blk in enumerate(self.blocks): | |
| layer_past = past_kvs[i] if past_kvs is not None else None | |
| x, present_kv = blk(x, past_kv=layer_past, use_cache=use_cache) | |
| if use_cache: | |
| present_kvs.append(present_kv) | |
| logits = self.lm_head(self.ln_f(x)) | |
| if use_cache: | |
| return logits, present_kvs | |
| return logits, None | |
| # ============================================================================= | |
| # TOKENIZER | |
| # ============================================================================= | |
| def get_tokenizer(vocab_name, doc_eos_token, pad_to_multiple, checkpoint_dir): | |
| if os.path.exists(os.path.join(checkpoint_dir, "tokenizer_config.json")): | |
| print(f"π Tokenizer from checkpoint dir: {checkpoint_dir}") | |
| tok = AutoTokenizer.from_pretrained(checkpoint_dir, use_fast=True) | |
| else: | |
| print(f"π Tokenizer from HuggingFace: {vocab_name}") | |
| tok = AutoTokenizer.from_pretrained(vocab_name, use_fast=True) | |
| if doc_eos_token not in tok.get_vocab(): | |
| tok.add_special_tokens({"additional_special_tokens": [doc_eos_token]}) | |
| if tok.pad_token is None: | |
| tok.pad_token = doc_eos_token | |
| tok.pad_token_id = tok.convert_tokens_to_ids(doc_eos_token) | |
| if pad_to_multiple and (len(tok) % pad_to_multiple != 0): | |
| n = pad_to_multiple - (len(tok) % pad_to_multiple) | |
| tok.add_tokens([f"<|dummy_{i}|>" for i in range(n)], special_tokens=False) | |
| tok.doc_eos_token = doc_eos_token | |
| tok.doc_eos_token_id = tok.convert_tokens_to_ids(doc_eos_token) | |
| if tok.pad_token_id is None: | |
| tok.pad_token = doc_eos_token | |
| tok.pad_token_id = tok.doc_eos_token_id | |
| tok.model_max_length = int(1e9) | |
| print(f" β Vocab: {len(tok):,} | doc EOS '{doc_eos_token}' = id {tok.doc_eos_token_id}") | |
| return tok | |
| # ============================================================================= | |
| # ARCHITECTURE AUTO-DETECTION FROM STATE DICT | |
| # ============================================================================= | |
| def infer_arch_from_state_dict(sd: Dict[str, torch.Tensor]) -> Dict[str, Any]: | |
| """ | |
| Every architectural dim in this MLA variant is recoverable from shapes: | |
| embed.weight β (vocab, d_model) | |
| blocks.N.* β n_layers (max N + 1) | |
| blocks.0.attn.q_norm.weight β (head_dim,) [only if qk_norm] | |
| blocks.0.attn.kv_down.weight β (kv_lora_rank, d_model) | |
| blocks.0.attn.k_rope_proj.weightβ (n_heads * qk_rope, d_model) | |
| blocks.0.attn.k_up.weight β (n_heads * qk_nope, kv_lora_rank) | |
| blocks.0.ff.gate_up_proj.weight β (2 * inner, d_model) | |
| Without qk_norm, head_dim is not directly observable, so n_heads falls back | |
| to the hardcoded constant and qk_rope is derived from it. | |
| """ | |
| out: Dict[str, Any] = {} | |
| vocab, d_model = sd["embed.weight"].shape | |
| out["vocab_size"] = int(vocab) | |
| out["d_model"] = int(d_model) | |
| layer_ids = set() | |
| for k in sd: | |
| m = re.match(r"blocks\.(\d+)\.", k) | |
| if m: | |
| layer_ids.add(int(m.group(1))) | |
| out["n_layers"] = max(layer_ids) + 1 if layer_ids else N_LAYERS | |
| out["qk_norm"] = "blocks.0.attn.q_norm.weight" in sd | |
| if out["qk_norm"]: | |
| head_dim = int(sd["blocks.0.attn.q_norm.weight"].shape[0]) | |
| out["n_heads"] = out["d_model"] // head_dim | |
| else: | |
| out["n_heads"] = N_HEADS | |
| head_dim = out["d_model"] // out["n_heads"] | |
| out["head_dim"] = head_dim | |
| out["kv_lora_rank"] = int(sd["blocks.0.attn.kv_down.weight"].shape[0]) | |
| out["qk_rope_head_dim"] = int(sd["blocks.0.attn.k_rope_proj.weight"].shape[0]) // out["n_heads"] | |
| out["qk_nope_head_dim"] = int(sd["blocks.0.attn.k_up.weight"].shape[0]) // out["n_heads"] | |
| inner = int(sd["blocks.0.ff.gate_up_proj.weight"].shape[0]) // 2 | |
| out["ff_inner"] = inner | |
| out["ff_mult"] = inner / out["d_model"] | |
| return out | |
| def _print_arch_table(a: Dict[str, Any], hardcoded: Dict[str, Any]): | |
| rows = [ | |
| ("d_model", a["d_model"], hardcoded["d_model"]), | |
| ("n_layers", a["n_layers"], hardcoded["n_layers"]), | |
| ("n_heads", a["n_heads"], hardcoded["n_heads"]), | |
| ("head_dim", a["head_dim"], hardcoded["head_dim"]), | |
| ("qk_rope_head_dim", a["qk_rope_head_dim"], hardcoded["qk_rope_head_dim"]), | |
| ("qk_nope_head_dim", a["qk_nope_head_dim"], hardcoded["qk_nope_head_dim"]), | |
| ("kv_lora_rank", a["kv_lora_rank"], hardcoded["kv_lora_rank"]), | |
| ("ff_mult", round(a["ff_mult"], 4), hardcoded["ff_mult"]), | |
| ("qk_norm", a["qk_norm"], hardcoded["qk_norm"]), | |
| ("vocab_size", a["vocab_size"], hardcoded["vocab_size"]), | |
| ] | |
| print(f" {'field':<20} {'checkpoint':>14} {'hardcoded':>14}") | |
| print(f" {'-'*20} {'-'*14} {'-'*14}") | |
| n_mismatch = 0 | |
| for name, ck, hc in rows: | |
| marker = "" if ck == hc else " β MISMATCH" | |
| if ck != hc: | |
| n_mismatch += 1 | |
| print(f" {name:<20} {str(ck):>14} {str(hc):>14}{marker}") | |
| if n_mismatch: | |
| print(f" β οΈ {n_mismatch} field(s) differ from the hardcoded constants.") | |
| else: | |
| print(f" β Checkpoint matches hardcoded constants exactly.") | |
| # ============================================================================= | |
| # MODEL LOADING | |
| # ============================================================================= | |
| def load_model(checkpoint_dir, checkpoint_file, tokenizer, device, dtype): | |
| global D_MODEL, N_LAYERS, N_HEADS, KV_LORA_RANK, QK_ROPE_HEAD_DIM | |
| global FF_MULT, QK_NORM, HEAD_DIM, QK_NOPE_HEAD_DIM | |
| ckpt_path = os.path.join(checkpoint_dir, checkpoint_file) | |
| print(f"\n{'='*68}\nπ§ Loading base model\n{'='*68}") | |
| if not os.path.exists(ckpt_path): | |
| avail = [] | |
| if os.path.isdir(checkpoint_dir): | |
| avail = sorted(f for f in os.listdir(checkpoint_dir) if f.endswith(".pt")) | |
| raise FileNotFoundError( | |
| f"Checkpoint not found: {ckpt_path}\n" | |
| f" .pt files in {checkpoint_dir}: {avail if avail else '(none)'}" | |
| ) | |
| print(f" File: {ckpt_path}") | |
| ckpt = torch.load(ckpt_path, map_location="cpu") | |
| state_dict = ckpt["model"] if isinstance(ckpt, dict) and "model" in ckpt else ckpt | |
| state_dict = {k.replace("_orig_mod.", ""): v for k, v in state_dict.items()} | |
| for key in ("step", "global_step", "epoch", "best_val_loss", "val_loss", "accuracy"): | |
| if isinstance(ckpt, dict) and key in ckpt: | |
| print(f" {key}: {ckpt[key]}") | |
| # ββ Verify / adopt architecture from the actual tensors | |
| print(f"\n π Architecture read from checkpoint tensors:") | |
| arch = infer_arch_from_state_dict(state_dict) | |
| _print_arch_table(arch, { | |
| "d_model": D_MODEL, "n_layers": N_LAYERS, "n_heads": N_HEADS, | |
| "head_dim": HEAD_DIM, "qk_rope_head_dim": QK_ROPE_HEAD_DIM, | |
| "qk_nope_head_dim": QK_NOPE_HEAD_DIM, "kv_lora_rank": KV_LORA_RANK, | |
| "ff_mult": FF_MULT, "qk_norm": QK_NORM, "vocab_size": len(tokenizer), | |
| }) | |
| if AUTO_ADOPT_CHECKPOINT_ARCH: | |
| D_MODEL = arch["d_model"] | |
| N_LAYERS = arch["n_layers"] | |
| N_HEADS = arch["n_heads"] | |
| KV_LORA_RANK = arch["kv_lora_rank"] | |
| QK_ROPE_HEAD_DIM = arch["qk_rope_head_dim"] | |
| FF_MULT = arch["ff_mult"] | |
| QK_NORM = arch["qk_norm"] | |
| HEAD_DIM = arch["head_dim"] | |
| QK_NOPE_HEAD_DIM = arch["qk_nope_head_dim"] | |
| print(f" βͺ AUTO_ADOPT_CHECKPOINT_ARCH=True β building from checkpoint dims.") | |
| vocab_size = arch["vocab_size"] | |
| if vocab_size != len(tokenizer): | |
| print(f" β οΈ Checkpoint vocab {vocab_size:,} != tokenizer {len(tokenizer):,}. " | |
| f"Building embedding at checkpoint size; token ids above " | |
| f"{vocab_size-1} will be out of range.") | |
| mcfg = ModelConfig( | |
| vocab_size=vocab_size, d_model=D_MODEL, n_layers=N_LAYERS, n_heads=N_HEADS, | |
| kv_lora_rank=KV_LORA_RANK, qk_rope_head_dim=QK_ROPE_HEAD_DIM, | |
| ff_hidden_mult=FF_MULT, qk_norm=QK_NORM, max_seq_len=CONTEXT_LEN, | |
| ) | |
| model = DenseLLM(mcfg) | |
| missing, unexpected = model.load_state_dict(state_dict, strict=False) | |
| if missing: | |
| raise RuntimeError(f"{len(missing)} missing keys β architecture mismatch. " | |
| f"First few: {missing[:6]}") | |
| if unexpected: | |
| print(f" β οΈ {len(unexpected)} unexpected keys ignored: {unexpected[:6]}") | |
| model = model.to(device=device, dtype=dtype).eval() | |
| for p in model.parameters(): | |
| p.requires_grad_(False) | |
| del ckpt, state_dict | |
| gc.collect() | |
| if torch.cuda.is_available(): | |
| torch.cuda.empty_cache() | |
| total = sum(p.numel() for p in model.parameters()) | |
| param_str = f"{total/1e9:.2f}B" if total >= 1e9 else f"{total/1e6:.0f}M" | |
| cache_d = N_HEADS * HEAD_DIM * 2 # this impl caches full k and v | |
| print(f"\n β {total:,} params ({param_str})") | |
| print(f" β {D_MODEL}d Γ {N_LAYERS}L Γ {N_HEADS}H | " | |
| f"head={HEAD_DIM} = nope({QK_NOPE_HEAD_DIM}) + rope({QK_ROPE_HEAD_DIM})") | |
| print(f" β kv_lora_rank={KV_LORA_RANK} | ff_inner={int(FF_MULT*D_MODEL)} | " | |
| f"qk_norm={QK_NORM}") | |
| print(f" β Decode cache: {cache_d}d/token/layer | RMSNorm: {_rmsnorm_source}") | |
| print(f" β Device: {device} | dtype: {dtype}") | |
| print(f"{'='*68}\n") | |
| return model, param_str | |
| # ============================================================================= | |
| # SAMPLING HELPERS | |
| # ============================================================================= | |
| def _apply_rep_penalty(logits: torch.Tensor, seen_ids: torch.Tensor, penalty: float): | |
| """In-place repetition penalty over already-emitted ids. logits: (B, V).""" | |
| if penalty == 1.0: | |
| return logits | |
| for b in range(logits.shape[0]): | |
| uniq = torch.unique(seen_ids[b]) | |
| s = logits[b, uniq] | |
| logits[b, uniq] = torch.where(s > 0, s / penalty, s * penalty) | |
| return logits | |
| def _filter_and_sample(logits, temperature, top_k, top_p): | |
| """logits: (B, V) β next tokens (B, 1).""" | |
| if temperature <= 0: | |
| return torch.argmax(logits, dim=-1, keepdim=True) | |
| logits = logits / max(temperature, 1e-8) | |
| if top_k and top_k > 0: | |
| v, _ = torch.topk(logits, min(int(top_k), logits.size(-1))) | |
| logits = logits.masked_fill(logits < v[:, [-1]], -float("inf")) | |
| if top_p and top_p < 1.0: | |
| sorted_logits, sorted_idx = torch.sort(logits, descending=True) | |
| cum = torch.cumsum(F.softmax(sorted_logits, dim=-1), dim=-1) | |
| remove = cum > top_p | |
| remove[:, 1:] = remove[:, :-1].clone() | |
| remove[:, 0] = False | |
| mask = torch.zeros_like(logits, dtype=torch.bool).scatter_(1, sorted_idx, remove) | |
| logits = logits.masked_fill(mask, -float("inf")) | |
| probs = F.softmax(logits, dim=-1) | |
| return torch.multinomial(probs, num_samples=1) | |
| # ============================================================================= | |
| # STREAMING GENERATION (batch of 1, KV cache) | |
| # ============================================================================= | |
| @torch.inference_mode() | |
| def generate_streaming(model, tokenizer, prompt, device, dtype, | |
| max_new_tokens=512, temperature=0.8, top_k=50, top_p=0.95, | |
| repetition_penalty=1.05, seed=-1, yield_every=4): | |
| """Yields (text_so_far, n_tokens, ended_with_eos, tok_per_sec, hit_limit).""" | |
| if seed is not None and seed >= 0: | |
| torch.manual_seed(int(seed)) | |
| random.seed(int(seed)) | |
| eos_id = tokenizer.doc_eos_token_id | |
| ids = tokenizer(prompt, return_tensors="pt")["input_ids"].to(device) | |
| if ids.shape[1] >= CONTEXT_LEN: | |
| ids = ids[:, -(CONTEXT_LEN - 8):] | |
| prompt_len = ids.shape[1] | |
| t0 = time.time() | |
| logits, past_kvs = model(ids, use_cache=True) | |
| next_logits = logits[:, -1, :].float() | |
| gen_ids: List[int] = [] | |
| ended_with_eos = False | |
| hit_limit = False | |
| for i in range(int(max_new_tokens)): | |
| if prompt_len + len(gen_ids) >= CONTEXT_LEN: | |
| hit_limit = True | |
| break | |
| seen = torch.cat([ids, torch.tensor([gen_ids], dtype=torch.long, device=device)], dim=1) \ | |
| if gen_ids else ids | |
| next_logits = _apply_rep_penalty(next_logits, seen, float(repetition_penalty)) | |
| next_token = _filter_and_sample(next_logits, temperature, top_k, top_p) | |
| tid = int(next_token.item()) | |
| if tid == eos_id: | |
| ended_with_eos = True | |
| break | |
| gen_ids.append(tid) | |
| if (i % yield_every == 0) or (i == int(max_new_tokens) - 1): | |
| text = tokenizer.decode(gen_ids, skip_special_tokens=True) | |
| yield text, len(gen_ids), False, len(gen_ids) / max(time.time() - t0, 1e-9), False | |
| logits, past_kvs = model(next_token, past_kvs=past_kvs, use_cache=True) | |
| next_logits = logits[:, -1, :].float() | |
| else: | |
| hit_limit = True | |
| del past_kvs | |
| text = tokenizer.decode(gen_ids, skip_special_tokens=True) | |
| yield text, len(gen_ids), ended_with_eos, len(gen_ids) / max(time.time() - t0, 1e-9), hit_limit | |
| # ============================================================================= | |
| # BATCHED GENERATION (maj@k self-consistency) | |
| # ============================================================================= | |
| @torch.inference_mode() | |
| def generate_batch(model, tokenizer, prompt, device, dtype, k=8, | |
| max_new_tokens=512, temperature=0.8, top_k=50, top_p=0.95, | |
| repetition_penalty=1.05, seed=-1): | |
| """Generate k independent samples in parallel. Returns list of dicts.""" | |
| if seed is not None and seed >= 0: | |
| torch.manual_seed(int(seed)) | |
| eos_id = tokenizer.doc_eos_token_id | |
| prompt_ids = tokenizer(prompt, return_tensors="pt")["input_ids"][0].tolist() | |
| if len(prompt_ids) >= CONTEXT_LEN: | |
| prompt_ids = prompt_ids[-(CONTEXT_LEN - 8):] | |
| prompt_len = len(prompt_ids) | |
| k = int(k) | |
| ids = torch.tensor([prompt_ids] * k, dtype=torch.long, device=device) | |
| finished = torch.zeros(k, dtype=torch.bool, device=device) | |
| finished_at_eos = torch.zeros(k, dtype=torch.bool, device=device) | |
| all_gen: List[List[int]] = [[] for _ in range(k)] | |
| t0 = time.time() | |
| logits, past_kvs = model(ids, use_cache=True) | |
| next_logits = logits[:, -1, :].float() | |
| for step in range(int(max_new_tokens)): | |
| if bool(finished.all()) or (prompt_len + step) >= CONTEXT_LEN: | |
| break | |
| if repetition_penalty != 1.0 and step > 0: | |
| seen = torch.tensor( | |
| [prompt_ids + g + [eos_id] * (step - len(g)) for g in all_gen], | |
| dtype=torch.long, device=device) | |
| next_logits = _apply_rep_penalty(next_logits, seen, float(repetition_penalty)) | |
| next_tokens = _filter_and_sample(next_logits, temperature, top_k, top_p) | |
| next_tokens[finished] = eos_id | |
| for g in range(k): | |
| if not finished[g]: | |
| tid = int(next_tokens[g].item()) | |
| if tid != eos_id: | |
| all_gen[g].append(tid) | |
| just_done = (next_tokens.squeeze(-1) == eos_id) & ~finished | |
| finished_at_eos |= just_done | |
| finished |= just_done | |
| logits, past_kvs = model(next_tokens, past_kvs=past_kvs, use_cache=True) | |
| next_logits = logits[:, -1, :].float() | |
| del past_kvs | |
| gc.collect() | |
| if torch.cuda.is_available(): | |
| torch.cuda.empty_cache() | |
| elapsed = time.time() - t0 | |
| total_tok = sum(len(g) for g in all_gen) | |
| results = [] | |
| for g in range(k): | |
| text = tokenizer.decode(all_gen[g], skip_special_tokens=True) | |
| results.append({ | |
| "idx": g, | |
| "text": text, | |
| "n_tokens": len(all_gen[g]), | |
| "ended_with_eos": bool(finished_at_eos[g].item()), | |
| "answer": extract_answer(text), | |
| }) | |
| return results, elapsed, total_tok / max(elapsed, 1e-9) | |
| # ============================================================================= | |
| # ANSWER EXTRACTION (same logic as the training/eval scripts) | |
| # ============================================================================= | |
| def extract_answer(text: str) -> Optional[str]: | |
| if not text: | |
| return None | |
| boxed = [] | |
| for m in re.finditer(r'\\boxed\{', text): | |
| start = m.end() | |
| depth, i = 1, start | |
| while i < len(text) and depth > 0: | |
| if text[i] == '{': depth += 1 | |
| elif text[i] == '}': depth -= 1 | |
| i += 1 | |
| if depth == 0: | |
| boxed.append(text[start:i-1].strip()) | |
| if boxed: | |
| return boxed[-1] | |
| m = re.search(r'(?:the\s+)?answer\s+is\s*[:\s]*([+-]?\d+(?:\.\d+)?(?:/\d+)?)', text, re.IGNORECASE) | |
| if m: return m.group(1).strip() | |
| m = re.search(r'answer\s*[:=]\s*([+-]?\d+(?:\.\d+)?(?:/\d+)?)', text, re.IGNORECASE) | |
| if m: return m.group(1).strip() | |
| eqs = re.findall(r'=\s*([+-]?\d+(?:\.\d+)?(?:/\d+)?)\s*[.\s]*$', text, re.MULTILINE) | |
| if eqs: return eqs[-1].strip() | |
| nums = re.findall(r'(?<![.\w])([+-]?\d+(?:\.\d+)?(?:/\d+)?)(?![.\w])', text) | |
| if nums: return nums[-1].strip() | |
| return None | |
| def normalize_answer(answer: Optional[str]) -> Optional[str]: | |
| if answer is None: return None | |
| a = answer.strip() | |
| a = re.sub(r'\\text\{([^}]*)\}', r'\1', a) | |
| a = re.sub(r'\\mathrm\{([^}]*)\}', r'\1', a) | |
| a = re.sub(r'[\$\\,\\;\\!\\>\\:]', '', a) | |
| a = a.replace(' ', '').replace(',', '') | |
| a = re.sub(r'%$', '', a) | |
| a = re.sub(r'\^\{?circ\}?$', '', a) | |
| a = re.sub(r'Β°$', '', a).strip() | |
| try: | |
| if '/' in a: | |
| parts = a.split('/') | |
| if len(parts) == 2: | |
| val = float(parts[0]) / float(parts[1]) | |
| return str(int(val)) if val == int(val) else f"{val:.10f}".rstrip('0').rstrip('.') | |
| val = float(a) | |
| return str(int(val)) if val == int(val) else f"{val:.10f}".rstrip('0').rstrip('.') | |
| except (ValueError, ZeroDivisionError, OverflowError): | |
| pass | |
| return a.lower().strip() | |
| # ============================================================================= | |
| # GRADIO INTERFACE | |
| # ============================================================================= | |
| def create_gradio_interface(model, tokenizer, device, dtype, param_str): | |
| import gradio as gr | |
| custom_css = """ | |
| .gradio-container { max-width: 1080px !important; margin: auto !important; } | |
| .mono-out textarea { | |
| font-family: 'JetBrains Mono','SF Mono','Courier New',monospace !important; | |
| font-size: 13.5px !important; line-height: 1.65 !important; | |
| } | |
| .metric-strip { | |
| font-family: 'JetBrains Mono','SF Mono',monospace; font-size: 13px; | |
| padding: 6px 2px; opacity: 0.9; | |
| } | |
| """ | |
| ff_inner = int(FF_MULT * D_MODEL) | |
| def _metrics_md(answer, n_tok, eos, tps, hit_limit): | |
| ans = f"`{answer}`" if answer else "β" | |
| stop = "EOS β " if eos else ("length cap β" if hit_limit else "β¦running") | |
| return (f"<div class='metric-strip'>" | |
| f"<b>answer:</b> {ans} Β· " | |
| f"<b>tokens:</b> {n_tok} Β· " | |
| f"<b>stop:</b> {stop} Β· " | |
| f"<b>speed:</b> {tps:.1f} tok/s</div>") | |
| def build_prompt(raw, template_name): | |
| tmpl = PROMPT_TEMPLATES.get(template_name, "{p}") | |
| return tmpl.format(p=raw.strip()) | |
| # ββ Tab 1: single completion, streaming | |
| def run_generate(raw_prompt, template_name, max_new, temp, tk, tp, rep, seed, echo): | |
| if not raw_prompt.strip(): | |
| yield "Enter a prompt first.", _metrics_md(None, 0, False, 0.0, False) | |
| return | |
| prompt = build_prompt(raw_prompt, template_name) | |
| head = prompt if echo else "" | |
| last = ("", 0, False, 0.0, False) | |
| for text, n_tok, eos, tps, hit in generate_streaming( | |
| model, tokenizer, prompt, device, dtype, | |
| max_new_tokens=max_new, temperature=temp, top_k=tk, top_p=tp, | |
| repetition_penalty=rep, seed=seed): | |
| last = (text, n_tok, eos, tps, hit) | |
| yield head + text, _metrics_md(None, n_tok, eos, tps, hit) | |
| text, n_tok, eos, tps, hit = last | |
| yield head + text, _metrics_md(extract_answer(text), n_tok, eos, tps, hit) | |
| # ββ Tab 2: maj@k self-consistency | |
| def run_majk(raw_prompt, template_name, k, max_new, temp, tk, tp, rep, seed): | |
| if not raw_prompt.strip(): | |
| return "Enter a prompt first.", "" | |
| prompt = build_prompt(raw_prompt, template_name) | |
| results, elapsed, tps = generate_batch( | |
| model, tokenizer, prompt, device, dtype, k=k, max_new_tokens=max_new, | |
| temperature=temp, top_k=tk, top_p=tp, repetition_penalty=rep, seed=seed) | |
| norm_counts = Counter() | |
| display_for_norm = {} | |
| for r in results: | |
| n = normalize_answer(r["answer"]) | |
| if n is None: | |
| continue | |
| norm_counts[n] += 1 | |
| display_for_norm.setdefault(n, r["answer"]) | |
| lines = [f"### maj@{int(k)} vote", ""] | |
| if norm_counts: | |
| top_norm, votes = norm_counts.most_common(1)[0] | |
| lines.append(f"**Consensus: `{display_for_norm[top_norm]}`** " | |
| f"({votes}/{int(k)} samples, {votes/int(k)*100:.0f}%)") | |
| lines.append("") | |
| lines.append("| answer | votes |") | |
| lines.append("|---|---|") | |
| for n, c in norm_counts.most_common(): | |
| lines.append(f"| `{display_for_norm[n]}` | {c} |") | |
| else: | |
| lines.append("No answer could be extracted from any sample.") | |
| n_eos = sum(1 for r in results if r["ended_with_eos"]) | |
| avg_len = sum(r["n_tokens"] for r in results) / max(len(results), 1) | |
| lines += ["", f"_{n_eos}/{int(k)} terminated on EOS Β· avg {avg_len:.0f} tokens Β· " | |
| f"{elapsed:.1f}s Β· {tps:.0f} tok/s aggregate_"] | |
| blocks = [] | |
| for r in results: | |
| flag = "EOS" if r["ended_with_eos"] else "cut" | |
| blocks.append(f"{'='*64}\n# sample {r['idx']+1}/{int(k)} Β· {r['n_tokens']} tok Β· " | |
| f"{flag} Β· answer: {r['answer']}\n{'='*64}\n{r['text']}\n") | |
| return "\n".join(lines), "\n".join(blocks) | |
| def set_preset(name): | |
| return { | |
| "Precise": (0.3, 40, 0.90, 1.05), | |
| "Balanced": (0.8, 50, 0.95, 1.05), | |
| "Creative": (1.0, 0, 0.98, 1.10), | |
| "Greedy": (0.0, 0, 1.00, 1.00), | |
| }.get(name, (0.8, 50, 0.95, 1.05)) | |
| with gr.Blocks(title="THE MATHEMATICIAN V2 β base", | |
| css=custom_css, | |
| theme=gr.themes.Soft(primary_hue="indigo", secondary_hue="blue")) as demo: | |
| gr.Markdown( | |
| f""" | |
| # π§ THE MATHEMATICIAN V2 β base ({param_str}) | |
| Pretrained base checkpoint, no instruction tuning and no RL. It completes | |
| text; it does not follow instructions. Give it a problem statement and let | |
| the template supply the continuation cue. | |
| `{D_MODEL}d Γ {N_LAYERS}L Γ {N_HEADS}H` Β· MLA rank-`{KV_LORA_RANK}` Β· | |
| head `{HEAD_DIM}` = nope(`{QK_NOPE_HEAD_DIM}`) + rope(`{QK_ROPE_HEAD_DIM}`) Β· | |
| ff_inner `{ff_inner}` Β· ctx `{CONTEXT_LEN}` Β· `{device}`/`{str(dtype).replace('torch.','')}` | |
| --- | |
| """ | |
| ) | |
| with gr.Row(): | |
| with gr.Column(scale=3): | |
| prompt_input = gr.Textbox( | |
| label="Prompt", lines=4, max_lines=12, | |
| placeholder="Find the remainder when 2^100 is divided by 7.", | |
| ) | |
| template_dd = gr.Dropdown( | |
| choices=list(PROMPT_TEMPLATES.keys()), value=DEFAULT_TEMPLATE, | |
| label="Prompt template", | |
| ) | |
| with gr.Column(scale=2): | |
| preset_dd = gr.Dropdown( | |
| choices=["Balanced", "Precise", "Creative", "Greedy"], | |
| value="Balanced", label="Preset", | |
| ) | |
| temp_sl = gr.Slider(0.0, 2.0, DEFAULT_TEMPERATURE, step=0.05, label="temperature") | |
| topk_sl = gr.Slider(0, 200, DEFAULT_TOP_K, step=5, label="top_k (0 = off)") | |
| topp_sl = gr.Slider(0.0, 1.0, DEFAULT_TOP_P, step=0.01, label="top_p") | |
| rep_sl = gr.Slider(1.0, 1.5, DEFAULT_REP_PENALTY, step=0.01, | |
| label="repetition penalty") | |
| maxtok_sl = gr.Slider(16, 2048, DEFAULT_MAX_NEW_TOKENS, step=16, | |
| label="max new tokens") | |
| seed_num = gr.Number(value=-1, precision=0, label="seed (-1 = random)") | |
| with gr.Tabs(): | |
| with gr.Tab("Completion"): | |
| echo_cb = gr.Checkbox(value=True, label="Echo prompt in output") | |
| out_box = gr.Textbox(label="Output", lines=20, max_lines=40, | |
| interactive=False, elem_classes=["mono-out"]) | |
| metrics_md = gr.Markdown(_metrics_md(None, 0, False, 0.0, False)) | |
| with gr.Row(): | |
| gen_btn = gr.Button("Generate", variant="primary", scale=2) | |
| stop_btn = gr.Button("Stop", variant="stop", scale=1) | |
| clear_btn = gr.Button("Clear", scale=1) | |
| gr.Examples( | |
| examples=[ | |
| "Find the remainder when 2^100 is divided by 7.", | |
| "In triangle ABC, AB = 13, BC = 14, and CA = 15. Find the area of triangle ABC.", | |
| "If x + y = 7 and xy = 10, find the value of x^3 + y^3.", | |
| "How many ways can the letters in the word BANANA be arranged?", | |
| "Find the sum of all positive divisors of 360.", | |
| "Find the number of trailing zeros in 50!.", | |
| "Let f(x) = x^2 - 4x + 3. Find the sum of all integers n such that f(f(n)) = 3.", | |
| ], | |
| inputs=prompt_input, | |
| ) | |
| with gr.Tab("maj@k (self-consistency)"): | |
| gr.Markdown( | |
| "Samples k completions in parallel from the same prompt and takes a " | |
| "majority vote over extracted answers. Temperature must be > 0 or " | |
| "every sample will be identical." | |
| ) | |
| k_sl = gr.Slider(2, 32, DEFAULT_MAJ_K, step=1, label="k samples") | |
| majk_btn = gr.Button("Run maj@k", variant="primary") | |
| vote_md = gr.Markdown() | |
| samples_box = gr.Textbox(label="All samples", lines=24, max_lines=60, | |
| interactive=False, elem_classes=["mono-out"]) | |
| with gr.Accordion("Model card", open=False): | |
| gr.Markdown( | |
| f""" | |
| | field | value | | |
| |---|---| | |
| | checkpoint | `{os.path.join(CHECKPOINT_DIR, CHECKPOINT_FILE)}` | | |
| | parameters | {param_str} | | |
| | d_model | {D_MODEL} | | |
| | layers | {N_LAYERS} | | |
| | heads | {N_HEADS} | | |
| | head_dim | {HEAD_DIM} = nope({QK_NOPE_HEAD_DIM}) + rope({QK_ROPE_HEAD_DIM}) | | |
| | kv_lora_rank | {KV_LORA_RANK} | | |
| | ff_mult | {FF_MULT} (inner {ff_inner}) | | |
| | qk_norm | {QK_NORM} (post-assembly, full head_dim) | | |
| | k_rope_proj | per-head ({N_HEADS} Γ {QK_ROPE_HEAD_DIM}) | | |
| | rope base | {ROPE_BASE:,.0f} | | |
| | context | {CONTEXT_LEN:,} | | |
| | tokenizer | {VOCAB_NAME} ({len(tokenizer):,} tokens) | | |
| | RMSNorm | {_rmsnorm_source} | | |
| **Notes on this variant.** `q_proj` is a single full projection split at | |
| runtime rather than the canonical q-compression path, `k_rope_proj` is | |
| per-head rather than shared-and-broadcast, and `k_up`/`v_up` are separate | |
| rather than a fused `kv_b_proj`. Because K is materialised in full for | |
| SDPA, decode caches `n_heads Γ head_dim Γ 2` per token per layer β the | |
| latent cache saving isn't realised in this implementation, which is fine | |
| at this scale but worth remembering when comparing memory numbers. | |
| **Base-model behaviour.** Expect textbook-style continuations, section | |
| headers, and no reliable `\\boxed{{}}` + EOS discipline β that comes from | |
| the SFT/GRPO stages. Repetition penalty above ~1.1 tends to hurt math | |
| output because formulas legitimately repeat tokens. | |
| """ | |
| ) | |
| gen_inputs = [prompt_input, template_dd, maxtok_sl, temp_sl, topk_sl, | |
| topp_sl, rep_sl, seed_num, echo_cb] | |
| gen_event = gen_btn.click(fn=run_generate, inputs=gen_inputs, | |
| outputs=[out_box, metrics_md]) | |
| submit_event = prompt_input.submit(fn=run_generate, inputs=gen_inputs, | |
| outputs=[out_box, metrics_md]) | |
| stop_btn.click(fn=None, inputs=None, outputs=None, | |
| cancels=[gen_event, submit_event]) | |
| clear_btn.click(fn=lambda: ("", "", _metrics_md(None, 0, False, 0.0, False)), | |
| outputs=[prompt_input, out_box, metrics_md]) | |
| preset_dd.change(fn=set_preset, inputs=preset_dd, | |
| outputs=[temp_sl, topk_sl, topp_sl, rep_sl]) | |
| majk_btn.click( | |
| fn=run_majk, | |
| inputs=[prompt_input, template_dd, k_sl, maxtok_sl, temp_sl, topk_sl, | |
| topp_sl, rep_sl, seed_num], | |
| outputs=[vote_md, samples_box], | |
| ) | |
| return demo | |
| # ============================================================================= | |
| # CLI MODE | |
| # ============================================================================= | |
| def run_cli(model, tokenizer, device, dtype, param_str): | |
| print(f"\n{'='*68}\nπ¬ CLI mode β THE MATHEMATICIAN V2 base ({param_str})") | |
| print(f" Template: {DEFAULT_TEMPLATE}") | |
| print(f" Commands: /quit /temp <f> /max <int> /raw (toggle template)") | |
| print(f"{'='*68}\n") | |
| temp = DEFAULT_TEMPERATURE | |
| max_new = DEFAULT_MAX_NEW_TOKENS | |
| use_template = True | |
| while True: | |
| try: | |
| raw = input("\n> ").strip() | |
| except (EOFError, KeyboardInterrupt): | |
| print("\nbye") | |
| return | |
| if not raw: | |
| continue | |
| if raw in ("/quit", "/exit", "/q"): | |
| print("bye") | |
| return | |
| if raw.startswith("/temp"): | |
| try: | |
| temp = float(raw.split()[1]); print(f" temperature = {temp}") | |
| except Exception: | |
| print(" usage: /temp 0.7") | |
| continue | |
| if raw.startswith("/max"): | |
| try: | |
| max_new = int(raw.split()[1]); print(f" max_new_tokens = {max_new}") | |
| except Exception: | |
| print(" usage: /max 512") | |
| continue | |
| if raw == "/raw": | |
| use_template = not use_template | |
| print(f" template = {'ON (' + DEFAULT_TEMPLATE + ')' if use_template else 'OFF'}") | |
| continue | |
| prompt = PROMPT_TEMPLATES[DEFAULT_TEMPLATE].format(p=raw) if use_template else raw | |
| print("-" * 68) | |
| printed = 0 | |
| final = ("", 0, False, 0.0, False) | |
| for text, n_tok, eos, tps, hit in generate_streaming( | |
| model, tokenizer, prompt, device, dtype, | |
| max_new_tokens=max_new, temperature=temp, top_k=DEFAULT_TOP_K, | |
| top_p=DEFAULT_TOP_P, repetition_penalty=DEFAULT_REP_PENALTY): | |
| print(text[printed:], end="", flush=True) | |
| printed = len(text) | |
| final = (text, n_tok, eos, tps, hit) | |
| text, n_tok, eos, tps, hit = final | |
| stop = "EOS" if eos else ("length cap" if hit else "?") | |
| print(f"\n{'-'*68}") | |
| print(f"answer={extract_answer(text)} | {n_tok} tok | stop={stop} | {tps:.1f} tok/s") | |
| # ============================================================================= | |
| # MAIN | |
| # ============================================================================= | |
| def main(): | |
| print(f""" | |
| ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| β π§ THE MATHEMATICIAN V2 β BASE (PRETRAINED) INFERENCE β | |
| β older MLA variant Β· per-head k_rope Β· split q_proj Β· no RL β | |
| ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| """) | |
| device = "cuda" if torch.cuda.is_available() else "cpu" | |
| if device == "cuda": | |
| dtype = torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float16 | |
| torch.set_float32_matmul_precision("high") | |
| torch.backends.cuda.matmul.allow_tf32 = True | |
| torch.backends.cudnn.allow_tf32 = True | |
| props = torch.cuda.get_device_properties(0) | |
| print(f"π₯οΈ {props.name} ({props.total_memory/(1024**3):.1f}GB)") | |
| else: | |
| dtype = torch.float32 | |
| print("β οΈ CPU only β generation will be slow.") | |
| tokenizer = get_tokenizer(VOCAB_NAME, DOC_EOS_TOKEN, VOCAB_PAD_MULTIPLE, CHECKPOINT_DIR) | |
| model, param_str = load_model(CHECKPOINT_DIR, CHECKPOINT_FILE, tokenizer, device, dtype) | |
| if RUN_CLI: | |
| run_cli(model, tokenizer, device, dtype, param_str) | |
| return | |
| print(f"π Launching Gradio on port {GRADIO_PORT} (share={GRADIO_SHARE})\n") | |
| demo = create_gradio_interface(model, tokenizer, device, dtype, param_str) | |
| demo.queue().launch(share=GRADIO_SHARE, server_port=GRADIO_PORT, | |
| debug=False, show_error=True) | |
| if __name__ == "__main__": | |
| main() | |
| ``` | |