alexwengg's picture
decision-modernbert-base Core ML (CC BY-NC 4.0): buckets, engine config, tokenizer, reference runtime
26b3403 verified
Raw
History Blame Contribute Delete
6.34 kB
"""Encoder rendering and windowing shared by training, evaluation and the Core ML engine.
A window is one encoder sequence:
[CLS] <type> <question> [SEP] ([MASK] <option>)* [SEP] <context slice> [SEP]
Each option is read at its [MASK] marker. Nothing is truncated: long questions keep a
head and tail in every window and their middle joins the scanned context; option lists
that do not fit are packed into several groups; an option too long for one window is
split into pieces, each with its own marker; the context (question middle + state) is
scanned in overlapping slices. An option's logit is the log-mean-exp over every marker
occurrence of it (all windows and pieces), so training and inference pool identically.
"""
from __future__ import annotations
import json
from dataclasses import dataclass, field
from typing import Any
from tokenizers import Tokenizer
CLS, SEP, PAD, MASK = 50281, 50282, 50283, 50284
MAX_OPTIONS = 255
TYPE_TEXT = {"choice": "choice:", "noul": "yes or no:", "score": "rate on the scale:"}
def describe(value: Any) -> str:
return value if isinstance(value, str) else json.dumps(value, ensure_ascii=False)
def option_list(question: dict[str, Any]) -> tuple[list[str], list[str]]:
"""Same keys/descriptions as the server's reference `option_list`."""
kind = question["type"]
criteria = question.get("criteria")
if kind == "choice":
keys = list(criteria)
return keys, [k if v is None else f"{k}: {describe(v)}" for k, v in criteria.items()]
if kind == "noul":
c = criteria or {}
return ["false", "true"], [describe(c.get("false") or "No / false"), describe(c.get("true") or "Yes / true")]
if kind == "score":
return [str(i) for i in range(len(criteria))], [describe(v) for v in criteria]
raise ValueError(f"Unknown question type: {kind}")
@dataclass(frozen=True)
class WindowConfig:
max_len: int = 512
buckets: tuple[int, ...] = (128, 256, 512)
max_slots: int = 64 # option markers per window
q_head: int = 64
q_tail: int = 128
min_context: int = 96 # context tokens reserved per window when options are packed
overlap: float = 0.25
@dataclass
class Window:
ids: list[int]
positions: list[int] = field(default_factory=list) # marker token positions
options: list[int] = field(default_factory=list) # option index per marker
@property
def length(self) -> int:
return len(self.ids)
class Unsupported(ValueError):
pass
class Renderer:
def __init__(self, tokenizer_path: str, config: WindowConfig = WindowConfig()):
self.tok = Tokenizer.from_file(tokenizer_path)
self.tok.no_padding()
self.tok.no_truncation()
self.cfg = config
self._type_ids = {k: self._enc(v) for k, v in TYPE_TEXT.items()}
def _enc(self, text: str) -> list[int]:
return self.tok.encode(text, add_special_tokens=False).ids if text else []
def _enc_batch(self, texts: list[str]) -> list[list[int]]:
return [e.ids for e in self.tok.encode_batch(texts, add_special_tokens=False)] if texts else []
def bucket(self, n: int) -> int:
for b in self.cfg.buckets:
if n <= b:
return b
raise Unsupported(f"window of {n} tokens exceeds {self.cfg.buckets[-1]}")
def windows(self, state: Any, question: dict[str, Any]) -> tuple[list[Window], int]:
"""All windows for one question and its option count."""
cfg = self.cfg
_, descriptions = option_list(question)
n_opt = len(descriptions)
if not 1 <= n_opt <= MAX_OPTIONS:
raise Unsupported(f"{n_opt} options exceeds the declared limit of {MAX_OPTIONS}")
state_text = describe(state) if state not in (None, "", {}) else ""
instr = describe(question.get("instructions") or "Choose the best matching option.")
q_ids, s_ids, *o_ids = self._enc_batch([instr, state_text] + descriptions)
# question: keep head + tail, move the middle into the scanned context
if len(q_ids) > cfg.q_head + cfg.q_tail:
middle = q_ids[cfg.q_head:len(q_ids) - cfg.q_tail]
q_ids = q_ids[:cfg.q_head] + q_ids[len(q_ids) - cfg.q_tail:]
context = middle + (s_ids if not s_ids else [SEP] + s_ids)
else:
context = s_ids
prefix = [CLS] + self._type_ids[question["type"]] + q_ids + [SEP]
fixed = len(prefix) + 2 # SEP after options, SEP at end
opt_budget = cfg.max_len - fixed - (cfg.min_context if context else 0)
if opt_budget < 8:
raise Unsupported("question head/tail leaves no room for options")
# option pieces: (option index, tokens), each piece fits opt_budget with its marker
pieces: list[tuple[int, list[int]]] = []
for i, ids in enumerate(o_ids):
ids = ids or [PAD] # empty description still gets a marker
step = opt_budget - 1
for s in range(0, len(ids), step):
pieces.append((i, ids[s:s + step]))
groups: list[list[tuple[int, list[int]]]] = [[]]
used = 0
for piece in pieces:
cost = 1 + len(piece[1])
if groups[-1] and (used + cost > opt_budget or len(groups[-1]) >= cfg.max_slots):
groups.append([])
used = 0
groups[-1].append(piece)
used += cost
out: list[Window] = []
for group in groups:
body = list(prefix)
positions, options = [], []
for i, ids in group:
positions.append(len(body))
options.append(i)
body.append(MASK)
body.extend(ids)
body.append(SEP)
room = cfg.max_len - len(body) - 1
if not context:
out.append(Window(body + [SEP], positions, options))
continue
if room <= 0:
raise Unsupported("no room for context")
stride = max(1, int(room * (1 - cfg.overlap)))
starts = [0] if len(context) <= room else list(range(0, len(context) - room, stride)) + [len(context) - room]
for s in starts:
out.append(Window(body + context[s:s + room] + [SEP], list(positions), list(options)))
return out, n_opt