File size: 7,921 Bytes
61b6fb9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
"""Decision reading for mini-v41: typed questions (choice / score / noul) answered by the
probability of each option LETTER at the first position of the assistant reply.

Same idea as letter-probability readouts of typed decisions, done here with a causal model: the options are
listed as A) B) C) ... and the logits of the letter tokens are read after the assistant
marker. The id of each letter comes from rendering a reply with the chat template, so it is
the token the model would really emit.

Common case format (one JSON line per case):
    {"id", "grupo", "lang": "en"|"es", "state", "questions": {qid: {type, instructions, criteria}},
     "gold": {qid: [valid keys]}, "suave": {qid: {key: prob}} (optional, soft target)}
Keys: choice -> its criteria ids; score -> "0".."n"; noul -> "true"/"false".
"""
from __future__ import annotations

import collections
import json
import random
from pathlib import Path

import torch

LETTERS = [chr(ord("A") + i) for i in range(26)]

TEMPLATES = {
    "es": {"case": "Caso", "question": "Pregunta", "options": "Opciones",
           "answer": "Responde solo con la letra de la opción correcta.", "yes": "sí", "no": "no"},
    "en": {"case": "Case", "question": "Question", "options": "Options",
           "answer": "Answer only with the letter of the correct option.", "yes": "yes", "no": "no"},
}


def read_jsonl(path):
    return [json.loads(line) for line in Path(path).read_text(encoding="utf-8").splitlines() if line.strip()]


def options(q, lang="es"):
    """[(key, text)] in presentation order. noul is always yes first, then no."""
    kind, crit = q["type"], q.get("criteria")
    if kind == "choice":
        return [(k, str(v)) for k, v in crit.items()]
    if kind == "score":
        return [(str(i), str(v)) for i, v in enumerate(crit)]
    if kind == "noul":
        t = TEMPLATES[lang]
        crit = crit or {}
        yes, no = crit.get("true"), crit.get("false")
        return [("true", "%s: %s" % (t["yes"], yes) if yes else t["yes"]),
                ("false", "%s: %s" % (t["no"], no) if no else t["no"])]
    raise ValueError("unsupported question type %r" % kind)


