"""Tokenizer: byte-level BPE, vocab 8192, trained ONLY on train-split text.
Reserved ids: 0=[PAD], 1=[BOS], 2..13=[EVT_k] event-type tokens, 14..31 spare.
Every event is encoded as [EVT_k] + BPE("+
m ") so both PNSR and the
transcript baseline see identical per-event token sequences (dt included).
"""
from __future__ import annotations
from pathlib import Path
from tokenizers import Tokenizer, decoders, models, pre_tokenizers, trainers
from ..common import tokenizer_path
from ..world.schema import ET
N_RESERVED = 32
PAD, BOS = 0, 1
EVT_BASE = 2 # EVT_BASE + int(etype)
VOCAB = 8192
def special_tokens() -> list[str]:
toks = ["[PAD]", "[BOS]"]
toks += [f"[EVT_{ET(i).name}]" for i in range(12)]
toks += [f"[SPARE{i}]" for i in range(N_RESERVED - len(toks) - 0)]
return toks[:N_RESERVED]
def train_tokenizer(texts, out: Path | None = None) -> Tokenizer:
tok = Tokenizer(models.BPE(unk_token=None, byte_fallback=True))
tok.pre_tokenizer = pre_tokenizers.ByteLevel(add_prefix_space=True)
tok.decoder = decoders.ByteLevel()
trainer = trainers.BpeTrainer(
vocab_size=VOCAB, special_tokens=special_tokens(), show_progress=False,
initial_alphabet=pre_tokenizers.ByteLevel.alphabet())
tok.train_from_iterator(texts, trainer)
out = out or tokenizer_path()
out.parent.mkdir(parents=True, exist_ok=True)
tok.save(str(out))
return tok
_CACHED: Tokenizer | None = None
def load_tokenizer() -> Tokenizer:
global _CACHED
if _CACHED is None:
_CACHED = Tokenizer.from_file(str(tokenizer_path()))
return _CACHED
def event_text(ev) -> str:
return f"+{ev.dt}m {ev.text}"
def encode_event(tok: Tokenizer, ev) -> list[int]:
ids = tok.encode(event_text(ev)).ids
return [EVT_BASE + int(ev.etype)] + ids