malleshmadapathi's picture
DECISIONx v2: Qwen2.5-1.5B + LoRA decision model, inference package, eval reports
065ec46 verified
Raw History Blame Contribute Delete
3.43 kB
"""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)