Ines-1 / training /code /decisions.py
Endikavi's picture
Ines-1 RC1 (private staging; release commit b7f5644)
61b6fb9 verified
Raw History Blame Contribute Delete
7.92 kB
"""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