Download modeling_bert4jev.py from ukung/bert4jev: direct link, hf CLI and curl.
- Browser
- Download file 8.18 kB
-
https://huggingface.co/ukung/bert4jev/resolve/main/modeling_bert4jev.py
- Command line
-
hf download hf://ukung/bert4jev/modeling_bert4jev.py
-
curl -L -o modeling_bert4jev.py https://huggingface.co/ukung/bert4jev/resolve/main/modeling_bert4jev.py
8.18 kB
| 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")) | |
| 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) | |
| 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 | |
| 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"])} | |
| 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 | |