"""Turn a Decision into token ids, and say where each candidate ends. The layout is state -> question -> candidates. The backbone is causal, so a candidate's last token has attended to the state, the question and every candidate before it; its hidden state is what the head scores. Candidate order is shuffled during training (for `choice`) so no position is favoured. Ids are built segment by segment rather than by tokenising one string and searching for offsets. Training and inference share this function, so the two can never disagree about where a candidate ends -- which is the whole contract. """ from __future__ import annotations from dataclasses import dataclass from .schema import Decision SYSTEM = ("You are a decision model. Read the state, then judge every candidate " "answer to the question.") _HEADER = { "choice": "Question (pick one): {q}\nCandidates:", "null": "Is this true? {q}\nCandidates:", "score": "Question (rate on the scale): {q}\nCandidates:", } @dataclass class Rendered: input_ids: list[int] option_positions: list[int] # index of each candidate's last token def _frame(tok) -> tuple[str, str]: """(prefix, suffix) of the chat frame around the user turn, or plain text for a tokenizer with no template.""" marker = "\x00BODY\x00" if getattr(tok, "chat_template", None): text = tok.apply_chat_template( [{"role": "system", "content": SYSTEM}, {"role": "user", "content": marker}], tokenize=False, add_generation_prompt=True) pre, post = text.split(marker) bos = getattr(tok, "bos_token", None) if bos and pre.startswith(bos): pre = pre[len(bos):] # added back explicitly below return pre, post return f"{SYSTEM}\n\n", "\n" class Renderer: def __init__(self, tok, max_state_tokens: int = 384, max_option_tokens: int = 32, max_question_tokens: int = 96): self.tok = tok self.max_state = max_state_tokens self.max_option = max_option_tokens self.max_question = max_question_tokens pre, post = _frame(tok) self._pre = self._ids(pre) self._post = self._ids(post) bos_id = getattr(tok, "bos_token_id", None) adds_bos = bool(tok("a").input_ids[:1] == [bos_id]) if bos_id is not None else False self._bos = [bos_id] if adds_bos else [] def _ids(self, s: str, limit: int | None = None) -> list[int]: ids = self.tok(s, add_special_tokens=False).input_ids return ids[:limit] if limit else ids def render(self, d: Decision, candidates: list[str] | None = None) -> Rendered: cands = candidates if candidates is not None else d.candidates() ids = self._bos + list(self._pre) if d.state.strip(): ids += self._ids("State:\n") ids += self._ids(d.state.strip(), self.max_state) ids += self._ids("\n\n") q = self._ids(d.question.strip(), self.max_question) head = _HEADER[d.type].split("{q}") ids += self._ids(head[0]) + q + self._ids(head[1]) positions = [] for i, c in enumerate(cands, 1): ids += self._ids(f"\n[{i}] ") ids += self._ids(c.strip(), self.max_option) or self._ids("(empty)") positions.append(len(ids) - 1) ids += self._post return Rendered(ids, positions)