"""(state, Question) -> token ids + option-marker positions, plus batching utilities. Layout of one sequence (one question per sequence):: [CLS] [SEP] [SEP] [MASK] key: text [MASK] key: text ... [SEP] Every option gets its own [MASK] marker; noul is rendered as ``[MASK] false: [MASK] true:``. Literal special-token strings inside user text (e.g. "[MASK]") are tokenized as plain text. Length budget (``max_len`` total, ``head_max_len`` reserved for the question part): * The question part is guaranteed up to ``head_max_len`` tokens, and may borrow whatever the state does not need (short state -> long question is fine). The state is truncated first: it gets ``max_len - 2 - len(question part)`` tokens, keeping its head (25%) and tail (75%). * If the question part is longer than its budget it is compressed in stages, never dropping a marker: (1) instructions capped at ``instr_max_len``; (2) option texts water-filled (every text truncated to a common cap, short texts untouched); (3) keys only; if the keys alone still do not fit, the question borrows from the state down to ``min_state_len`` tokens and the keys are water-filled, but never below ``key_min_len`` tokens (shorter keys stop being distinguishable). (4) If even that does not fit (hundreds of options), the sequence is allowed to exceed ``max_len`` (ModernBERT handles 8192 positions; 255 options with short keys give ~1.1k tokens). ``Encoded.trunc`` records the most aggressive stage reached (T_* codes). The middle cut works at any ``max_len`` (up to ModernBERT's 8192 positions). For states longer than the state budget, ``encode_chunked`` instead splits the state into overlapping windows that each fit (the question part is identical in every window); ``SystemOne.predict(..., long_state="chunk")`` scores every window and pools the answers. """ from __future__ import annotations import math from dataclasses import dataclass from typing import Iterator, Sequence import numpy as np import torch from torch.utils.data import Sampler from jevlike.types import QTYPES, Question, Record TYPE_TAGS = {"choice": "Choose one:", "score": "Rate on the scale:", "noul": "True or false:"} TYPE_IDS = {t: i for i, t in enumerate(QTYPES)} NOUL_KEYS = ("false", "true") PAD_MULTIPLE = 8 # Encoded.trunc codes T_NONE, T_STATE, T_INSTR, T_TEXTS, T_KEYS, T_OVERFLOW = range(6) # T_STATE < question stages @dataclass class Encoded: input_ids: list[int] markers: list[int] # position of each option's [MASK], in option order qtype: int # index into QTYPES trunc: int = T_NONE # most aggressive truncation stage applied (T_* codes) def __len__(self) -> int: return len(self.input_ids) def _fit(seqs: list[list[int]], budget: int, floor: int = 0) -> list[list[int]]: """Water-fill: truncate every sequence to a common cap (short ones untouched) so the total fits ``budget``; leftover tokens go one each to the first truncated sequences. Never cuts below ``floor`` tokens (so the result may exceed the budget).""" lengths = [len(x) for x in seqs] if sum(lengths) <= budget: return seqs cap, used, rest = 0, 0, len(lengths) for l in sorted(lengths): if used + l * rest > budget: cap = max(0, (budget - used) // rest) break used, rest = used + l, rest - 1 if cap < floor: return [x[:floor] for x in seqs] spare = budget - sum(min(l, cap) for l in lengths) out = [] for x in seqs: extra = 1 if len(x) > cap and spare > 0 else 0 spare -= extra out.append(x[:cap + extra]) return out class Serializer: """Tokenizer wrapper implementing the input layout above.""" def __init__(self, tokenizer, max_len: int = 512, head_max_len: int = 192, instr_max_len: int = 96, min_state_len: int = 64, key_min_len: int = 4, state_head_frac: float = 0.25): assert head_max_len < max_len self.tok = tokenizer self.max_len, self.head_max_len = max_len, head_max_len self.instr_max_len, self.min_state_len, self.key_min_len = instr_max_len, min_state_len, key_min_len self.state_head_frac = state_head_frac self.cls, self.sep, self.mask, self.pad = (tokenizer.cls_token_id, tokenizer.sep_token_id, tokenizer.mask_token_id, tokenizer.pad_token_id) self.ellipsis = self._tokenize([" ..."])[0] def _tokenize(self, texts: list[str]) -> list[list[int]]: if not texts: return [] return self.tok(texts, add_special_tokens=False, split_special_tokens=True)["input_ids"] @staticmethod def _pieces(q: Question) -> tuple[str, list[str], list[str]]: prefix = f"{TYPE_TAGS[q.type]} {q.instructions.strip()}" if q.type == "noul": return prefix, [f" {k}:" for k in NOUL_KEYS], ["", ""] return prefix, [f" {o.key.strip()}:" for o in q.options], [f" {o.text.strip()}" if o.text.strip() else "" for o in q.options] def encode(self, state: str, q: Question) -> Encoded: return self.encode_many([(state, q)])[0] def _prepare(self, pairs: Sequence[tuple[str, Question]]) -> list[tuple]: """Tokenize all pairs with one batched tokenizer call over the unique strings.""" pieces = [(s, *self._pieces(q)) for s, q in pairs] uniq: dict[str, int] = {} for s, prefix, keys, texts in pieces: for t in (s, prefix, *keys, *texts): uniq.setdefault(t, len(uniq)) ids = self._tokenize(list(uniq)) get = lambda t: ids[uniq[t]] if t else [] # noqa: E731 return [(get(s), get(prefix), [get(k) for k in keys], [get(t) for t in texts], TYPE_IDS[q.type]) for (s, prefix, keys, texts), (_, q) in zip(pieces, pairs)] def encode_many(self, pairs: Sequence[tuple[str, Question]]) -> list[Encoded]: """Encode many pairs (long states cut in the middle).""" return [self._assemble(*p) for p in self._prepare(pairs)] def encode_chunked(self, pairs: Sequence[tuple[str, Question]], overlap: int = 128, max_chunks: int = 64) -> list[list[Encoded]]: """Like ``encode_many`` but a state longer than its budget is split into overlapping windows instead of being cut: one ``Encoded`` per window, all with the same question part. Pairs whose state fits give a single-element list identical to ``encode_many``. Windows advance by ``window - overlap`` tokens (overlap capped at a quarter of the window); the last window is aligned to the end of the state. Windows that start/end inside the state carry a " ..." marker at the cut, like the middle cut. If more than ``max_chunks`` windows would be needed, ``max_chunks`` windows are spread evenly over the state (first and last kept, gaps in between); ``max_chunks <= 1`` falls back to the middle cut. """ return [self._assemble_chunks(*p, overlap=overlap, max_chunks=max_chunks) for p in self._prepare(pairs)] def _fit_question(self, n_state: int, prefix: list[int], keys: list[list[int]], texts: list[list[int]]): """Compress the question part for a state of ``n_state`` tokens. Returns (prefix, keys, texts, trunc, state_room).""" n = len(keys) budget = max(self.head_max_len, self.max_len - 2 - n_state) # question-part budget room = budget - 2 - n # tokens for prefix + keys + texts trunc = T_NONE def total(p, ks, ts): return len(p) + sum(map(len, ks)) + sum(map(len, ts)) if total(prefix, keys, texts) > room: trunc, prefix = T_INSTR, prefix[:self.instr_max_len] if total(prefix, keys, texts) > room: trunc = T_TEXTS texts = _fit(texts, room - len(prefix) - sum(map(len, keys))) if not any(texts): trunc = T_KEYS if len(prefix) + sum(map(len, keys)) > room: # keys alone do not fit: borrow from the state room = max(room, self.max_len - 4 - n - min(n_state, self.min_state_len)) keys = _fit(keys, room - len(prefix), floor=self.key_min_len) q_len = 2 + n + total(prefix, keys, texts) state_room = max(self.max_len - 2 - q_len, min(self.min_state_len, n_state)) return prefix, keys, texts, trunc, state_room def _build(self, state: list[int], prefix: list[int], keys: list[list[int]], texts: list[list[int]], qtype: int, trunc: int) -> Encoded: ids = [self.cls, *state, self.sep, *prefix, self.sep] markers = [] for k, t in zip(keys, texts): markers.append(len(ids)) ids += [self.mask, *k, *t] ids.append(self.sep) assert len(ids) == 4 + len(keys) + len(state) + len(prefix) + sum(map(len, keys)) + sum(map(len, texts)) if len(ids) > self.max_len: trunc = T_OVERFLOW return Encoded(ids, markers, qtype, trunc) def _assemble(self, state: list[int], prefix: list[int], keys: list[list[int]], texts: list[list[int]], qtype: int) -> Encoded: prefix, keys, texts, trunc, state_room = self._fit_question(len(state), prefix, keys, texts) if len(state) > state_room: trunc = max(trunc, T_STATE) head = int(state_room * self.state_head_frac) tail = state_room - head - len(self.ellipsis) state = state[:head] + self.ellipsis + state[len(state) - tail:] if tail > 0 else state[:state_room] return self._build(state, prefix, keys, texts, qtype, trunc) def _assemble_chunks(self, state: list[int], prefix: list[int], keys: list[list[int]], texts: list[list[int]], qtype: int, overlap: int, max_chunks: int) -> list[Encoded]: fitted = self._fit_question(len(state), prefix, keys, texts) prefix_, keys_, texts_, trunc, state_room = fitted if len(state) <= state_room or max_chunks <= 1: return [self._assemble(state, prefix, keys, texts, qtype)] ell = self.ellipsis return [self._build((ell if a > 0 else []) + state[a:b] + (ell if b < len(state) else []), prefix_, keys_, texts_, qtype, trunc) for a, b in chunk_windows(len(state), max(1, state_room - 2 * len(ell)), overlap, max_chunks)] def chunk_windows(n: int, window: int, overlap: int, max_chunks: int) -> list[tuple[int, int]]: """[start, end) token windows of size ``window`` covering ``n`` tokens (see encode_chunked).""" if n <= window: return [(0, n)] stride = window - min(max(0, overlap), window // 4) k = math.ceil((n - window) / stride) + 1 if k <= max_chunks: starts = [i * stride for i in range(k - 1)] + [n - window] else: starts = np.linspace(0, n - window, max_chunks).round().astype(int).tolist() return [(a, a + window) for a in starts] # ---------------------------------------------------------------- batching def collate(encs: Sequence[Encoded], labels: Sequence[int] | None, pad_id: int) -> dict[str, torch.Tensor]: """Dynamically pad a list of encoded questions (one per row). Returns input_ids/attention_mask [B, T]; flat marker index ``marker_idx`` [M] into the flattened [B*T] hidden states, ``marker_seg`` [M] (question id = row) and ``marker_rank`` [M] (option index within its question); ``n_options``, ``qtype`` and ``labels`` [B]. """ B = len(encs) T = math.ceil(max(len(e) for e in encs) / PAD_MULTIPLE) * PAD_MULTIPLE ids = np.full((B, T), pad_id, dtype=np.int64) att = np.zeros((B, T), dtype=np.int64) for i, e in enumerate(encs): ids[i, :len(e)] = e.input_ids att[i, :len(e)] = 1 n_opt = [len(e.markers) for e in encs] seg = np.repeat(np.arange(B), n_opt) pos = np.concatenate([e.markers for e in encs]) out = { "input_ids": torch.from_numpy(ids), "attention_mask": torch.from_numpy(att), "marker_idx": torch.from_numpy(seg * T + pos), "marker_seg": torch.from_numpy(seg), "marker_rank": torch.from_numpy(np.concatenate([np.arange(k) for k in n_opt])), "n_options": torch.tensor(n_opt), "qtype": torch.tensor([e.qtype for e in encs]), } if labels is not None: out["labels"] = torch.tensor(list(labels)) return out class EncodedDataset(torch.utils.data.Dataset): """Pre-encoded records (tokenization happens once, up front).""" def __init__(self, records: Sequence[Record], serializer: Serializer, chunk: int = 20000): self.records = list(records) self.encs: list[Encoded] = [] for i in range(0, len(self.records), chunk): self.encs += serializer.encode_many([(r.state, r.question) for r in self.records[i:i + chunk]]) self.labels = [r.label for r in self.records] self.lengths = np.array([len(e) for e in self.encs]) self.pad_id = serializer.pad def __len__(self) -> int: return len(self.encs) def __getitem__(self, i: int) -> int: return i def collate(self, idx: Sequence[int]) -> dict[str, torch.Tensor]: b = collate([self.encs[i] for i in idx], [self.labels[i] for i in idx], self.pad_id) b["index"] = torch.tensor(list(idx)) return b class TokenBudgetBatchSampler(Sampler[list[int]]): """Length-bucketed batches bounded by padded tokens (B * T_max <= max_tokens) and max_batch. Shuffles indices, sorts them by length inside pools of ~``pool`` batches, cuts batches, then shuffles batch order. Deterministic per (seed, epoch); ``skip`` resumes mid-epoch. """ def __init__(self, lengths: Sequence[int], max_tokens: int = 8192, max_batch: int = 64, shuffle: bool = True, seed: int = 0, pool: int = 50): self.lengths = np.asarray(lengths) self.max_tokens, self.max_batch, self.shuffle, self.seed, self.pool = max_tokens, max_batch, shuffle, seed, pool self.epoch, self.skip = 0, 0 def set_epoch(self, epoch: int, skip: int = 0) -> None: self.epoch, self.skip = epoch, skip def batches(self) -> list[list[int]]: rng = np.random.default_rng(self.seed + 1000 * self.epoch) n = len(self.lengths) idx = rng.permutation(n) if self.shuffle else np.arange(n) padded = np.ceil(self.lengths / PAD_MULTIPLE).astype(int) * PAD_MULTIPLE pool_size = self.max_batch * self.pool if self.shuffle else n out: list[list[int]] = [] for p in range(0, n, pool_size): chunk = idx[p:p + pool_size] chunk = chunk[np.argsort(-padded[chunk], kind="stable")] cur: list[int] = [] for i in chunk.tolist(): t_max = padded[cur[0]] if cur else padded[i] # sorted descending: first is longest if cur and ((len(cur) + 1) * t_max > self.max_tokens or len(cur) >= self.max_batch): out.append(cur) cur = [] cur.append(i) if cur: out.append(cur) if self.shuffle: rng.shuffle(out) return out def __iter__(self) -> Iterator[list[int]]: yield from self.batches()[self.skip:] def __len__(self) -> int: return len(self.batches()) - self.skip