Files
Maven/cmd/semantic-router-experiment/slice19_bpe.py
T

81 lines
2.5 KiB
Python

#!/usr/bin/env python3
"""
Slice 19 tokenizers, derived from the development corpus only.
A. CharVocab — codepoint ids over the dev corpus (deterministic order)
B. BpeVocab2048 — byte-level BPE, vocab ~2048, trained on dev corpus only
Records for §4 of the brief: vocab size, OOV behaviour, serialized tokenizer size.
"""
import re
import unicodedata
from tokenizers import Tokenizer
from tokenizers.decoders import ByteLevel as ByteLevelDecoder
from tokenizers.models import BPE
from tokenizers.pre_tokenizers import ByteLevel as ByteLevelPreTokenizer
from tokenizers.trainers import BpeTrainer
class CharVocab:
"""Codepoint ids from the dev corpus, sorted by codepoint value."""
def __init__(self, texts):
chars = set()
for t in texts:
chars.update(t)
self.id_to_char = sorted(chars)
self.char_to_id = {c: i + 1 for i, c in enumerate(self.id_to_char)} # 0 = PAD
self.pad = 0
@property
def size(self):
return len(self.id_to_char) + 1
def encode(self, text, max_len):
ids = [self.char_to_id.get(c, 0) for c in text] # 0 doubles as UNK/OOV
return ids[:max_len]
class BpeVocab:
"""Byte-level BPE, trained only on the strings it is given."""
def __init__(self, texts, vocab_size=2048, sep=""):
self.tok = Tokenizer(BPE())
self.tok.pre_tokenizer = ByteLevelPreTokenizer(trim_offsets=False)
self.tok.decoder = ByteLevelDecoder()
trainer = BpeTrainer(vocab_size=vocab_size, special_tokens=["[PAD]"],
show_progress=False)
# train on the corpus *strings*, byte-level BPE handles all codepoints
self.tok.train_from_iterator(texts, trainer=trainer)
self.pad_id = self.tok.token_to_id("[PAD]")
self._vocab = self.tok.get_vocab()
self._n = len(self._vocab)
@property
def size(self):
return self._n
def encode(self, text):
return self.tok.encode(text).ids
def serialized_bytes(self):
# measure the serialized tokenizer size on disk
import os
d = self.tok.to_str()
return len(d.encode("utf-8"))
def normalize_match_text(s: str) -> str:
"""NFKC → lowercase → collapse whitespace. Punctuation kept."""
out = unicodedata.normalize("NFKC", s).strip().lower()
out = re.sub(r"\s+", " ", out)
return out
def strip_punct(text: str) -> str:
"""Remove safe punctuation from an already-normalized text."""
t = re.sub(r"[^\w\s]", " ", text)
t = re.sub(r"\s+", " ", t).strip()
return t