import json import os import torch import torch.nn as nn from transformers import PreTrainedModel, DebertaV2Model, AutoTokenizer from safetensors.torch import load_file try: from .configuration_bert4jev import Bert4JevConfig except Exception: from configuration_bert4jev import Bert4JevConfig MARKERS = ["[STATE]", "[Q]", "[OPT]"] NOUL_OPTIONS = ("no", "yes") OPT_SLOTS = 256 class _Collator: def __init__(self, tok, max_state_tokens=256, max_len=512): self.tok = tok self.max_state = max_state_tokens self.max_len = max_len self.cls, self.sep = tok.cls_token_id, tok.sep_token_id self.m_state, self.m_q, self.m_opt = tok.convert_tokens_to_ids(MARKERS) self.OM = OPT_SLOTS def _ids(self, text): return self.tok(text, add_special_tokens=False)["input_ids"] def encode_one(self, state, questions): ids = [self.cls, self.m_state] + self._ids(state)[: self.max_state] positions, q_positions, spans = [], [], [] for qi, q in enumerate(questions): q_positions.append(len(ids)) t = self._ids(q["instructions"]) spans.append((-(qi + 1), len(ids) + 1, len(ids) + 1 + len(t))) ids += [self.m_q] + t pos = [] for oi, o in enumerate(q["options"]): pos.append(len(ids)) t = self._ids(o) spans.append((qi * self.OM + oi, len(ids) + 1, len(ids) + 1 + len(t))) ids += [self.m_opt] + t positions.append(pos) ids.append(self.sep) if len(ids) > self.max_len: raise ValueError("sequence %d > max_len %d" % (len(ids), self.max_len)) return ids, positions, q_positions, spans def __call__(self, items, device=None): enc = [self.encode_one(s, qs) for s, qs in items] L = max(len(e[0]) for e in enc) Qm = max(len(e[1]) for e in enc) Om = max(len(p) for e in enc for p in e[1]) pad = self.tok.pad_token_id input_ids = torch.full((len(items), L), pad, dtype=torch.long) attn = torch.zeros((len(items), L), dtype=torch.long) opt_pos = torch.full((len(items), Qm, Om), -1, dtype=torch.long) q_pos = torch.full((len(items), Qm), -1, dtype=torch.long) n_slots = Qm * Om + Qm seg = torch.full((len(items), L), n_slots, dtype=torch.long) for b, (ids, positions, q_positions, spans) in enumerate(enc): input_ids[b, : len(ids)] = torch.tensor(ids) attn[b, : len(ids)] = 1 q_pos[b, : len(q_positions)] = torch.tensor(q_positions) for qi, pos in enumerate(positions): opt_pos[b, qi, : len(pos)] = torch.tensor(pos) for sid, st, en in spans: slot = (Qm * Om + (-sid - 1)) if sid < 0 else ((sid // self.OM) * self.OM + sid % self.OM) seg[b, st:en] = slot out = {"input_ids": input_ids, "attention_mask": attn, "opt_pos": opt_pos, "opt_mask": opt_pos >= 0, "q_pos": q_pos, "seg": seg} return {k: (v.to(device) if device is not None else v) for k, v in out.items()} def _readout(kind, probs): k = max(range(len(probs)), key=lambda i: probs[i]) conf = float(probs[k]) if kind == "choice": return {"choice_idx": k, "confidence": conf} if kind == "score": return {"score": float(sum(i * p for i, p in enumerate(probs))), "confidence": conf} return {"noul": float(probs[1])} class Bert4JevModel(PreTrainedModel): config_class = Bert4JevConfig base_model_prefix = "deberta" def __init__(self, config): super().__init__(config) self.deberta = DebertaV2Model(config) h = int(getattr(config, "hidden", getattr(config, "hidden_size", 1024))) self.head = nn.Sequential(nn.Linear(3 * h, h), nn.GELU(), nn.Linear(h, 1)) self.temperature = float(getattr(config, "temperature", 1.0)) self.pool = getattr(config, "pool", "span") self.markers = list(getattr(config, "markers", MARKERS)) self.max_state_tokens = int(getattr(config, "max_state_tokens", 256)) self.max_len = int(getattr(config, "max_len", 512)) self._tokenizer = None self._collator = None def _ensure_tokenizer(self, path): if self._tokenizer is None: self._tokenizer = AutoTokenizer.from_pretrained(path) self._collator = _Collator(self._tokenizer, self.max_state_tokens, self.max_len) return self._collator def forward(self, input_ids, attention_mask, opt_pos, opt_mask, q_pos, seg): h = self.deberta(input_ids=input_ids, attention_mask=attention_mask).last_hidden_state B, Qm, Om = opt_pos.shape H = h.size(-1) hp = h.float() n_slots = Qm * Om + Qm sums = hp.new_zeros((B, n_slots + 1, H)).scatter_add_(1, seg.unsqueeze(-1).expand(-1, -1, H), hp) cnt = hp.new_zeros((B, n_slots + 1)).scatter_add_(1, seg, torch.ones_like(seg, dtype=hp.dtype)).clamp(min=1).unsqueeze(-1) mean = sums / cnt g = mean[:, : Qm * Om].reshape(B, Qm, Om, H) q = mean[:, Qm * Om: Qm * Om + Qm].unsqueeze(2).expand(-1, -1, Om, -1) g = torch.cat([q, g, q * g], dim=-1).to(self.head[0].weight.dtype) logits = self.head(g).squeeze(-1) return logits.masked_fill(~opt_mask, float("-inf")) @classmethod def _resolve_dir(cls, name, token=None): if os.path.isdir(name): return name from huggingface_hub import snapshot_download return snapshot_download(name, token=token) @classmethod def from_pretrained(cls, pretrained_model_name_or_path, *model_args, **kwargs): token = kwargs.pop("token", None) or kwargs.pop("use_auth_token", None) d = cls._resolve_dir(pretrained_model_name_or_path, token=token) config = Bert4JevConfig.from_pretrained(d) dtype = torch.float16 if str(getattr(config, "dtype", "")).endswith("float16") else torch.float32 ai = getattr(config, "attn_implementation", "eager") try: backbone = DebertaV2Model.from_pretrained(d, attn_implementation=ai, dtype=dtype) except TypeError: backbone = DebertaV2Model.from_pretrained(d, attn_implementation=ai, torch_dtype=dtype) model = cls(config) model.deberta = backbone model.head.load_state_dict(load_file(os.path.join(d, "head.safetensors"))) model.to(dtype) model.temperature = float(getattr(config, "temperature", 1.0)) dev = kwargs.get("device") or ("cuda" if torch.cuda.is_available() else "cpu") model.to(dev).eval() model._ensure_tokenizer(d) return model @staticmethod def _question(q): kind = q["type"] if kind == "noul": return {"kind": "noul", "instructions": q["instructions"], "options": list(NOUL_OPTIONS)} return {"kind": kind, "instructions": q["instructions"], "options": list(q["options"])} @torch.no_grad() def decide(self, state, questions): if self._collator is None: raise ValueError("tokenizer not initialized; load with from_pretrained()") qs = [self._question(q) for q in questions] b = self._collator([(state, qs)], self.deberta.device) logits = self.forward(b["input_ids"], b["attention_mask"], b["opt_pos"], b["opt_mask"], b["q_pos"], b["seg"]).float() probs = (logits / self.temperature).softmax(-1)[0] out = [] for qi, q in enumerate(qs): p = probs[qi, : len(q["options"])].tolist() r = _readout(q["kind"], p) if q["kind"] == "choice": out.append({"choice": q["options"][r["choice_idx"]], "probabilities": dict(zip(q["options"], p)), "confidence": r["confidence"]}) elif q["kind"] == "score": out.append({"score": r["score"], "probabilities": dict(zip(q["options"], p)), "confidence": r["confidence"]}) else: out.append({"noul": r["noul"]}) return out