etanlightstone's picture
Upload Domino Decision Model checkpoint
1f9a4df verified
Raw History Blame Contribute Delete
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
# --------------------------------------------------------------------------------------
@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["<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", {})