Cem13's picture
Release Kodama Core with weights, runtime, attribution, and evaluation
c7893fa verified
Raw History Blame Contribute Delete
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
@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()