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
|