| """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 |
| q_head: int = 64 |
| q_tail: int = 128 |
| min_context: int = 96 |
| overlap: float = 0.25 |
|
|
|
|
| @dataclass |
| class Window: |
| ids: list[int] |
| positions: list[int] = field(default_factory=list) |
| options: list[int] = field(default_factory=list) |
|
|
| @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) |
| |
| 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 |
| 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") |
| |
| pieces: list[tuple[int, list[int]]] = [] |
| for i, ids in enumerate(o_ids): |
| ids = ids or [PAD] |
| 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 |
|
|