Download runtime/jevlike/serialize.py from Cem13/kodama-core: direct link, hf CLI and curl.
- Browser
- Download file 15.5 kB
-
https://huggingface.co/Cem13/kodama-core/resolve/main/runtime/jevlike/serialize.py
- Command line
-
hf download hf://Cem13/kodama-core/runtime/jevlike/serialize.py
-
curl -L -o serialize.py https://huggingface.co/Cem13/kodama-core/resolve/main/runtime/jevlike/serialize.py
15.5 kB
| """(state, Question) -> token ids + option-marker positions, plus batching utilities. | |
| Layout of one sequence (one question per sequence):: | |
| [CLS] <state> [SEP] <type tag> <instructions> [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 | |
| 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"] | |
| 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 | |