Download runtime/jevlike/model.py from Cem13/kodama-core: direct link, hf CLI and curl.
- Browser
- Download file 8.37 kB
-
https://huggingface.co/Cem13/kodama-core/resolve/main/runtime/jevlike/model.py
- Command line
-
hf download hf://Cem13/kodama-core/runtime/jevlike/model.py
-
curl -L -o model.py https://huggingface.co/Cem13/kodama-core/resolve/main/runtime/jevlike/model.py
8.37 kB
| """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 | |
| 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" | |
| 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) | |
| 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) | |
| 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() | |