# Extracted from PostHog/jeeves loader/dataloader.py (MIT). # Only training imports/unused nodes were removed; see LICENSE_JEEVES_CODE. from __future__ import annotations import re from dataclasses import dataclass from transformers import AutoTokenizer from dataformat import DataFormat, Question STATE, Q, OPT, OPT_END, DECIDE = ('<|fim_prefix|>', '<|fim_middle|>', '<|box_start|>', '<|box_end|>', '<|fim_suffix|>') THINK, THINK_END = ('', '') CONTROL_RE = re.compile('<\\|([A-Za-z0-9_]+)\\|>') def sanitize(text: str) -> str: return CONTROL_RE.sub('<¦\\1¦>', text) @dataclass class Example: prompt: list[int] remainder: list[int] label: int target: list[float] | None n_options: int source: str source_id: int record_id: str question_id: str think: list[int] | None = None plain_prompt: list[int] | None = None @property def full(self) -> list[int]: return self.prompt + self.remainder @dataclass(frozen=True) class Markers: pad_id: int opt_end_id: int think_end_id: int empty_think: tuple[int, ...] pad_multiple: int class Encoder: FORMAT = 'markers-v3-plainchains' def __init__(self, tokenizer_repo: str='Qwen/Qwen3.5-9B', pad_multiple: int=128): self.tok = AutoTokenizer.from_pretrained(tokenizer_repo) self.pad_multiple = pad_multiple ids = self.tok.convert_tokens_to_ids([STATE, Q, OPT, OPT_END, DECIDE, THINK, THINK_END]) if any((i is None or i == self.tok.unk_token_id for i in ids)): raise ValueError('marker tokens missing from tokenizer') self.state_id, self.q_id, self.opt_id, self.opt_end_id, self.decide_id, self.think_id, self.think_end_id = ids self.pad_id = self.tok.pad_token_id special = set(self.tok.all_special_ids) | set(self.tok.get_added_vocab().values()) self.banned_ids = sorted(special - {self.think_end_id}) self.empty_think = self.encode('\n') @property def markers(self) -> Markers: return Markers(self.pad_id, self.opt_end_id, self.think_end_id, tuple(self.empty_think), self.pad_multiple) def encode(self, text: str) -> list[int]: return self.tok(text, add_special_tokens=False).input_ids @staticmethod def option_block(q: Question) -> str: return ''.join((f'{OPT}{sanitize(o)}{OPT_END}\n' for o in q.options())) def chat(self, content: str) -> str: return self.tok.apply_chat_template([{'role': 'user', 'content': content}], tokenize=False, add_generation_prompt=True, enable_thinking=True) def prompt_text(self, record: DataFormat, q: Question) -> str: return self.chat(f'{STATE}{sanitize(record.state_text())}\n{Q}{sanitize(q.instruction_text())}\n{self.option_block(q)}') def plain_prompt_text(self, record: DataFormat, q: Question) -> str: options = '\n'.join((sanitize(o) for o in q.options())) return self.chat(f'Context:\n{sanitize(record.state_text())}\n\nQuestion: {sanitize(q.instruction_text())}\n\nOptions:\n{options}') def remainder_text(self, q: Question) -> str: return f'{THINK_END}\n\n{self.option_block(q)}{DECIDE}' def example(self, record: DataFormat, q: Question, source_id: int) -> Example: prompt = self.encode(self.prompt_text(record, q)) remainder = self.encode(self.remainder_text(q)) if prompt[-2:] != self.encode(f'{THINK}\n') or remainder[0] != self.think_end_id or remainder[-1] != self.decide_id: raise ValueError(f'unexpected layout for {record.id}/{q.id}') if remainder.count(self.opt_end_id) != len(q.options()): raise ValueError(f'option boundary count mismatch for {record.id}/{q.id}') return Example(prompt=prompt, remainder=remainder, label=q.label_index(), target=q.target_vector(), n_options=len(q.options()), source=record.source, source_id=source_id, record_id=record.id, question_id=q.id, plain_prompt=self.encode(self.plain_prompt_text(record, q))) def readout_positions(ids: list[int], m: Markers) -> tuple[list[int], int]: start = len(ids) - 1 - ids[::-1].index(m.think_end_id) opts = [i for i in range(start, len(ids)) if ids[i] == m.opt_end_id] return (opts, len(ids) - 1)