kodama-core / runtime /jevlike /serialize.py
Cem13's picture
Release Kodama Core with weights, runtime, attribution, and evaluation
c7893fa verified
Raw History Blame Contribute Delete
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
@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