"""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