""" 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 + ][ Q2 + options + ] ... <- 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 is compared with the hidden state at each option's -> 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 -> 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 = ["", "", "", "", "", "", "", "", "", ""] 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[""] 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[""] 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:')") # -------------------------------------------------------------------------------------- # Config # -------------------------------------------------------------------------------------- @dataclass 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) @staticmethod 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 @staticmethod 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 @torch.no_grad() 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[""], tok.special[""]] + tok.encode(state_text) def encode_branch(tok, q: dict) -> tuple[list[int], list[int], int]: """Returns (token ids, offsets, offset) for one question branch.""" sp = tok.special ids = [sp[""], sp[f"<{q['type']}>"]] + tok.encode(q["instructions"]) opt_end = [] for text in q["texts"]: ids += [sp[""]] + tok.encode(text) + [sp[""]] opt_end.append(len(ids) - 1) ids.append(sp[""]) 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", {})