jeeves-mlx / jeeves_format.py
cowWhySo's picture
Publish Jeeves MLX 4bit (CPU diagnostics; not benchmarked)
6f4220a verified
Raw History Blame Contribute Delete
4.24 kB
# 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 = ('<think>', '</think>')
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)