Download common.py from usrnotfound101/anlp-a2-decoding: direct link, hf CLI and curl.
- Browser
- Download file 34.2 kB
-
https://huggingface.co/usrnotfound101/anlp-a2-decoding/resolve/main/common.py
- Command line
-
hf download hf://usrnotfound101/anlp-a2-decoding/common.py
-
curl -L -o common.py https://huggingface.co/usrnotfound101/anlp-a2-decoding/resolve/main/common.py
34.2 kB
| # ============================================================================= | |
| # common.py - shared utilities for ANLP Assignment 2 | |
| # * device / AMP / seeding helpers | |
| # * Hugging Face Hub helpers (token from Kaggle secrets, push, download) | |
| # * decoder-only Transformer with pluggable FFN (dense MLP or MoE) | |
| # * KV-cache + batched greedy/sampled generation (uses model.forward only) | |
| # ============================================================================= | |
| import os, math, json, time, random, contextlib, shutil, re, sys, glob, hashlib | |
| from dataclasses import dataclass, asdict, field, fields | |
| from typing import Optional, List, Dict | |
| import numpy as np | |
| import torch | |
| import torch.nn as nn | |
| from torch.nn import functional as F | |
| # ----------------------------------------------------------------------------- | |
| # basic helpers | |
| # ----------------------------------------------------------------------------- | |
| def seed_everything(seed: int): | |
| random.seed(seed); np.random.seed(seed); torch.manual_seed(seed) | |
| if torch.cuda.is_available(): | |
| torch.cuda.manual_seed_all(seed) | |
| def get_device(): | |
| return torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| def amp_dtype(device): | |
| """bf16 on Ampere+ (A100/H100/L4), fp16 on Kaggle's T4/P100, None on CPU.""" | |
| if device.type != "cuda": | |
| return None | |
| major, _ = torch.cuda.get_device_capability(device) | |
| return torch.bfloat16 if major >= 8 else torch.float16 | |
| def autocast_ctx(device, dtype): | |
| if dtype is None: | |
| return contextlib.nullcontext() | |
| return torch.autocast(device_type="cuda", dtype=dtype) | |
| def make_grad_scaler(device, dtype): | |
| enabled = (device.type == "cuda" and dtype == torch.float16) | |
| try: | |
| return torch.amp.GradScaler("cuda", enabled=enabled) | |
| except Exception: # older torch | |
| return torch.cuda.amp.GradScaler(enabled=enabled) | |
| def fmt_num(n): | |
| for unit, div in (("B", 1e9), ("M", 1e6), ("K", 1e3)): | |
| if abs(n) >= div: | |
| return f"{n/div:.2f}{unit}" | |
| return str(n) | |
| def save_json(obj, path): | |
| os.makedirs(os.path.dirname(os.path.abspath(path)), exist_ok=True) | |
| with open(path, "w", encoding="utf-8") as f: | |
| json.dump(obj, f, indent=2, ensure_ascii=False, default=_json_default) | |
| def load_json(path): | |
| with open(path, encoding="utf-8") as f: | |
| return json.load(f) | |
| def _json_default(o): | |
| if isinstance(o, (np.integer,)): return int(o) | |
| if isinstance(o, (np.floating,)): return float(o) | |
| if isinstance(o, np.ndarray): return o.tolist() | |
| if isinstance(o, torch.Tensor): return o.tolist() | |
| if isinstance(o, torch.dtype): return str(o) | |
| return str(o) | |
| def gpu_name(): | |
| return torch.cuda.get_device_name(0) if torch.cuda.is_available() else "cpu" | |
| # ----------------------------------------------------------------------------- | |
| # Hugging Face Hub helpers | |
| # ----------------------------------------------------------------------------- | |
| TOKENIZER_REPO = "anlp-a2-tokenizers" # holds mt_spm.model and lm_spm.model | |
| def get_hf_token(): | |
| """HF token from env var HF_TOKEN or Kaggle secret named HF_TOKEN.""" | |
| tok = os.environ.get("HF_TOKEN") | |
| if tok: | |
| return tok | |
| try: | |
| from kaggle_secrets import UserSecretsClient | |
| tok = UserSecretsClient().get_secret("HF_TOKEN") | |
| if tok: | |
| os.environ["HF_TOKEN"] = tok | |
| return tok | |
| except Exception: | |
| pass | |
| return None | |
| def hf_login(token): | |
| if not token: | |
| print("[hf] no HF token found - add a Kaggle secret named HF_TOKEN (write access).") | |
| return | |
| try: | |
| from huggingface_hub import login | |
| login(token=token, add_to_git_credential=False) | |
| except Exception as e: | |
| print("[hf] login failed:", e) | |
| def hf_whoami(token): | |
| from huggingface_hub import HfApi | |
| return HfApi(token=token).whoami()["name"] | |
| def hf_resolve_user(cfg_user, token): | |
| if cfg_user: | |
| return cfg_user | |
| if token: | |
| try: | |
| return hf_whoami(token) | |
| except Exception as e: | |
| print("[hf] whoami failed:", e) | |
| return None | |
| def _retry(fn, tries=4, wait=10): | |
| last = None | |
| for i in range(tries): | |
| try: | |
| return fn() | |
| except Exception as e: | |
| last = e | |
| print(f"[hf] attempt {i+1}/{tries} failed: {e}") | |
| time.sleep(wait * (i + 1)) | |
| raise last | |
| def hf_create_repo(repo_id, token, private=False, repo_type="model"): | |
| from huggingface_hub import HfApi | |
| api = HfApi(token=token) | |
| _retry(lambda: api.create_repo(repo_id, private=private, exist_ok=True, repo_type=repo_type)) | |
| return api | |
| def hf_push_folder(folder, repo_id, token, private=False, message="upload", | |
| path_in_repo=None, run_as_future=False, repo_type="model"): | |
| api = hf_create_repo(repo_id, token, private, repo_type) | |
| kw = dict(folder_path=folder, repo_id=repo_id, repo_type=repo_type, | |
| commit_message=message) | |
| if path_in_repo: | |
| kw["path_in_repo"] = path_in_repo | |
| if run_as_future: | |
| return api.upload_folder(run_as_future=True, **kw) | |
| return _retry(lambda: api.upload_folder(**kw)) | |
| def hf_push_file(local_path, path_in_repo, repo_id, token, private=False, | |
| message="upload", run_as_future=False, repo_type="model"): | |
| api = hf_create_repo(repo_id, token, private, repo_type) | |
| kw = dict(path_or_fileobj=local_path, path_in_repo=path_in_repo, repo_id=repo_id, | |
| repo_type=repo_type, commit_message=message) | |
| if run_as_future: | |
| return api.upload_file(run_as_future=True, **kw) | |
| return _retry(lambda: api.upload_file(**kw)) | |
| def hf_try_download(repo_id, filename, token=None, repo_type="model", force=False): | |
| """Returns a local path, or None if the repo/file does not exist.""" | |
| try: | |
| from huggingface_hub import hf_hub_download | |
| return hf_hub_download(repo_id=repo_id, filename=filename, token=token, | |
| repo_type=repo_type, force_download=force) | |
| except Exception as e: | |
| print(f"[hf] could not download {repo_id}/{filename}: {type(e).__name__}") | |
| return None | |
| def hf_list_files(repo_id, token=None, repo_type="model"): | |
| try: | |
| from huggingface_hub import HfApi | |
| return HfApi(token=token).list_repo_files(repo_id, repo_type=repo_type) | |
| except Exception: | |
| return [] | |
| # ----------------------------------------------------------------------------- | |
| # Weights & Biases (assignment: upload training logs to WandB, public links in README/report) | |
| # Fail-soft: if wandb is missing / not logged in, training continues and logs only to JSON. | |
| # Auth: `wandb login` or WANDB_API_KEY (Kaggle secret of that name is picked up too). | |
| # WANDB_MODE=offline logs locally (sync later with `wandb sync`); cfg["wandb"]=False disables. | |
| # ----------------------------------------------------------------------------- | |
| WANDB_PROJECT = "anlp-a2" | |
| def _wandb_key(): | |
| key = os.environ.get("WANDB_API_KEY") | |
| if key: | |
| return key | |
| try: | |
| from kaggle_secrets import UserSecretsClient | |
| key = UserSecretsClient().get_secret("WANDB_API_KEY") | |
| if key: | |
| os.environ["WANDB_API_KEY"] = key | |
| return key | |
| except Exception: | |
| return None | |
| def wandb_init(cfg, name, config, tags=(), group=None, notes=None, stable_id=False): | |
| """Start (or resume, with the same stable id -> requeued jobs continue one run) a W&B run. | |
| Returns the run, or None when disabled / unavailable.""" | |
| if not cfg.get("wandb", True): | |
| return None | |
| try: | |
| import wandb | |
| _wandb_key() | |
| run = wandb.init(project=cfg.get("wandb_project") or WANDB_PROJECT, entity=cfg.get("wandb_entity"), | |
| name=name, config=config, | |
| **(dict(id=re.sub(r"[^A-Za-z0-9_-]", "-", name), resume="allow") if stable_id else {}), | |
| tags=list(tags), group=group, notes=notes, | |
| settings=wandb.Settings(init_timeout=120, console="off")) | |
| print("[wandb] run:", getattr(run, "url", None) or "(offline - `wandb sync <dir>` later)") | |
| return run | |
| except Exception as e: | |
| print(f"[wandb] disabled ({type(e).__name__}: {str(e)[:120]}) - logging to JSON only") | |
| return None | |
| def wandb_log(run, data, step=None): | |
| if run is None: | |
| return | |
| try: | |
| run.log({k: v for k, v in data.items() if v is not None}, step=step) | |
| except Exception as e: | |
| print("[wandb] log failed:", e) | |
| def wandb_images(run, paths, step=None): | |
| """paths: {key: png_path}.""" | |
| if run is None: | |
| return | |
| try: | |
| import wandb | |
| wandb_log(run, {k: wandb.Image(v) for k, v in paths.items() if os.path.exists(v)}, step) | |
| except Exception as e: | |
| print("[wandb] image log failed:", e) | |
| def wandb_finish(run, summary=None): | |
| if run is None: | |
| return | |
| try: | |
| for k, v in (summary or {}).items(): | |
| run.summary[k] = v | |
| run.finish() | |
| except Exception as e: | |
| print("[wandb] finish failed:", e) | |
| def expert_load_scalars(prefix, loads): | |
| """loads: [layer][expert] fractions -> flat {prefix/L{l}_E{e}: x} for W&B.""" | |
| return {f"{prefix}/L{l}_E{e}": float(x) for l, row in enumerate(loads) for e, x in enumerate(row)} | |
| # ----------------------------------------------------------------------------- | |
| # Model | |
| # ----------------------------------------------------------------------------- | |
| class ModelConfig: | |
| vocab_size: int = 32000 | |
| block_size: int = 256 | |
| n_layer: int = 6 | |
| n_head: int = 8 | |
| n_embd: int = 512 | |
| dropout: float = 0.0 | |
| bias: bool = False | |
| # ---- feed-forward ---- | |
| ffn_type: str = "dense" # "dense" | "moe" | |
| mlp_hidden: int = 0 # dense hidden size (0 -> 4 * n_embd) | |
| n_experts: int = 4 # number of *routed* experts | |
| n_shared: int = 0 # number of always-on shared experts | |
| top_k: int = 1 # routed experts active per token | |
| expert_hidden: int = 0 # hidden size of each expert (routed & shared) | |
| aux_loss_coef: float = 0.01 # Switch-style load-balancing loss | |
| z_loss_coef: float = 1e-3 # router z-loss (ST-MoE) | |
| def to_dict(self): | |
| return asdict(self) | |
| def from_dict(cls, d): | |
| names = {f.name for f in fields(cls)} | |
| return cls(**{k: v for k, v in d.items() if k in names}) | |
| # The five FFN variants of Part 1. H = 4 * d is the dense hidden size. | |
| # dense : 1 MLP (d -> H -> d) total = active = 2dH | |
| # moe_4e_top1 : 4 experts of hidden H/4, top-1 total = 2dH, active = 2dH/4 | |
| # moe_4e_top2 : 4 experts of hidden H/4, top-2 total = 2dH, active = 2dH/2 | |
| # moe_1s3r_top1 : 1 shared + 3 routed (hidden H/4), top-1 total = 2dH, active = 2dH/2 | |
| # moe_4e_top2_active: 4 experts of hidden H/2, top-2 total = 4dH, active = 2dH (= dense) | |
| FFN_VARIANTS = ["dense", "moe_4e_top1", "moe_4e_top2", "moe_1s3r_top1", "moe_4e_top2_active"] | |
| FFN_VARIANT_DESCRIPTIONS = { | |
| "dense": "Standard 2-layer MLP (d -> 4d -> d)", | |
| "moe_4e_top1": "MoE: 4 experts (hidden d), top-1 routing - same total params as dense", | |
| "moe_4e_top2": "MoE: 4 experts (hidden d), top-2 routing - same total params as dense", | |
| "moe_1s3r_top1": "MoE: 1 shared + 3 routed experts (hidden d), top-1 routed - same total params as dense", | |
| "moe_4e_top2_active": "MoE: 4 experts (hidden 2d), top-2 routing - same ACTIVE params as dense", | |
| } | |
| def ffn_variant_kwargs(name, n_embd): | |
| H = 4 * n_embd | |
| if name == "dense": | |
| return dict(ffn_type="dense", mlp_hidden=H) | |
| if name == "moe_4e_top1": | |
| return dict(ffn_type="moe", n_experts=4, n_shared=0, top_k=1, expert_hidden=H // 4) | |
| if name == "moe_4e_top2": | |
| return dict(ffn_type="moe", n_experts=4, n_shared=0, top_k=2, expert_hidden=H // 4) | |
| if name == "moe_1s3r_top1": | |
| return dict(ffn_type="moe", n_experts=3, n_shared=1, top_k=1, expert_hidden=H // 4) | |
| if name == "moe_4e_top2_active": | |
| return dict(ffn_type="moe", n_experts=4, n_shared=0, top_k=2, expert_hidden=H // 2) | |
| raise ValueError(f"unknown FFN variant {name}") | |
| class LayerNorm(nn.Module): | |
| def __init__(self, ndim, bias): | |
| super().__init__() | |
| self.weight = nn.Parameter(torch.ones(ndim)) | |
| self.bias = nn.Parameter(torch.zeros(ndim)) if bias else None | |
| def forward(self, x): | |
| return F.layer_norm(x, self.weight.shape, self.weight, self.bias, 1e-5) | |
| class KVCache: | |
| """Simple per-layer key/value cache. Tensors are (B, n_head, T, head_dim).""" | |
| def __init__(self, n_layer): | |
| self.k = [None] * n_layer | |
| self.v = [None] * n_layer | |
| def length(self): | |
| return 0 if self.k[0] is None else self.k[0].shape[2] | |
| def update(self, i, k, v): | |
| if self.k[i] is None: | |
| self.k[i], self.v[i] = k, v | |
| else: | |
| self.k[i] = torch.cat([self.k[i], k], dim=2) | |
| self.v[i] = torch.cat([self.v[i], v], dim=2) | |
| return self.k[i], self.v[i] | |
| def reorder(self, idx): | |
| for i in range(len(self.k)): | |
| if self.k[i] is not None: | |
| self.k[i] = self.k[i].index_select(0, idx) | |
| self.v[i] = self.v[i].index_select(0, idx) | |
| class CausalSelfAttention(nn.Module): | |
| def __init__(self, cfg: ModelConfig): | |
| super().__init__() | |
| assert cfg.n_embd % cfg.n_head == 0 | |
| self.c_attn = nn.Linear(cfg.n_embd, 3 * cfg.n_embd, bias=cfg.bias) | |
| self.c_proj = nn.Linear(cfg.n_embd, cfg.n_embd, bias=cfg.bias) | |
| self.resid_dropout = nn.Dropout(cfg.dropout) | |
| self.n_head, self.n_embd, self.dropout = cfg.n_head, cfg.n_embd, cfg.dropout | |
| def forward(self, x, attn_mask=None, cache=None, layer_idx=0): | |
| B, T, C = x.size() | |
| q, k, v = self.c_attn(x).split(self.n_embd, dim=2) | |
| hs = C // self.n_head | |
| q = q.view(B, T, self.n_head, hs).transpose(1, 2) | |
| k = k.view(B, T, self.n_head, hs).transpose(1, 2) | |
| v = v.view(B, T, self.n_head, hs).transpose(1, 2) | |
| if cache is not None: | |
| k, v = cache.update(layer_idx, k, v) | |
| dp = self.dropout if self.training else 0.0 | |
| if attn_mask is None: | |
| # plain causal attention; valid only when queries and keys are aligned | |
| assert T == 1 or k.shape[2] == T, "pass an explicit attn_mask when prefilling a non-empty cache" | |
| y = F.scaled_dot_product_attention(q, k, v, dropout_p=dp, is_causal=(T > 1)) | |
| else: | |
| y = F.scaled_dot_product_attention(q, k, v, attn_mask=attn_mask, dropout_p=dp) | |
| y = y.transpose(1, 2).contiguous().view(B, T, C) | |
| return self.resid_dropout(self.c_proj(y)) | |
| class MLP(nn.Module): | |
| """Standard 2-layer GELU MLP: d -> hidden -> d.""" | |
| def __init__(self, n_embd, hidden, bias=False, dropout=0.0): | |
| super().__init__() | |
| self.c_fc = nn.Linear(n_embd, hidden, bias=bias) | |
| self.gelu = nn.GELU() | |
| self.c_proj = nn.Linear(hidden, n_embd, bias=bias) | |
| self.dropout = nn.Dropout(dropout) | |
| def forward(self, x): | |
| return self.dropout(self.c_proj(self.gelu(self.c_fc(x)))) | |
| class MoE(nn.Module): | |
| """Token-choice top-k Mixture-of-Experts, drop-in replacement for MLP. | |
| * router: linear d -> n_experts, softmax over routed experts | |
| * top-1: output = p_e * E_e(x) (Switch Transformer; keeps router gradient) | |
| * top-k (k>1): gates renormalised over the chosen k (Mixtral-style) | |
| * shared experts (DeepSeek-MoE style) are applied to every token and added | |
| * aux loss = coef * E * sum_e f_e * P_e (+ router z-loss), computed on real tokens only | |
| * optional tracking of which experts are chosen, per token group (language / segment) | |
| """ | |
| def __init__(self, cfg: ModelConfig): | |
| super().__init__() | |
| h = cfg.expert_hidden or (4 * cfg.n_embd // max(cfg.n_experts + cfg.n_shared, 1)) | |
| self.n_experts, self.top_k, self.n_shared = cfg.n_experts, cfg.top_k, cfg.n_shared | |
| self.aux_coef, self.z_coef = cfg.aux_loss_coef, cfg.z_loss_coef | |
| self.experts = nn.ModuleList([MLP(cfg.n_embd, h, cfg.bias, cfg.dropout) for _ in range(cfg.n_experts)]) | |
| self.shared = nn.ModuleList([MLP(cfg.n_embd, h, cfg.bias, cfg.dropout) for _ in range(cfg.n_shared)]) | |
| self.router = nn.Linear(cfg.n_embd, cfg.n_experts, bias=False) | |
| self.aux_loss = torch.zeros(()) | |
| # tracking state | |
| self.track = False | |
| self.n_groups = 0 | |
| self.counts = None # (n_groups, n_experts) long | |
| self.token_groups = None # (N,) long, -1 = ignore | |
| self.load_accum = None # (n_experts,) running fraction for training logs | |
| self.load_steps = 0 | |
| def reset_tracking(self, n_groups): | |
| self.n_groups = n_groups | |
| self.counts = torch.zeros(n_groups, self.n_experts, dtype=torch.long, device=self.router.weight.device) | |
| def forward(self, x, tok_mask=None): | |
| B, T, C = x.shape | |
| xf = x.reshape(-1, C) | |
| N = xf.shape[0] | |
| if tok_mask is not None: | |
| sel = tok_mask.reshape(-1).nonzero(as_tuple=True)[0] | |
| xs = xf.index_select(0, sel) | |
| else: | |
| sel, xs = None, xf | |
| logits = self.router(xs).float() # (n, E) | |
| probs = F.softmax(logits, dim=-1) | |
| topv, topi = probs.topk(self.top_k, dim=-1) # (n, k) | |
| if self.top_k > 1: | |
| topv = topv / topv.sum(-1, keepdim=True) | |
| out = torch.zeros(xs.shape, dtype=torch.float32, device=xs.device) | |
| for e, expert in enumerate(self.experts): | |
| tok, slot = (topi == e).nonzero(as_tuple=True) | |
| if tok.numel() == 0: | |
| continue | |
| y = expert(xs.index_select(0, tok)).float() * topv[tok, slot].unsqueeze(-1) | |
| out.index_add_(0, tok, y) | |
| for s in self.shared: | |
| out = out + s(xs).float() | |
| if self.training and xs.shape[0] > 0: | |
| E = self.n_experts | |
| f = F.one_hot(topi, E).sum(1).float().mean(0) / self.top_k # fraction of assignments | |
| P = probs.mean(0) # mean router prob | |
| aux = E * (f * P).sum() | |
| z = torch.logsumexp(logits, dim=-1).pow(2).mean() | |
| self.aux_loss = self.aux_coef * aux + self.z_coef * z | |
| with torch.no_grad(): | |
| self.load_accum = f.detach() if self.load_accum is None else self.load_accum + f.detach() | |
| self.load_steps += 1 | |
| else: | |
| self.aux_loss = torch.zeros((), device=x.device) | |
| if self.track and self.token_groups is not None: | |
| with torch.no_grad(): | |
| g = self.token_groups.reshape(-1) | |
| if sel is not None: | |
| g = g.index_select(0, sel) | |
| g = g.unsqueeze(1).expand_as(topi) | |
| ok = g >= 0 | |
| flat = (g[ok] * self.n_experts + topi[ok]).reshape(-1) | |
| self.counts += torch.bincount(flat, minlength=self.n_groups * self.n_experts).view( | |
| self.n_groups, self.n_experts) | |
| out = out.to(x.dtype) | |
| if sel is not None: | |
| full = torch.zeros_like(xf) | |
| full.index_copy_(0, sel, out) | |
| out = full | |
| return out.view(B, T, C) | |
| class Block(nn.Module): | |
| def __init__(self, cfg: ModelConfig): | |
| super().__init__() | |
| self.ln_1 = LayerNorm(cfg.n_embd, cfg.bias) | |
| self.attn = CausalSelfAttention(cfg) | |
| self.ln_2 = LayerNorm(cfg.n_embd, cfg.bias) | |
| if cfg.ffn_type == "moe": | |
| self.ffn = MoE(cfg) | |
| else: | |
| self.ffn = MLP(cfg.n_embd, cfg.mlp_hidden or 4 * cfg.n_embd, cfg.bias, cfg.dropout) | |
| self.is_moe = cfg.ffn_type == "moe" | |
| def forward(self, x, attn_mask=None, cache=None, layer_idx=0, tok_mask=None): | |
| x = x + self.attn(self.ln_1(x), attn_mask, cache, layer_idx) | |
| h = self.ln_2(x) | |
| x = x + (self.ffn(h, tok_mask) if self.is_moe else self.ffn(h)) | |
| return x | |
| class Transformer(nn.Module): | |
| def __init__(self, cfg: ModelConfig): | |
| super().__init__() | |
| self.config = cfg | |
| self.transformer = nn.ModuleDict(dict( | |
| wte=nn.Embedding(cfg.vocab_size, cfg.n_embd), | |
| wpe=nn.Embedding(cfg.block_size, cfg.n_embd), | |
| drop=nn.Dropout(cfg.dropout), | |
| h=nn.ModuleList([Block(cfg) for _ in range(cfg.n_layer)]), | |
| ln_f=LayerNorm(cfg.n_embd, cfg.bias), | |
| )) | |
| self.lm_head = nn.Linear(cfg.n_embd, cfg.vocab_size, bias=False) | |
| self.transformer.wte.weight = self.lm_head.weight # weight tying | |
| self.apply(self._init_weights) | |
| for pn, p in self.named_parameters(): | |
| if pn.endswith("c_proj.weight"): | |
| torch.nn.init.normal_(p, mean=0.0, std=0.02 / math.sqrt(2 * cfg.n_layer)) | |
| def _init_weights(self, m): | |
| if isinstance(m, nn.Linear): | |
| torch.nn.init.normal_(m.weight, mean=0.0, std=0.02) | |
| if m.bias is not None: | |
| torch.nn.init.zeros_(m.bias) | |
| elif isinstance(m, nn.Embedding): | |
| torch.nn.init.normal_(m.weight, mean=0.0, std=0.02) | |
| # ---- MoE helpers ---- | |
| def moe_layers(self): | |
| return [b.ffn for b in self.transformer.h if isinstance(b.ffn, MoE)] | |
| # ---- forward ---- | |
| def hidden_states(self, idx, pos=None, attn_mask=None, cache=None, tok_mask=None): | |
| B, T = idx.shape | |
| if pos is None: | |
| start = cache.length if cache is not None else 0 | |
| assert start + T <= self.config.block_size, f"sequence length {start+T} > block_size" | |
| pos = torch.arange(start, start + T, device=idx.device).unsqueeze(0) | |
| x = self.transformer.drop(self.transformer.wte(idx) + self.transformer.wpe(pos)) | |
| for i, block in enumerate(self.transformer.h): | |
| x = block(x, attn_mask, cache, i, tok_mask) | |
| return self.transformer.ln_f(x) | |
| def aux_loss(self): | |
| layers = self.moe_layers() | |
| if not layers: | |
| return None | |
| return sum(m.aux_loss for m in layers) | |
| def forward(self, idx, targets=None, pos=None, attn_mask=None, cache=None, | |
| tok_mask=None, last_only=False): | |
| """targets: (B,T) with -1 = ignore. Returns (logits, loss, aux_loss). | |
| When targets are given, logits are only computed on supervised positions | |
| (saves the big vocab projection on source / padding tokens) and None is returned.""" | |
| x = self.hidden_states(idx, pos, attn_mask, cache, tok_mask) | |
| aux = self.aux_loss() | |
| if targets is not None: | |
| valid = targets != -1 | |
| logits = self.lm_head(x[valid]) | |
| loss = F.cross_entropy(logits.float(), targets[valid]) | |
| return None, loss, aux | |
| if last_only: | |
| x = x[:, -1:, :] | |
| return self.lm_head(x), None, aux | |
| def token_nll(self, idx, targets, tok_mask=None): | |
| """Per-token NLL (B,T), 0 where targets == -1.""" | |
| x = self.hidden_states(idx, tok_mask=tok_mask) | |
| valid = targets != -1 | |
| out = torch.zeros(targets.shape, dtype=torch.float32, device=idx.device) | |
| logits = self.lm_head(x[valid]).float() | |
| out[valid] = F.cross_entropy(logits, targets[valid], reduction="none") | |
| return out | |
| # ---- parameter accounting ---- | |
| def _numel(m): | |
| return sum(p.numel() for p in m.parameters()) | |
| def param_report(model: Transformer): | |
| cfg = model.config | |
| seen, total = set(), 0 | |
| for p in model.parameters(): | |
| if id(p) not in seen: | |
| seen.add(id(p)); total += p.numel() | |
| emb = model.transformer.wte.weight.numel() + model.transformer.wpe.weight.numel() | |
| ffn_total = ffn_active = router = 0 | |
| for b in model.transformer.h: | |
| f = b.ffn | |
| if isinstance(f, MoE): | |
| per_exp = _numel(f.experts[0]) | |
| r = _numel(f.router) | |
| sh = sum(_numel(s) for s in f.shared) | |
| ffn_total += per_exp * f.n_experts + sh + r | |
| ffn_active += per_exp * f.top_k + sh + r | |
| router += r | |
| else: | |
| ffn_total += _numel(f); ffn_active += _numel(f) | |
| non_ffn = total - ffn_total | |
| return { | |
| "total_params": total, | |
| "active_params_per_token": non_ffn + ffn_active, | |
| "embedding_params": emb, | |
| "total_non_embedding": total - emb, | |
| "active_non_embedding": non_ffn + ffn_active - emb, | |
| "ffn_total_params": ffn_total, | |
| "ffn_active_params": ffn_active, | |
| "router_params": router, | |
| } | |
| def print_param_report(rep, title=""): | |
| print(f"--- parameters {title} ---") | |
| for k, v in rep.items(): | |
| print(f" {k:26s} {v:>12,d} ({fmt_num(v)})") | |
| # ---- save / load (safetensors handles the tied embedding) ---- | |
| def save_model_dir(model, out_dir, extra_config=None): | |
| os.makedirs(out_dir, exist_ok=True) | |
| try: | |
| from safetensors.torch import save_model | |
| save_model(model, os.path.join(out_dir, "model.safetensors")) | |
| except Exception as e: | |
| print("[save] safetensors failed, falling back to torch.save:", e) | |
| torch.save(model.state_dict(), os.path.join(out_dir, "pytorch_model.bin")) | |
| cfg = model.config.to_dict() | |
| if extra_config: | |
| cfg.update(extra_config) | |
| save_json(cfg, os.path.join(out_dir, "config.json")) | |
| def load_model_dir(path_or_repo, token=None, device="cpu"): | |
| """Load a model saved by save_model_dir from a local dir or a HF repo id.""" | |
| if os.path.isdir(path_or_repo): | |
| d = path_or_repo | |
| cfg_path = os.path.join(d, "config.json") | |
| st = os.path.join(d, "model.safetensors") | |
| else: | |
| cfg_path = hf_try_download(path_or_repo, "config.json", token) | |
| st = hf_try_download(path_or_repo, "model.safetensors", token) | |
| cfg = ModelConfig.from_dict(load_json(cfg_path)) | |
| model = Transformer(cfg) | |
| from safetensors.torch import load_model | |
| load_model(model, st, device="cpu") | |
| return model.to(device) | |
| # ----------------------------------------------------------------------------- | |
| # Generation (KV cache, left padding, batched). Uses model.forward only. | |
| # ----------------------------------------------------------------------------- | |
| def generate(model, prompts: List[List[int]], max_new_tokens: int, eos_id: Optional[int], | |
| pad_id: int = 0, temperature: float = 0.0, top_k: Optional[int] = None, | |
| device=None, amp=None, generator=None): | |
| """Batched autoregressive generation. temperature=0 -> greedy. | |
| Returns list of generated token lists (EOS stripped).""" | |
| device = device or next(model.parameters()).device | |
| model.eval() | |
| B = len(prompts) | |
| L = max(len(p) for p in prompts) | |
| max_new_tokens = min(max_new_tokens, model.config.block_size - L) | |
| idx = torch.full((B, L), pad_id, dtype=torch.long) | |
| valid = torch.zeros((B, L), dtype=torch.bool) | |
| for i, p in enumerate(prompts): | |
| idx[i, L - len(p):] = torch.tensor(p, dtype=torch.long) | |
| valid[i, L - len(p):] = True | |
| idx, valid = idx.to(device), valid.to(device) | |
| pos = (valid.long().cumsum(1) - 1).clamp(min=0) | |
| causal = torch.tril(torch.ones(L, L, dtype=torch.bool, device=device)) | |
| eye = torch.eye(L, dtype=torch.bool, device=device) | |
| # key must be valid (not left padding); always allow self-attention so pad rows never go NaN | |
| mask = (causal[None, None] & valid[:, None, None, :]) | eye[None, None] | |
| cache = KVCache(model.config.n_layer) | |
| with autocast_ctx(device, amp): | |
| logits, _, _ = model(idx, pos=pos, attn_mask=mask, cache=cache, last_only=True) | |
| next_pos = pos[:, -1] + 1 | |
| finished = torch.zeros(B, dtype=torch.bool, device=device) | |
| outs = [] | |
| for _ in range(max_new_tokens): | |
| lg = logits[:, -1, :].float() | |
| if temperature and temperature > 0: | |
| lg = lg / temperature | |
| if top_k: | |
| v, _ = torch.topk(lg, min(top_k, lg.size(-1))) | |
| lg[lg < v[:, [-1]]] = -float("inf") | |
| nxt = torch.multinomial(F.softmax(lg, -1), 1, generator=generator).squeeze(1) | |
| else: | |
| nxt = lg.argmax(-1) | |
| nxt = torch.where(finished, torch.full_like(nxt, pad_id), nxt) | |
| outs.append(nxt) | |
| if eos_id is not None: | |
| finished |= nxt == eos_id | |
| if bool(finished.all()): | |
| break | |
| valid = torch.cat([valid, torch.ones(B, 1, dtype=torch.bool, device=device)], 1) | |
| with autocast_ctx(device, amp): | |
| logits, _, _ = model(nxt[:, None], pos=next_pos[:, None], attn_mask=valid[:, None, None, :], | |
| cache=cache, last_only=True) | |
| next_pos = next_pos + 1 | |
| if not outs: | |
| return [[] for _ in range(B)] | |
| gen = torch.stack(outs, 1).tolist() | |
| res = [] | |
| for row in gen: | |
| if eos_id is not None and eos_id in row: | |
| row = row[:row.index(eos_id)] | |
| res.append(row) | |
| return res | |
| def generate_sorted(model, prompts, max_new_tokens, eos_id, pad_id=0, batch_size=64, | |
| device=None, amp=None, progress=False, **kw): | |
| """Sort prompts by length for efficient batching, restore the original order.""" | |
| order = sorted(range(len(prompts)), key=lambda i: len(prompts[i])) | |
| out = [None] * len(prompts) | |
| t0 = time.time() | |
| for bi, s in enumerate(range(0, len(order), batch_size)): | |
| ids = order[s:s + batch_size] | |
| res = generate(model, [prompts[i] for i in ids], max_new_tokens, eos_id, pad_id, | |
| device=device, amp=amp, **kw) | |
| for i, r in zip(ids, res): | |
| out[i] = r | |
| if progress and bi % 20 == 0: | |
| print(f" generated {min(s+batch_size, len(order))}/{len(order)} ({time.time()-t0:.0f}s)") | |
| return out | |
| # ----------------------------------------------------------------------------- | |
| # HAP-E split (Part 2) - shared by the tokenizer notebook and the LM runs | |
| # ----------------------------------------------------------------------------- | |
| def hape_split(doc_ids, seed=0, val_frac=0.05, test_frac=0.05): | |
| """Deterministic split by base document id. Returns dict base -> 'train'|'val'|'test'.""" | |
| bases = sorted({d.split("@")[0] for d in doc_ids}) | |
| rng = np.random.default_rng(seed) | |
| perm = rng.permutation(len(bases)) | |
| n_val, n_test = int(len(bases) * val_frac), int(len(bases) * test_frac) | |
| split = {} | |
| for rank, i in enumerate(perm): | |
| split[bases[i]] = "val" if rank < n_val else ("test" if rank < n_val + n_test else "train") | |
| return split | |
| # ----------------------------------------------------------------------------- | |
| # LR schedule | |
| # ----------------------------------------------------------------------------- | |
| def lr_factor(step, total_steps, warmup_steps, min_ratio=0.1): | |
| """Linear warmup then cosine decay to min_ratio.""" | |
| if step < warmup_steps: | |
| return (step + 1) / max(1, warmup_steps) | |
| prog = (step - warmup_steps) / max(1, total_steps - warmup_steps) | |
| prog = min(max(prog, 0.0), 1.0) | |
| return min_ratio + (1 - min_ratio) * 0.5 * (1 + math.cos(math.pi * prog)) | |
| def set_lr(optimizer, factor): | |
| for g in optimizer.param_groups: | |
| if "base_lr" not in g: | |
| g["base_lr"] = g["lr"] | |
| g["lr"] = g["base_lr"] * factor | |
| # ----------------------------------------------------------------------------- | |
| # run several jobs at once, one per GPU (e.g. Kaggle "GPU T4 x2") | |
| # ----------------------------------------------------------------------------- | |
| def launch_parallel(jobs, poll_s=180, tail=3): | |
| """jobs: list of (name, python_code). Job i runs in its own process on GPU i. | |
| Output goes to <name>.log; the last lines of every log are printed every poll_s seconds.""" | |
| import subprocess | |
| get_hf_token() # puts HF_TOKEN into os.environ for the children | |
| n_gpu = max(torch.cuda.device_count(), 1) | |
| assert len(jobs) <= n_gpu, f"{len(jobs)} jobs but only {n_gpu} GPU(s)" | |
| procs = [] | |
| for gpu, (name, code) in enumerate(jobs): | |
| env = dict(os.environ, CUDA_VISIBLE_DEVICES=str(gpu), PYTHONUNBUFFERED="1", | |
| PYTORCH_CUDA_ALLOC_CONF="expandable_segments:True") | |
| logf = open(f"{name}.log", "w") | |
| procs.append((name, subprocess.Popen([sys.executable, "-u", "-c", code], env=env, | |
| stdout=logf, stderr=subprocess.STDOUT), logf)) | |
| print(f"started {name} on GPU {gpu} (log: {name}.log)") | |
| while True: | |
| time.sleep(poll_s) | |
| alive = False | |
| for name, p, _ in procs: | |
| alive |= p.poll() is None | |
| try: | |
| lines = open(f"{name}.log").read().splitlines()[-tail:] | |
| except Exception: | |
| lines = [] | |
| print(f"--- {name} ({'running' if p.poll() is None else 'exit ' + str(p.returncode)}) ---") | |
| for l in lines: | |
| print(" ", l[:200]) | |
| if not alive: | |
| break | |
| for _, _, f in procs: | |
| f.close() | |
| codes = {name: p.returncode for name, p, _ in procs} | |
| print("exit codes:", codes) | |
| return codes | |
| # ----------------------------------------------------------------------------- | |
| # Plot style (validated categorical palette, fixed order) | |
| # ----------------------------------------------------------------------------- | |
| PALETTE = ["#2a78d6", "#eb6834", "#1baf7a", "#eda100", "#e87ba4", "#008300", "#4a3aa7", "#e34948"] | |
| def setup_plot_style(): | |
| import matplotlib as mpl | |
| mpl.rcParams.update({ | |
| "figure.dpi": 110, "savefig.dpi": 150, "savefig.bbox": "tight", | |
| "axes.spines.top": False, "axes.spines.right": False, | |
| "axes.grid": True, "grid.color": "#e4e3df", "grid.linewidth": 0.8, | |
| "axes.edgecolor": "#8a8984", "axes.labelcolor": "#2b2a27", | |
| "xtick.color": "#52514e", "ytick.color": "#52514e", | |
| "axes.prop_cycle": mpl.cycler(color=PALETTE), | |
| "lines.linewidth": 2.0, "legend.frameon": False, "font.size": 10, | |
| }) | |