File size: 6,340 Bytes
26b3403 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 | """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
|