File size: 1,811 Bytes
f930dac
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""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("+<dt>m <text>") 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