bert4jev / modeling_bert4jev.py
ukung's picture
Upload modeling_bert4jev.py with huggingface_hub
39eca46 verified
Raw History Blame Contribute Delete
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"))
@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