"""DecisionModel: ModernBERT backbone -> 2-layer decision head -> per-marker scorer -> softmax per question, plus an escalate head predicting whether the argmax answer is correct.""" from __future__ import annotations import json import math import os from dataclasses import asdict, dataclass, field import torch import torch.nn.functional as F from safetensors.torch import load_file, save_file from torch import nn from transformers import AutoConfig, AutoModel from jevlike.serialize import TYPE_IDS N_TYPES = len(TYPE_IDS) N_ESC_FEATS = 3 + N_TYPES # max prob, normalized entropy, log n_options, one-hot type @dataclass class DecisionConfig: backbone: str = "answerdotai/ModernBERT-base" backbone_config: dict = field(default_factory=dict) # filled from the backbone; makes ckpts self-contained attn_implementation: str = "sdpa" head_layers: int = 2 head_heads: int = 12 head_ff: int = 2048 dropout: float = 0.1 esc_hidden: int = 256 max_len: int = 512 head_max_len: int = 192 escalate_threshold: float = 0.5 # escalate when escalate_prob = P(argmax wrong) > threshold name: str = "jevlike-base" @classmethod def load(cls, path: str) -> "DecisionConfig": with open(os.path.join(path, "config.json")) as f: d = json.load(f) return cls(**{k: v for k, v in d.items() if k in cls.__dataclass_fields__}) class DecisionModel(nn.Module): def __init__(self, cfg: DecisionConfig, backbone: nn.Module | None = None): super().__init__() if backbone is None: # architecture only (weights come from a checkpoint) bc = dict(cfg.backbone_config) backbone = AutoModel.from_config(AutoConfig.for_model(bc.pop("model_type"), **bc), attn_implementation=cfg.attn_implementation) self.cfg, self.backbone = cfg, backbone cfg.backbone_config = backbone.config.to_dict() D = backbone.config.hidden_size if D % cfg.head_heads: # e.g. ModernBERT-large (1024): use 64-dim heads cfg.head_heads = D // 64 self.head = nn.ModuleList(nn.TransformerEncoderLayer( D, cfg.head_heads, cfg.head_ff, cfg.dropout, activation="gelu", batch_first=True, norm_first=True) for _ in range(cfg.head_layers)) self.head_norm = nn.LayerNorm(D) self.scorer = nn.Sequential(nn.Linear(D, D), nn.GELU(), nn.Dropout(cfg.dropout), nn.Linear(D, 1)) self.escalate = nn.Sequential(nn.Linear(D + N_ESC_FEATS, cfg.esc_hidden), nn.GELU(), nn.Dropout(cfg.dropout), nn.Linear(cfg.esc_hidden, 1)) # skip path: escalate logit starts at logit(max prob), i.e. "P(correct) = own confidence" self.esc_skip = nn.Parameter(torch.tensor([1.0, 0.0])) for lin in (self.scorer[-1], self.escalate[-1]): # start near-uniform / uninformed nn.init.normal_(lin.weight, std=0.02) nn.init.zeros_(lin.bias) @classmethod def from_backbone(cls, cfg: DecisionConfig) -> "DecisionModel": """New model with pretrained backbone weights and freshly initialized heads.""" bb = AutoModel.from_pretrained(cfg.backbone, attn_implementation=cfg.attn_implementation) return cls(cfg, bb) def head_parameters(self): return [p for n, p in self.named_parameters() if not n.startswith("backbone.")] def forward(self, input_ids: torch.Tensor, attention_mask: torch.Tensor, marker_idx: torch.Tensor, marker_seg: torch.Tensor, marker_rank: torch.Tensor, n_options: torch.Tensor, qtype: torch.Tensor, k_max: int | None = None, inputs_embeds: torch.Tensor | None = None, **_) -> dict[str, torch.Tensor]: """Returns ``logits`` [B, K] (-inf beyond each question's n_options; one question per row) and ``escalate_logit`` [B] (logit of P(argmax correct)). The escalate head sees the [CLS] state plus detached answer-distribution features, so it cannot move the answers. ``inputs_embeds`` [B, T, D] (optional) replaces the token-embedding lookup of ``input_ids`` (used by the multimodal prototype, jevlike/mm.py, to splice in image tokens).""" if inputs_embeds is not None: h = self.backbone(inputs_embeds=inputs_embeds, attention_mask=attention_mask).last_hidden_state else: h = self.backbone(input_ids=input_ids, attention_mask=attention_mask).last_hidden_state pad = attention_mask == 0 for layer in self.head: h = layer(h, src_key_padding_mask=pad) h = self.head_norm(h) B, T, D = h.shape marker_logit = self.scorer(h.reshape(B * T, D)[marker_idx]).squeeze(-1).float() K = k_max or int(n_options.max()) logits = torch.full((B, K), -math.inf, device=h.device).index_put((marker_seg, marker_rank), marker_logit) with torch.no_grad(): # answer-distribution features, detached logp = logits.log_softmax(-1) p = logp.exp() ent = -(p * logp.masked_fill(p == 0, 0)).sum(-1) / n_options.float().log() maxp = p.max(-1).values feats = torch.cat([maxp[:, None], ent[:, None], (n_options.float().log() / math.log(256))[:, None], F.one_hot(qtype, N_TYPES).float()], -1) conf_logit = torch.logit(maxp, eps=1e-6) esc = self.escalate(torch.cat([h[:, 0].float(), feats], -1)).squeeze(-1).float() esc = esc + self.esc_skip[0] * conf_logit + self.esc_skip[1] return {"logits": logits, "escalate_logit": esc} # ------------------------------------------------------------ io def save_pretrained(self, path: str, tokenizer=None) -> None: os.makedirs(path, exist_ok=True) with open(os.path.join(path, "config.json"), "w") as f: json.dump(asdict(self.cfg), f, indent=1, default=str) save_file({k: v.detach().contiguous().cpu() for k, v in self.state_dict().items()}, os.path.join(path, "model.safetensors")) if tokenizer is not None: tokenizer.save_pretrained(path) @classmethod def from_pretrained(cls, path: str, device: str | torch.device = "cpu") -> "DecisionModel": model = cls(DecisionConfig.load(path)) model.load_state_dict(load_file(os.path.join(path, "model.safetensors"), device=str(device))) return model.to(device) # ---------------------------------------------------------------- losses & metrics def compute_loss(out: dict[str, torch.Tensor], batch: dict[str, torch.Tensor], rps_weight: float = 1.0, esc_weight: float = 0.2) -> tuple[torch.Tensor, dict[str, float]]: """Per-question loss = CE + rps_weight * RPS (score questions) + esc_weight * BCE(escalate). Returns (sum over questions, stats); callers normalize by the number of questions per step. No label smoothing: all terms are strictly proper scoring rules. """ logits, labels, n = out["logits"], batch["labels"], batch["n_options"] logp = logits.log_softmax(-1) nll = -logp.gather(1, labels[:, None]).squeeze(1) cdf_err = logp.exp().cumsum(-1) - F.one_hot(labels, logits.shape[1]).float().cumsum(-1) rps = cdf_err.pow(2).sum(-1) / (n - 1).clamp(min=1) is_score = (batch["qtype"] == TYPE_IDS["score"]).float() correct = (logits.detach().argmax(-1) == labels).float() bce = F.binary_cross_entropy_with_logits(out["escalate_logit"], correct, reduction="none") per = nll + rps_weight * rps * is_score + esc_weight * bce stats = {"nll": nll.sum().item(), "rps": (rps * is_score).sum().item(), "n_score": is_score.sum().item(), "bce": bce.sum().item(), "correct": correct.sum().item(), "n": float(len(labels))} return per.sum(), stats def ece(conf: torch.Tensor, correct: torch.Tensor, n_bins: int = 15) -> float: """Expected calibration error of top-1 confidence (equal-width bins).""" if len(conf) == 0: return float("nan") bins = (conf.float() * n_bins).long().clamp(0, n_bins - 1) tot = torch.zeros(n_bins).index_add_(0, bins, torch.ones_like(conf)) c = torch.zeros(n_bins).index_add_(0, bins, conf.float()) a = torch.zeros(n_bins).index_add_(0, bins, correct.float()) return ((c - a).abs().sum() / len(conf)).item()