undyne / example.py
tiagozip's picture
undyne: 13.5M answer-span highlighter distilled from an 82M teacher
f97b180 verified
Raw History Blame Contribute Delete
4.85 kB
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?"))