File size: 4,849 Bytes
f97b180
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
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?"))