class Reader:
    """Builds the prompt for one question and scores the option letters."""

    def __init__(self, im, max_len=None):
        from mini_v41.chat import ChatTemplate, Message
        self.im, self.tok, self.Message = im, im.tokenizer, Message
        self.max_len = int(max_len or im.max_sequence_length)
        self.tpl = ChatTemplate(im.tokenizer, max_length=self.max_len)
        self._letter_ids = {}

    def letter_ids(self, n):
        if n not in self._letter_ids:
            user = [self.Message("user", "x")]
            prefix = self.tpl.render(user, for_generation=True).input_ids
            ids = []
            for letter in LETTERS[:n]:
                full = self.tpl.render(user + [self.Message("assistant", letter)]).input_ids
                assert full[:len(prefix)] == prefix, "chat template prefix is not stable"
                ids.append(full[len(prefix)])
            assert len(set(ids)) == n, "two letters share a token: %s" % ids
            self._letter_ids[n] = ids
        return self._letter_ids[n]

    def prompt(self, state, q, lang="es", order=None):
        """(input ids, keys in presentation order). `order` permutes the options (training)."""
        t = TEMPLATES[lang]
        opts = options(q, lang)
        if order is not None:
            opts = [opts[i] for i in order]
        if len(opts) > len(LETTERS):
            raise ValueError("%d options; at most %d" % (len(opts), len(LETTERS)))
        text = state if isinstance(state, str) else json.dumps(state, ensure_ascii=False)
        st_ids = self.tok.encode(text)
        cap = 400  # tokens per option; lowered if it does not fit
        while True:
            lines = []
            for letter, (_, d) in zip(LETTERS, opts):
                di = self.tok.encode(d)
                lines.append("%s) %s" % (letter, self.tok.decode(di[:cap]) + ("…" if len(di) > cap else "")))
            head = "%s: %s\n%s:\n%s\n\n%s" % (t["question"], q["instructions"], t["options"], "\n".join(lines), t["answer"])
            fixed = len(self.tpl.render([self.Message("user", "%s:\n\n\n%s" % (t["case"], head))],
                                        for_generation=True).input_ids)
            room = self.max_len - fixed - 8
            if room >= min(len(st_ids), 256) or cap <= 24:
                break
            cap = max(24, int(cap * 0.7))
        if len(st_ids) > room:  # keep the start and the end of the case
            half = max(room // 2, 1)
            text = self.tok.decode(st_ids[:half]) + "\n[…]\n" + self.tok.decode(st_ids[-(room - half - 4):])
        rendered = self.tpl.render([self.Message("user", "%s:\n%s\n\n%s" % (t["case"], text, head))],
                                   for_generation=True)
        if rendered is None:
            raise ValueError("does not fit even after trimming")
        return rendered.input_ids[-self.max_len:], [k for k, _ in opts]

    def logits_train(self, ids, n):
        """Full forward with grad (training). Returns (letter logits, model output)."""
        x = torch.tensor([ids], device=self.im.device)
        with torch.autocast("cuda", dtype=torch.bfloat16):
            out = self.im.model(x)
        return out.logits[0, -1].float()[self.letter_ids(n)], out

    @torch.no_grad()
    def logits_eval(self, ids, n):
        """prefill(num_logits=1): the decoder only runs over the last positions (constant cost)."""
        x = torch.tensor([ids], device=self.im.device)
        with torch.autocast("cuda", dtype=torch.bfloat16):
            logits, _ = self.im.model.prefill(x, num_logits=1)
        return logits[0, -1].float()[self.letter_ids(n)]


def evaluate(reader, cases, details=None):
    """Accuracy per group (and per group/question for multi-question groups)."""
    was_training = reader.im.model.training
    reader.im.model.eval()
    by = collections.defaultdict(lambda: [0, 0])
    for c in cases:
        lang = c.get("lang", "es")
        for qid, q in c["questions"].items():
            ids, keys = reader.prompt(c["state"], q, lang)
            z = reader.logits_eval(ids, len(keys))
            pred = keys[int(z.argmax())]
            ok = pred in c["gold"][qid]
            for k in {c["grupo"], c["grupo"].split("/")[0]}:
                by[k][0] += ok
                by[k][1] += 1
            if details is not None:
                details.append({"id": c["id"], "grupo": c["grupo"], "q": qid, "pred": pred,
                                "gold": c["gold"][qid], "ok": ok})
    reader.im.model.train(was_training)
    return {k: "%d/%d (%.1f %%)" % (a, b, 100 * a / b) for k, (a, b) in sorted(by.items())}


def accuracy(res, key):
    v = res.get(key)
    if not v:
        return None
    a, b = v.split(" ")[0].split("/")
    return int(a) / int(b)


def examples(reader, cases, rng, shuffle_choice=True):
    """(ids, n options, target distribution over presented options) per question."""
    out = []
    for c in cases:
        lang = c.get("lang", "es")
        for qid, q in c["questions"].items():
            gold = c["gold"][qid]
            soft = (c.get("suave") or {}).get(qid)
            if not soft and len(gold) != 1:
                continue
            n = len(options(q, lang))
            order = list(range(n))
            if shuffle_choice and q["type"] == "choice":
                rng.shuffle(order)
            ids, keys = reader.prompt(c["state"], q, lang, order)
            if soft:
                p = [float(soft.get(k, 0.0)) for k in keys]
                s = sum(p)
                target = [x / s for x in p] if s > 0 else None
            else:
                target = [1.0 if k == gold[0] else 0.0 for k in keys]
            if target:
                out.append((ids, n, target))
    rng.shuffle(out)
    return out