Download model.py from etanlightstone/ddm-medium-injection: direct link, hf CLI and curl.
- Browser
- Download file 24.5 kB
-
https://huggingface.co/etanlightstone/ddm-medium-injection/resolve/main/model.py
- Command line
-
hf download hf://etanlightstone/ddm-medium-injection/model.py
-
curl -L -o model.py https://huggingface.co/etanlightstone/ddm-medium-injection/resolve/main/model.py
24.5 kB
| """ | |
| model.py - the Domino Decision Model (DDM), a "System One"-style decision model | |
| reconstructed from public evidence about | |
| TypeSafe's Jev (mainly Archer Hume's black-box study, "Jev's Architecture Unmasked"). | |
| This is the most likely architecture given that evidence, NOT TypeSafe's actual code. | |
| [ state ][ Q1 + options + <dec> ][ Q2 + options + <dec> ] ... <- ONE packed sequence | |
| * Causal transformer backbone (RMSNorm, RoPE, GQA, SwiGLU; optional sparse MoE). | |
| * Block attention mask: the state attends causally to itself; each question branch | |
| attends to the state + itself, never to sibling branches -> questions are isolated | |
| but the state is encoded only once per forward pass. | |
| * Each branch restarts its position ids right after the state, so a question "sees" | |
| exactly what it would see if it were sent alone. | |
| * Options inside a branch are read together as a list (listwise), so options can | |
| influence each other ("none of the above" works, fake options can't be forged: | |
| option boundaries are special tokens that text can never produce). | |
| * Pointer readout: the hidden state at <dec> is compared with the hidden state at | |
| each option's <opt_end> -> one logit per option. The output size follows the | |
| number of options supplied (activation shape, not a fixed weight shape). | |
| * Noul (yes/no) uses a scalar head on <dec> -> sigmoid. | |
| * No text generation loop: one forward pass -> all probabilities. | |
| * Per-question-type temperature buffers, fitted after training (post-hoc calibration), | |
| are stored inside the checkpoint. | |
| Note: the dense boolean mask used here is O(L^2) memory. That is fine for the tiny/ | |
| small/medium presets; for 32k/64k contexts switch to FlexAttention (torch>=2.5) with a | |
| block mask, or a prefix-KV-cache serving setup. | |
| """ | |
| from __future__ import annotations | |
| import json | |
| import math | |
| from dataclasses import asdict, dataclass | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| # -------------------------------------------------------------------------------------- | |
| # Question types & special tokens | |
| # -------------------------------------------------------------------------------------- | |
| QTYPES = {"choice": 0, "score": 1, "noul": 2} | |
| QTYPE_NAMES = {v: k for k, v in QTYPES.items()} | |
| # Special tokens live OUTSIDE the text vocabulary, so no user text can ever produce them. | |
| SPECIAL_TOKENS = ["<pad>", "<bos>", "<state>", "<q>", "<choice>", "<score>", "<noul>", | |
| "<opt>", "<opt_end>", "<dec>"] | |
| NEG_INF = -1e9 # used for masking option logits (logits are computed in float32) | |
| class RequestError(ValueError): | |
| """Raised for invalid requests (the equivalent of an HTTP 422).""" | |
| # -------------------------------------------------------------------------------------- | |
| # Tokenizers | |
| # -------------------------------------------------------------------------------------- | |
| class ByteTokenizer: | |
| """UTF-8 bytes + special tokens. No dependencies; great for tests and small models.""" | |
| def __init__(self): | |
| self.name = "byte" | |
| self.n_special = len(SPECIAL_TOKENS) | |
| self.special = {t: i for i, t in enumerate(SPECIAL_TOKENS)} | |
| self.vocab_size = self.n_special + 256 | |
| self.pad_id = self.special["<pad>"] | |
| def encode(self, text: str) -> list[int]: | |
| return [b + self.n_special for b in text.encode("utf-8")] | |
| class TiktokenTokenizer: | |
| """BPE via tiktoken (e.g. o200k_base, the vocabulary Jev's token counts most resemble).""" | |
| def __init__(self, encoding: str = "o200k_base"): | |
| import tiktoken # pip install tiktoken | |
| self.enc = tiktoken.get_encoding(encoding) | |
| self.name = f"tiktoken:{encoding}" | |
| n = self.enc.n_vocab | |
| self.special = {t: n + i for i, t in enumerate(SPECIAL_TOKENS)} | |
| self.vocab_size = n + len(SPECIAL_TOKENS) | |
| self.pad_id = self.special["<pad>"] | |
| def encode(self, text: str) -> list[int]: | |
| return self.enc.encode_ordinary(text) # never emits special tokens | |
| def get_tokenizer(name: str): | |
| if name == "byte": | |
| return ByteTokenizer() | |
| if name.startswith("tiktoken:"): | |
| return TiktokenTokenizer(name.split(":", 1)[1]) | |
| raise ValueError(f"Unknown tokenizer '{name}' (use 'byte' or 'tiktoken:<encoding>')") | |
| # -------------------------------------------------------------------------------------- | |
| # Config | |
| # -------------------------------------------------------------------------------------- | |
| class ModelConfig: | |
| tokenizer: str = "byte" # "byte" or "tiktoken:o200k_base" | |
| vocab_size: int = 0 # 0 -> derived from tokenizer | |
| d_model: int = 256 # hidden size | |
| n_layers: int = 4 # transformer blocks | |
| n_heads: int = 4 # query heads | |
| n_kv_heads: int = 0 # 0 -> same as n_heads (no GQA) | |
| d_ff: int = 0 # 0 -> ~8/3 * d_model (SwiGLU), rounded to 64 | |
| n_experts: int = 0 # 0 -> dense FFN; >0 -> sparse MoE FFN in every block | |
| experts_top_k: int = 2 # experts active per token (MoE only) | |
| dropout: float = 0.0 | |
| rope_theta: float = 10000.0 | |
| readout_dim: int = 0 # 0 -> d_model; size of the pointer comparison space | |
| max_branch_tokens: int = 2048 # state + longest single question (Jev: 32k) | |
| max_total_tokens: int = 4096 # state + all questions (Jev: 64k) | |
| max_options: int = 255 # per Choice (Jev: 255) | |
| min_score_levels: int = 2 # Jev: 2 | |
| max_score_levels: int = 10 # Jev: 10 | |
| def __post_init__(self): | |
| if self.vocab_size == 0: | |
| self.vocab_size = get_tokenizer(self.tokenizer).vocab_size | |
| if self.n_kv_heads == 0: | |
| self.n_kv_heads = self.n_heads | |
| if self.d_ff == 0: | |
| self.d_ff = max(64, round(self.d_model * 8 / 3 / 64) * 64) | |
| if self.readout_dim == 0: | |
| self.readout_dim = self.d_model | |
| assert self.d_model % self.n_heads == 0, "d_model must be divisible by n_heads" | |
| assert self.n_heads % self.n_kv_heads == 0, "n_heads must be divisible by n_kv_heads" | |
| assert (self.d_model // self.n_heads) % 2 == 0, "head_dim must be even (RoPE)" | |
| # Size presets. Pick one in train.py with --preset, or override any field on the CLI. | |
| # | |
| # ModelConfig(**PRESETS["tiny"]) ~0.4M params. CPU smoke tests; trains in minutes. | |
| # ModelConfig(**PRESETS["small"]) ~10M params. One consumer GPU; narrow domains. | |
| # ModelConfig(**PRESETS["medium"]) ~230M params (≈155M are the 200k-vocab embedding). | |
| # A single 24-80GB GPU; multi-domain experiments. | |
| # ModelConfig(**PRESETS["large"]) ~13.6B total / ~2.3B active params, 16-expert MoE. | |
| # Multi-GPU. Realistically this size only becomes | |
| # useful when initialised from a pretrained LLM | |
| # (see DATASET.md) - from scratch it has no world | |
| # knowledge beyond your dataset. | |
| # | |
| # Custom example: | |
| # ModelConfig(tokenizer="byte", d_model=512, n_layers=8, n_heads=8, n_kv_heads=2, | |
| # n_experts=8, experts_top_k=2, max_branch_tokens=4096, max_total_tokens=8192) | |
| PRESETS: dict[str, dict] = { | |
| "tiny": dict(tokenizer="byte", d_model=128, n_layers=2, n_heads=4, | |
| max_branch_tokens=1024, max_total_tokens=2048), | |
| "small": dict(tokenizer="byte", d_model=384, n_layers=6, n_heads=6, n_kv_heads=2, | |
| max_branch_tokens=2048, max_total_tokens=4096), | |
| "medium": dict(tokenizer="tiktoken:o200k_base", d_model=768, n_layers=12, n_heads=12, | |
| n_kv_heads=4, max_branch_tokens=8192, max_total_tokens=16384), | |
| "large": dict(tokenizer="tiktoken:o200k_base", d_model=2048, n_layers=24, n_heads=16, | |
| n_kv_heads=8, n_experts=16, experts_top_k=2, | |
| max_branch_tokens=32768, max_total_tokens=65536), | |
| } | |
| # -------------------------------------------------------------------------------------- | |
| # Building blocks | |
| # -------------------------------------------------------------------------------------- | |
| class RMSNorm(nn.Module): | |
| def __init__(self, d: int, eps: float = 1e-6): | |
| super().__init__() | |
| self.weight = nn.Parameter(torch.ones(d)) | |
| self.eps = eps | |
| def forward(self, x): | |
| xf = x.float() | |
| out = xf * torch.rsqrt(xf.pow(2).mean(-1, keepdim=True) + self.eps) | |
| return out.type_as(x) * self.weight | |
| def rope_cos_sin(pos: torch.Tensor, head_dim: int, theta: float): | |
| """pos: [B, L] explicit position ids (branches restart after the state).""" | |
| inv = 1.0 / (theta ** (torch.arange(0, head_dim, 2, device=pos.device).float() / head_dim)) | |
| freqs = pos.float()[..., None] * inv # [B, L, hd/2] | |
| emb = torch.cat([freqs, freqs], dim=-1) # [B, L, hd] | |
| return emb.cos()[:, None], emb.sin()[:, None] # [B, 1, L, hd] | |
| def apply_rope(x, cos, sin): | |
| x1, x2 = x.float().chunk(2, dim=-1) | |
| rotated = torch.cat([-x2, x1], dim=-1) | |
| return (x.float() * cos + rotated * sin).type_as(x) | |
| class Attention(nn.Module): | |
| def __init__(self, cfg: ModelConfig): | |
| super().__init__() | |
| self.nh, self.nkv = cfg.n_heads, cfg.n_kv_heads | |
| self.hd = cfg.d_model // cfg.n_heads | |
| self.q = nn.Linear(cfg.d_model, self.nh * self.hd, bias=False) | |
| self.k = nn.Linear(cfg.d_model, self.nkv * self.hd, bias=False) | |
| self.v = nn.Linear(cfg.d_model, self.nkv * self.hd, bias=False) | |
| self.o = nn.Linear(self.nh * self.hd, cfg.d_model, bias=False) | |
| self.dropout = cfg.dropout | |
| def forward(self, x, cos, sin, mask): | |
| B, L, _ = x.shape | |
| q = self.q(x).view(B, L, self.nh, self.hd).transpose(1, 2) | |
| k = self.k(x).view(B, L, self.nkv, self.hd).transpose(1, 2) | |
| v = self.v(x).view(B, L, self.nkv, self.hd).transpose(1, 2) | |
| q, k = apply_rope(q, cos, sin), apply_rope(k, cos, sin) | |
| if self.nkv != self.nh: # grouped-query attention | |
| rep = self.nh // self.nkv | |
| k, v = k.repeat_interleave(rep, dim=1), v.repeat_interleave(rep, dim=1) | |
| out = F.scaled_dot_product_attention( | |
| q, k, v, attn_mask=mask, dropout_p=self.dropout if self.training else 0.0) | |
| return self.o(out.transpose(1, 2).reshape(B, L, self.nh * self.hd)) | |
| class SwiGLU(nn.Module): | |
| def __init__(self, d: int, d_ff: int): | |
| super().__init__() | |
| self.w1 = nn.Linear(d, d_ff, bias=False) | |
| self.w3 = nn.Linear(d, d_ff, bias=False) | |
| self.w2 = nn.Linear(d_ff, d, bias=False) | |
| def forward(self, x): | |
| return self.w2(F.silu(self.w1(x)) * self.w3(x)) | |
| class MoE(nn.Module): | |
| """Top-k routed sparse experts with a Switch-style load-balancing loss. | |
| Simple loop-over-experts dispatch: clear, not fast. Use a grouped-GEMM kernel at scale.""" | |
| def __init__(self, cfg: ModelConfig): | |
| super().__init__() | |
| self.n_experts, self.k = cfg.n_experts, cfg.experts_top_k | |
| self.router = nn.Linear(cfg.d_model, cfg.n_experts, bias=False) | |
| self.experts = nn.ModuleList(SwiGLU(cfg.d_model, cfg.d_ff) for _ in range(cfg.n_experts)) | |
| self.aux_loss = torch.tensor(0.0) | |
| def forward(self, x): | |
| B, L, D = x.shape | |
| xf = x.reshape(-1, D) | |
| probs = self.router(xf).float().softmax(dim=-1) # [N, E] | |
| top_p, top_i = probs.topk(self.k, dim=-1) # [N, k] | |
| top_p = top_p / top_p.sum(dim=-1, keepdim=True) | |
| out = torch.zeros_like(xf) | |
| for e, expert in enumerate(self.experts): | |
| tok, slot = (top_i == e).nonzero(as_tuple=True) | |
| if tok.numel() == 0: | |
| continue | |
| y = expert(xf[tok]) * top_p[tok, slot, None] | |
| out.index_add_(0, tok, y.to(out.dtype)) | |
| routed_frac = F.one_hot(top_i, self.n_experts).float().sum(1).mean(0) / self.k | |
| self.aux_loss = self.n_experts * (routed_frac * probs.mean(0)).sum() | |
| return out.view(B, L, D) | |
| class Block(nn.Module): | |
| def __init__(self, cfg: ModelConfig): | |
| super().__init__() | |
| self.norm1, self.norm2 = RMSNorm(cfg.d_model), RMSNorm(cfg.d_model) | |
| self.attn = Attention(cfg) | |
| self.ffn = MoE(cfg) if cfg.n_experts > 0 else SwiGLU(cfg.d_model, cfg.d_ff) | |
| self.drop = nn.Dropout(cfg.dropout) | |
| def forward(self, x, cos, sin, mask): | |
| x = x + self.drop(self.attn(self.norm1(x), cos, sin, mask)) | |
| return x + self.drop(self.ffn(self.norm2(x))) | |
| # -------------------------------------------------------------------------------------- | |
| # The model | |
| # -------------------------------------------------------------------------------------- | |
| class DominoDecisionModel(nn.Module): | |
| def __init__(self, cfg: ModelConfig): | |
| super().__init__() | |
| self.config = cfg | |
| self.embed = nn.Embedding(cfg.vocab_size, cfg.d_model) | |
| self.blocks = nn.ModuleList(Block(cfg) for _ in range(cfg.n_layers)) | |
| self.norm = RMSNorm(cfg.d_model) | |
| # Pointer readout (Choice + Score) and scalar readout (Noul). | |
| self.dec_proj = nn.Linear(cfg.d_model, cfg.readout_dim, bias=False) | |
| self.opt_proj = nn.Linear(cfg.d_model, cfg.readout_dim, bias=False) | |
| self.noul_head = nn.Linear(cfg.d_model, 1) | |
| # Post-hoc calibration: one temperature per question type, saved in the checkpoint. | |
| self.register_buffer("log_temperature", torch.zeros(len(QTYPES))) | |
| self.apply(self._init_weights) | |
| def _init_weights(m): | |
| if isinstance(m, nn.Linear): | |
| nn.init.normal_(m.weight, std=0.02) | |
| if m.bias is not None: | |
| nn.init.zeros_(m.bias) | |
| elif isinstance(m, nn.Embedding): | |
| nn.init.normal_(m.weight, std=0.02) | |
| def num_params(self) -> tuple[int, int]: | |
| """(total, active per token). Active counts only top-k experts.""" | |
| total = sum(p.numel() for p in self.parameters()) | |
| cfg = self.config | |
| if cfg.n_experts == 0: | |
| return total, total | |
| per_expert = 3 * cfg.d_model * cfg.d_ff | |
| inactive = cfg.n_layers * (cfg.n_experts - cfg.experts_top_k) * per_expert | |
| return total, total - inactive | |
| def build_mask(seg: torch.Tensor) -> torch.Tensor: | |
| """seg: [B, L] with 0 = state, k>=1 = question branch k, -1 = padding. | |
| Returns bool [B, 1, L, L], True = may attend.""" | |
| L = seg.shape[1] | |
| causal = torch.ones(L, L, dtype=torch.bool, device=seg.device).tril() | |
| sq, sk = seg[:, :, None], seg[:, None, :] | |
| allowed = causal[None] & ((sk == 0) | (sk == sq)) & (sk >= 0) | |
| eye = torch.eye(L, dtype=torch.bool, device=seg.device)[None] | |
| allowed = allowed | (eye & (sq < 0)) # padding rows attend to themselves (no NaNs) | |
| return allowed[:, None] | |
| def encode(self, ids, pos, seg): | |
| x = self.embed(ids) | |
| cos, sin = rope_cos_sin(pos, self.config.d_model // self.config.n_heads, | |
| self.config.rope_theta) | |
| mask = self.build_mask(seg) | |
| for block in self.blocks: | |
| x = block(x, cos, sin, mask) | |
| return self.norm(x) | |
| def readout(self, h, q_batch, q_dec, q_opt_pos, q_opt_mask): | |
| h_dec = h[q_batch, q_dec] # [Q, d] | |
| h_opt = h[q_batch[:, None], q_opt_pos] # [Q, K, d] | |
| dec = self.dec_proj(h_dec).float() | |
| opt = self.opt_proj(h_opt).float() | |
| opt_logits = torch.einsum("qr,qkr->qk", dec, opt) / math.sqrt(self.config.readout_dim) | |
| opt_logits = opt_logits.masked_fill(~q_opt_mask, NEG_INF) | |
| noul_logit = self.noul_head(h_dec).float().squeeze(-1) | |
| return opt_logits, noul_logit | |
| def forward(self, batch: dict): | |
| """Returns (opt_logits [Q, K], noul_logit [Q], hidden [B, L, d]). Uncalibrated.""" | |
| h = self.encode(batch["ids"], batch["pos"], batch["seg"]) | |
| opt_logits, noul_logit = self.readout(h, batch["q_batch"], batch["q_dec"], | |
| batch["q_opt_pos"], batch["q_opt_mask"]) | |
| return opt_logits, noul_logit, h | |
| def aux_loss(self) -> torch.Tensor: | |
| losses = [b.ffn.aux_loss for b in self.blocks if isinstance(b.ffn, MoE)] | |
| return torch.stack(losses).mean() if losses else torch.tensor(0.0) | |
| def lm_logits(self, h): | |
| """Tied next-token head, used only for the optional language-modelling aux loss.""" | |
| return h @ self.embed.weight.T | |
| def apply_temperature(self, opt_logits, noul_logit, q_type): | |
| t = self.log_temperature.exp()[q_type] | |
| return opt_logits / t[:, None], noul_logit / t | |
| def predict(self, batch: dict): | |
| """Calibrated probabilities: (option_probs [Q, K], noul_prob [Q]).""" | |
| opt_logits, noul_logit, _ = self.forward(batch) | |
| opt_logits, noul_logit = self.apply_temperature(opt_logits, noul_logit, batch["q_type"]) | |
| return opt_logits.softmax(-1), noul_logit.sigmoid() | |
| # -------------------------------------------------------------------------------------- | |
| # Request parsing & packing (shared by train.py and invoke.py) | |
| # -------------------------------------------------------------------------------------- | |
| def state_to_text(state) -> str: | |
| """State may be a string, a JSON object, or an array. Non-strings are serialised.""" | |
| if isinstance(state, str): | |
| return state | |
| if isinstance(state, (dict, list)): | |
| return json.dumps(state, ensure_ascii=False) | |
| raise RequestError("state must be a string, JSON object or array") | |
| def normalize_questions(questions: dict, cfg: ModelConfig) -> list[dict]: | |
| """Validates the Jev-style `questions` map and returns an ordered list of | |
| {id, type, instructions, keys, texts}. The question id is NOT sent to the model.""" | |
| if not isinstance(questions, dict) or not questions: | |
| raise RequestError("questions must be a non-empty object") | |
| out = [] | |
| for qid, q in questions.items(): | |
| if not isinstance(q, dict): | |
| raise RequestError(f"question '{qid}' must be an object") | |
| qtype = q.get("type") | |
| if qtype not in QTYPES: | |
| raise RequestError(f"question '{qid}': type must be one of {list(QTYPES)}") | |
| instr = q.get("instructions") | |
| if isinstance(instr, (dict, list)): | |
| instr = json.dumps(instr, ensure_ascii=False) | |
| if not isinstance(instr, str) or not instr.strip(): | |
| raise RequestError(f"question '{qid}': instructions are required") | |
| crit = q.get("criteria") | |
| if qtype == "choice": | |
| if not isinstance(crit, dict) or not crit: | |
| raise RequestError(f"question '{qid}': choice criteria must be a non-empty object") | |
| if len(crit) > cfg.max_options: | |
| raise RequestError(f"question '{qid}': at most {cfg.max_options} options") | |
| keys = [str(k) for k in crit] | |
| texts = [f"{k}: {v}" if v else str(k) for k, v in crit.items()] | |
| elif qtype == "score": | |
| if not isinstance(crit, list) or not all(isinstance(c, str) for c in crit): | |
| raise RequestError(f"question '{qid}': score criteria must be a list of strings") | |
| if not cfg.min_score_levels <= len(crit) <= cfg.max_score_levels: | |
| raise RequestError(f"question '{qid}': score needs {cfg.min_score_levels}-" | |
| f"{cfg.max_score_levels} levels") | |
| keys = [str(i) for i in range(len(crit))] | |
| texts = [f"{i}: {c}" for i, c in enumerate(crit)] | |
| else: # noul | |
| if crit is not None: | |
| raise RequestError(f"question '{qid}': noul questions take no criteria") | |
| keys, texts = [], [] | |
| out.append(dict(id=str(qid), type=qtype, instructions=instr, keys=keys, texts=texts)) | |
| return out | |
| def encode_state(tok, state_text: str) -> list[int]: | |
| return [tok.special["<bos>"], tok.special["<state>"]] + tok.encode(state_text) | |
| def encode_branch(tok, q: dict) -> tuple[list[int], list[int], int]: | |
| """Returns (token ids, <opt_end> offsets, <dec> offset) for one question branch.""" | |
| sp = tok.special | |
| ids = [sp["<q>"], sp[f"<{q['type']}>"]] + tok.encode(q["instructions"]) | |
| opt_end = [] | |
| for text in q["texts"]: | |
| ids += [sp["<opt>"]] + tok.encode(text) + [sp["<opt_end>"]] | |
| opt_end.append(len(ids) - 1) | |
| ids.append(sp["<dec>"]) | |
| return ids, opt_end, len(ids) - 1 | |
| def assemble(state_ids: list[int], branches: list, qtypes: list[str]) -> dict: | |
| """Packs one state + several branches into a single sequence.""" | |
| S = len(state_ids) | |
| ids, seg, pos = list(state_ids), [0] * S, list(range(S)) | |
| dec, opts = [], [] | |
| for j, (b_ids, opt_end, d) in enumerate(branches, start=1): | |
| off = len(ids) | |
| ids += b_ids | |
| seg += [j] * len(b_ids) | |
| pos += list(range(S, S + len(b_ids))) # every branch restarts right after the state | |
| dec.append(off + d) | |
| opts.append([off + o for o in opt_end]) | |
| return dict(ids=ids, seg=seg, pos=pos, dec=dec, opts=opts, | |
| types=[QTYPES[t] for t in qtypes]) | |
| def encode_example(tok, state_text: str, qlist: list[dict], cfg: ModelConfig, | |
| truncate_state: bool = False): | |
| """Encodes a whole example into one packed sequence. With truncate_state=True | |
| (training) the state is cut to fit the budgets; returns None if impossible.""" | |
| state_ids = encode_state(tok, state_text) | |
| branches = [encode_branch(tok, q) for q in qlist] | |
| longest = max(len(b[0]) for b in branches) | |
| all_q = sum(len(b[0]) for b in branches) | |
| budget = min(cfg.max_branch_tokens - longest, cfg.max_total_tokens - all_q) | |
| if len(state_ids) > budget: | |
| if not truncate_state: | |
| raise RequestError("request exceeds the token budget") | |
| if budget < 3: | |
| return None | |
| state_ids = state_ids[:budget] | |
| return assemble(state_ids, branches, [q["type"] for q in qlist]) | |
| def collate(encoded: list[dict], pad_id: int) -> dict: | |
| """Pads packed sequences into a batch and flattens all questions into index tensors.""" | |
| B, L = len(encoded), max(len(e["ids"]) for e in encoded) | |
| ids = torch.full((B, L), pad_id, dtype=torch.long) | |
| seg = torch.full((B, L), -1, dtype=torch.long) | |
| pos = torch.zeros((B, L), dtype=torch.long) | |
| q_batch, q_dec, q_type, q_opts = [], [], [], [] | |
| for b, e in enumerate(encoded): | |
| n = len(e["ids"]) | |
| ids[b, :n] = torch.tensor(e["ids"]) | |
| seg[b, :n] = torch.tensor(e["seg"]) | |
| pos[b, :n] = torch.tensor(e["pos"]) | |
| for d, o, t in zip(e["dec"], e["opts"], e["types"]): | |
| q_batch.append(b); q_dec.append(d); q_type.append(t); q_opts.append(o) | |
| K = max(1, max(len(o) for o in q_opts)) | |
| q_opt_pos = torch.zeros((len(q_opts), K), dtype=torch.long) | |
| q_opt_mask = torch.zeros((len(q_opts), K), dtype=torch.bool) | |
| for i, o in enumerate(q_opts): | |
| if o: | |
| q_opt_pos[i, :len(o)] = torch.tensor(o) | |
| q_opt_mask[i, :len(o)] = True | |
| return dict(ids=ids, seg=seg, pos=pos, q_batch=torch.tensor(q_batch), | |
| q_dec=torch.tensor(q_dec), q_type=torch.tensor(q_type), | |
| q_opt_pos=q_opt_pos, q_opt_mask=q_opt_mask) | |
| def batch_to(batch: dict, device) -> dict: | |
| return {k: v.to(device) for k, v in batch.items()} | |
| # -------------------------------------------------------------------------------------- | |
| # Checkpoints: config + weights + calibration + metadata in one file | |
| # -------------------------------------------------------------------------------------- | |
| CHECKPOINT_FORMAT = "ddm-v1" | |
| LEGACY_FORMATS = ("system-one-v1",) # checkpoints saved before the rename still load | |
| def save_checkpoint(path: str, model: DominoDecisionModel, extra: dict | None = None): | |
| torch.save({"format": CHECKPOINT_FORMAT, | |
| "config": asdict(model.config), | |
| "state_dict": model.state_dict(), # includes log_temperature | |
| "extra": extra or {}}, path) | |
| def load_checkpoint(path: str, map_location="cpu"): | |
| ckpt = torch.load(path, map_location=map_location, weights_only=True) | |
| if ckpt.get("format") not in (CHECKPOINT_FORMAT, *LEGACY_FORMATS): | |
| raise ValueError(f"{path} is not a {CHECKPOINT_FORMAT} checkpoint") | |
| model = DominoDecisionModel(ModelConfig(**ckpt["config"])) | |
| model.load_state_dict(ckpt["state_dict"]) | |
| return model.to(map_location), ckpt.get("extra", {}) | |