"""Typed-decision core for spikewhale models (Laya / Jev style, System 1). Port of convaiinnovations/laya `laya/common.py` onto our causal decoder trunk. What is the same as Laya ------------------------ * Question types `choice` / `score` / `noul`, option rendering, criteria rules. * DecisionModel head: `type_emb` added to every position, a 2-layer BIDIRECTIONAL transformer head over the whole sequence (padding-masked), a scorer MLP read at one marker position per option, the `act_head`, and a per-type `temperature` buffer. All options of a question are scored together in ONE forward pass. * Temperature handling, ECE and the answer/confidence definitions. What changes for a causal decoder --------------------------------- Laya's encoder reads the state AFTER the options, and every token sees every other token. Our trunk is causal, so an option's hidden state can only see what comes before it. The sequence is therefore reordered so everything an option needs is to its left: State:\n{state} \n\n{type} question: {ins}\nOptions: \n- opt0 \n- opt1 ... The marker for option i is the LAST token of its "\n- opt_i" segment (it has read state, question, earlier options and its own text). The bidirectional head then lets every marker see every other option, which is what Laya's encoder gives for free. Batches are RIGHT-padded, so padding never affects a real token under causal attention and the trunk needs no attention mask (it keeps the FlashAttention path); only the bidirectional head uses the pad mask. """ import json import math import os import random import sys from typing import Dict, List, Optional, Union import numpy as np import torch import torch.nn as nn from torch.utils.checkpoint import checkpoint QTYPES = {"choice": 0, "score": 1, "noul": 2} QTYPE_NAMES = {v: k for k, v in QTYPES.items()} _DEFAULT_NOUL_LABELS = {"false": "false", "true": "true"} OPT_MAX_TOKENS = 48 TEMP_MIN, TEMP_MAX = 0.5, 5.0 # Laya's runtime clamp on fitted temperatures # --------------------------------------------------------------------------- # # Question rendering (verbatim Laya semantics) # --------------------------------------------------------------------------- # def serialize_state(state: Union[str, dict, list]) -> str: if isinstance(state, str): return state return json.dumps(state, ensure_ascii=False) def render_criterion(value) -> str: if isinstance(value, str): return value return json.dumps(value, ensure_ascii=False, separators=(", ", ": "), default=str) def _resolve_noul_labels(labels=None): if labels is None: labels = _DEFAULT_NOUL_LABELS if not isinstance(labels, dict) or set(labels) != {"false", "true"}: raise ValueError("noul labels must map exactly 'false' and 'true' to distinct non-empty strings") f, t = labels["false"], labels["true"] if not isinstance(f, str) or not isinstance(t, str): raise ValueError("noul labels must map exactly 'false' and 'true' to distinct non-empty strings") f, t = f.strip(), t.strip() if not f or not t or f == t: raise ValueError("noul labels must map exactly 'false' and 'true' to distinct non-empty strings") return f, t def render_options(q: Dict) -> List[str]: """Option texts in label-index order. Noul order is always [false, true].""" t, crit = q["t"], q.get("crit") if t == "choice": return [k if v is None or v == "" else "%s: %s" % (k, render_criterion(v)) for k, v in crit.items()] if t == "score": return ["level %d: %s" % (i, render_criterion(c)) for i, c in enumerate(crit)] crit = crit or {} f, tr = _resolve_noul_labels(q.get("labels")) fc, tc = crit.get("false"), crit.get("true") return [f + ": " + (render_criterion(fc) if fc not in (None, "") else "no, the statement does not hold"), tr + ": " + (render_criterion(tc) if tc not in (None, "") else "yes, the statement holds")] def check_question(qid: str, qdef) -> None: """Reject a malformed question with a message naming it (Laya's rules).""" if not isinstance(qdef, dict): raise ValueError("question %r: definition must be a dict" % (qid,)) t = qdef.get("type") if t not in QTYPES: raise ValueError("question %r: unknown type %r; use one of %s" % (qid, t, sorted(QTYPES))) if "instructions" not in qdef: raise ValueError("question %r: no 'instructions'" % (qid,)) crit = qdef.get("criteria") if t == "choice": if not isinstance(crit, (dict, list)) or not crit: raise ValueError("question %r: choice needs non-empty 'criteria' (dict label->desc or list)" % (qid,)) elif t == "score": if not isinstance(crit, list) or not crit: raise ValueError("question %r: score needs 'criteria' as a non-empty list of levels" % (qid,)) elif crit is not None: if not isinstance(crit, dict) or not {str(k).lower() for k in crit} <= {"true", "false"}: raise ValueError("question %r: noul 'criteria' must be keyed only 'true'/'false'" % (qid,)) if "labels" in qdef: if t != "noul": raise ValueError("question %r: 'labels' is only supported for noul questions" % (qid,)) _resolve_noul_labels(qdef["labels"]) def to_internal(qdef: Dict) -> Dict: t, crit = qdef["type"], qdef.get("criteria") if t == "choice" and isinstance(crit, list): crit = {c: None for c in crit} elif t == "noul" and isinstance(crit, dict): crit = {str(k).lower(): v for k, v in crit.items()} ins = qdef["instructions"] if not isinstance(ins, str): ins = json.dumps(ins, ensure_ascii=False) q = {"t": t, "ins": ins, "crit": crit} if "labels" in qdef: q["labels"] = qdef["labels"] return q # --------------------------------------------------------------------------- # # Tokenization / causal sequence layout # --------------------------------------------------------------------------- # def encode(tok, text: str) -> List[int]: ids = tok.encode(text, add_special_tokens=False) return ids.tolist() if hasattr(ids, "tolist") else list(ids) def tokenize_question(tok, q: Dict) -> Dict: """Pre-tokenize the question part once: head ids + one id list per option.""" head = encode(tok, "\n\n%s question: %s\nOptions:" % (q["t"], q["ins"])) opts = [encode(tok, "\n- " + o)[:OPT_MAX_TOKENS] for o in render_options(q)] return {"head": head, "opts": opts} def tokenize_state(tok, state) -> List[int]: return encode(tok, "State:\n" + serialize_state(state)) def assemble(bos_id: int, state_ids: List[int], qtok: Dict, max_len: int, head_max_len: int, order: Optional[List[int]] = None, truncate_left: bool = False): """ state | question | options(in `order`). Returns (ids, markers). Same budgeting as Laya's build_sequence: options + question share `head_max_len`; if options crowd it out they are cut evenly; the state gets whatever room is left of `max_len` (truncated right, or left for chat lists). markers[j] is the position of the option `order[j]`. """ opts = qtok["opts"] order = list(range(len(opts))) if order is None else order opt_ids = [opts[i] for i in order] budget = head_max_len - sum(len(o) for o in opt_ids) if budget < 16: per = max(4, (head_max_len - 16) // max(1, len(opt_ids))) opt_ids = [o[:per] for o in opt_ids] budget = head_max_len - sum(len(o) for o in opt_ids) head = qtok["head"][: max(8, budget)] tail = list(head) rel_markers = [] for o in opt_ids: tail.extend(o) rel_markers.append(len(tail) - 1) room = max(0, max_len - 1 - len(tail)) st = state_ids[max(0, len(state_ids) - room):] if truncate_left else state_ids[:room] ids = [bos_id] + st + tail off = 1 + len(st) markers = [off + m for m in rel_markers] ids = ids[:max_len] return ids, [m for m in markers if m < max_len] def collate(rows: List[Dict], pad_id: int): """rows: dicts with ids, markers, qtype, optional target. Right padding.""" n, L = len(rows), max(len(r["ids"]) for r in rows) kmax = max(len(r["markers"]) for r in rows) ids = torch.full((n, L), pad_id, dtype=torch.long) att = torch.zeros((n, L), dtype=torch.long) mpos = torch.zeros((n, kmax), dtype=torch.long) mmask = torch.zeros((n, kmax), dtype=torch.bool) target = torch.zeros((n, kmax), dtype=torch.float32) for i, r in enumerate(rows): ids[i, :len(r["ids"])] = torch.tensor(r["ids"]) att[i, :len(r["ids"])] = 1 k = len(r["markers"]) mpos[i, :k] = torch.tensor(r["markers"]) mmask[i, :k] = True if "target" in r: target[i, :len(r["target"])] = torch.tensor(r["target"], dtype=torch.float32) return {"input_ids": ids, "attention_mask": att, "marker_pos": mpos, "marker_mask": mmask, "target": target, "qtype": torch.tensor([r["qtype"] for r in rows])} # --------------------------------------------------------------------------- # # Model # --------------------------------------------------------------------------- # class SpikeDecisionModel(nn.Module): """Causal spikewhale trunk + Laya's typed decision head.""" def __init__(self, lm: nn.Module, head_layers: int = 2, n_act: int = 2, dropout: float = 0.1): super().__init__() self.lm = lm d = lm.config.hidden_size nhead = max(1, d // 64) layer = nn.TransformerEncoderLayer(d, nhead, 4 * d, dropout, batch_first=True, norm_first=True) self.head = nn.TransformerEncoder(layer, head_layers, enable_nested_tensor=False) if head_layers > 0 else None self.type_emb = nn.Embedding(3, d) self.scorer = nn.Sequential(nn.LayerNorm(d), nn.Linear(d, d), nn.GELU(), nn.Linear(d, 1)) self.act_head = nn.Sequential(nn.Linear(d + 4, 256), nn.GELU(), nn.Linear(256, n_act)) self.register_buffer("temperature", torch.ones(3)) self.head_checkpointing = False def forward(self, input_ids, attention_mask, marker_pos, marker_mask, qtype, detach_encoder: bool = False): # Right-padded + causal: pads come after every real token, so no mask needed. h = self.lm.model(input_ids=input_ids, use_cache=False)[0] if detach_encoder: h = h.detach() h = h + self.type_emb(qtype)[:, None, :].to(h.dtype) if self.head is not None: pad = ~attention_mask.bool() for layer in self.head.layers: if self.head_checkpointing and self.training and torch.is_grad_enabled(): h = checkpoint(layer, h, src_key_padding_mask=pad, use_reentrant=False) else: h = layer(h, src_key_padding_mask=pad) idx = marker_pos.clamp(min=0)[:, :, None].expand(-1, -1, h.size(-1)) m = torch.gather(h, 1, idx) logits = self.scorer(m).squeeze(-1).float().masked_fill(~marker_mask, -1e4) p = torch.softmax(logits.detach(), -1) k = marker_mask.sum(-1).clamp(min=2).float() ent = -(p * torch.log(p.clamp_min(1e-9))).sum(-1) / torch.log(k) if p.size(-1) >= 2: top2 = p.topk(2, -1).values else: top1 = p.topk(1, -1).values top2 = torch.cat([top1, torch.zeros_like(top1)], dim=-1) feats = torch.stack([top2[:, 0], top2[:, 0] - top2[:, 1], ent, k / 255.0], -1) pooled = h[:, 0].float() # position (causal: sees only itself; kept for API parity) act_logits = self.act_head(torch.cat([pooled, feats], -1)) return logits, act_logits def load_tokenizer(code_dir: str): if code_dir not in sys.path: sys.path.insert(0, code_dir) from spike_tokenizer import SpikeTokenizer return SpikeTokenizer(vocab_file=os.path.join(code_dir, "tokenizer.json")) def load_decision_model(code_dir: str, path: str, device): c = torch.load(path, map_location="cpu", weights_only=False) if code_dir not in sys.path: sys.path.insert(0, code_dir) from config import SpikeWhaleConfig from model_v2 import SpikeWhaleLM lm = SpikeWhaleLM(SpikeWhaleConfig(**c["lm_config"])) dcfg = c["decision_cfg"] model = SpikeDecisionModel(lm, dcfg.get("head_layers", 2), dcfg.get("n_act", 2)) model.load_state_dict(c["model_state"], strict=True) if getattr(lm.config, "tie_word_embeddings", True): lm.tie_weights() return model.to(device).eval(), dcfg, c["lm_config"] # --------------------------------------------------------------------------- # # Calibration and metrics (as in Laya) # --------------------------------------------------------------------------- # def temp_bucket(qtype: int, k: int) -> str: size = "2" if k <= 2 else "3-5" if k <= 5 else "6-10" if k <= 10 else "11+" return "%s:%s" % (QTYPE_NAMES[int(qtype)], size) def clamp_temperature(t, lo: float = TEMP_MIN, hi: float = TEMP_MAX) -> float: try: t = float(t) except (TypeError, ValueError): return 1.0 if t != t or t in (float("inf"), float("-inf")): return 1.0 return min(hi, max(lo, t)) def ece_score(conf, correct, bins: int = 15) -> float: conf, correct = np.asarray(conf), np.asarray(correct) if len(conf) == 0: return float("nan") edges = np.linspace(0, 1, bins + 1) e = 0.0 for i, (lo, hi) in enumerate(zip(edges[:-1], edges[1:])): sel = (conf >= lo if i == 0 else conf > lo) & (conf <= hi) if sel.any(): e += sel.mean() * abs(conf[sel].mean() - correct[sel].mean()) return float(e) def answer_confidence(p, k: int) -> float: return 1.0 if k < 1 else float(np.clip(np.max(p[:k]), 0.0, 1.0)) def confidence_from_probs(p, k: int) -> float: if k < 2: return 1.0 p = p[:k] ent = -(p * np.log(np.clip(p, 1e-12, 1.0))).sum() return float(np.clip(1.0 - ent / math.log(k), 0.0, 1.0)) # --------------------------------------------------------------------------- #