import torch from transformers import AutoModelForTokenClassification, AutoTokenizer MAXLEN = 384 STRIDE = 96 SEED_T = 0.35 EXT_T = 0.21 class Undyne: def __init__(self, path="tiagozip/undyne", device="cpu"): self.tok = AutoTokenizer.from_pretrained(path) self.model = AutoModelForTokenClassification.from_pretrained(path).eval().to(device) self.device = device def _window(self, question, answer, start, end): chunk = answer[start:end] enc = self.tok(question, chunk, truncation="only_second", max_length=MAXLEN, return_offsets_mapping=True, return_tensors="pt") if question.strip() \ else self.tok(chunk, truncation=True, max_length=MAXLEN, return_offsets_mapping=True, return_tensors="pt") ans_seq = 1 if question.strip() else 0 offsets = enc.pop("offset_mapping")[0].tolist() with torch.inference_mode(): probs = self.model(**{k: v.to(self.device) for k, v in enc.items()}).logits[0].softmax(-1).cpu() inspan = (probs[:, 1] + probs[:, 2]).tolist() seq = enc.sequence_ids() return [(b + start, e + start, inspan[i]) for i, ((b, e), s) in enumerate(zip(offsets, seq)) if s == ans_seq and e > b] def spans(self, answer, question=""): if not answer.strip(): return [] budget = MAXLEN - len(self.tok(question)["input_ids"]) - 8 if question.strip() else MAXLEN - 4 offs = self.tok(answer, add_special_tokens=False, return_offsets_mapping=True)["offset_mapping"] if len(offs) <= budget: toks = self._window(question, answer, 0, len(answer)) else: step, seen = max(1, budget - STRIDE), {} for s0 in range(0, len(offs), step): chunk = offs[s0:s0 + budget] if not chunk: break for b, e, p in self._window(question, answer, chunk[0][0], chunk[-1][1]): seen[(b, e)] = max(seen.get((b, e), 0.0), p) if s0 + budget >= len(offs): break toks = [(b, e, p) for (b, e), p in sorted(seen.items())] return self._decode(toks, answer) def _decode(self, toks, answer): idx = range(len(toks)) on = {i for i in idx if toks[i][2] > SEED_T} for k in list(on): for step in (-1, 1): j = k + step while 0 <= j < len(toks) and j not in on and toks[j][2] > EXT_T: on.add(j) j += step raw, cur = [], None for i in idx: b, e, _ = toks[i] if i in on: if cur and b - cur[1] <= 1 and "\n" not in answer[cur[1]:b]: cur[1] = e else: if cur: raw.append(cur) cur = [b, e] elif cur: raw.append(cur) cur = None if cur: raw.append(cur) out = [] for b, e in raw: for pb, pe in self._split_lines(answer, b, e): while pb < pe and answer[pb] in "-*• \t": pb += 1 while pb > 0 and answer[pb - 1].isalnum(): pb -= 1 while pe < len(answer) and answer[pe].isalnum(): pe += 1 if out and pb - out[-1][1] <= 2 and "\n" not in answer[out[-1][1]:pb] and not any(c in ".;" for c in answer[out[-1][1]:pb]): out[-1][1] = pe else: out.append([pb, pe]) return [(b, e) for b, e in out if len(answer[b:e].strip()) >= 3] @staticmethod def _split_lines(answer, b, e): parts, start = [], b for i in range(b, e): if answer[i] == "\n": if i > start: parts.append((start, i)) start = i + 1 if e > start: parts.append((start, e)) return parts def highlight(self, answer, question="", fmt="**{}**"): out, last = "", 0 for b, e in self.spans(answer, question): out += answer[last:b] + fmt.format(answer[b:e]) last = e return out + answer[last:] if __name__ == "__main__": m = Undyne(".") a = ("The sky is blue because of a phenomenon called Rayleigh scattering, named after the 19th-century British " "physicist Lord Rayleigh, who also discovered argon. Sunlight contains all colors of the visible spectrum, " "and when it hits molecules in Earth's atmosphere, shorter wavelengths like blue and violet scatter far more " "than longer wavelengths like red and orange.") print(m.highlight(a, "Why is the sky blue?")) print() print(m.highlight(a, "Who is Rayleigh scattering named after?"))