BiGRU_T_version / src /bigru_t /tokenizer /bbpe_tokenizer.py
PowerMachine's picture
V6.7: upload src/bigru_t/tokenizer/bbpe_tokenizer.py (BBPE serial mode + OomGuard V7)
13cc5a0 verified
Raw History Blame Contribute Delete
39.3 kB
"""bbpe_tokenizer.py — BBPE (Byte-Level BPE) Tokenizer PARALELO (Map-Reduce).
═══════════════════════════════════════════════════════════════════════════════
REFATORAÇÃO: Algoritmo Paralelo Map-Reduce (substitui o treinamento sequencial)
═══════════════════════════════════════════════════════════════════════════════
O treinamento do BBPE foi refatorado para um modelo Map-Reduce paralelo,
substituindo o BpeTrainer sequencial da HuggingFace. O algoritmo:
ETAPA 0: Inicialização
- distribute_texts(): converte o iterador em shards balanceados
- pre_tokenize_shard(): pré-tokeniza byte-level cada shard
- Estruturas globais: vocab, token_to_id, merges
LAÇO PRINCIPAL DE MERGES:
FASE 1 — MAP: count_pairs_in_shard() em paralelo (ProcessPoolExecutor)
FASE 2 — REDUCE: agrega contagens locais em global_counts
FASE 3 — CHOICE: escolhe o par de maior frequência (desempate lexicográfico)
FASE 4 — APPLY: apply_merge_in_shard() em paralelo
FASE 5 — UPDATE: atualiza vocab, token_to_id, merges
ETAPA FINAL: build_bpe_from_merges() constrói o tokenizer HF a partir
dos merges + vocab calculados.
INFERÊNCIA PARALELA: encode_batch_parallel() usa ThreadPoolExecutor.
═══════════════════════════════════════════════════════════════════════════════
TEOREMA 20 (BBPE Universal Coverage) — mantido
═══════════════════════════════════════════════════════════════════════════════
BBPE opera no espaço de BYTES UTF-8 (256 símbolos base). Garante:
1. COBERTURA UNIVERSAL: qualquer string UTF-8 é tokenizável sem <unk>.
2. COMPATIBILIDADE MULTILINGUE: mesmo vocab para PT-BR, Xavante, EN, código.
3. COMPRESSÃO ÓTIMA: BPE greedy + byte-level = merges no espaço total.
4. ESCALABILIDADE: V = 16K-250K tokens cobre eficientemente múltiplos idiomas.
INTEGRIDADE (Semântica Preservada):
dec(enc(s)) = s para todo s ∈ Σ_UTF-8*.
Decorre de ByteLevel ser bijeção reversível no espaço de bytes.
═══════════════════════════════════════════════════════════════════════════════
COMPATIBILIDADE
═══════════════════════════════════════════════════════════════════════════════
Mantém a API pública da versão anterior:
- encode(text) -> List[int]
- decode(ids) -> str
- encode_batch(texts) -> List[List[int]]
- encode_batch_parallel(texts) -> List[List[int]] [NOVO]
- save(path) / load(path)
- train_from_stream(iter, path) [delega para train_parallel_from_stream]
- train_parallel_from_stream(iter, ...) [NOVO — algoritmo Map-Reduce]
- train_from_files(files, path)
- encode_tensor(texts, max_length)
- validate_roundtrip(test_texts)
"""
from __future__ import annotations
import os
import json
import logging
import gc
from collections import defaultdict
from concurrent.futures import ProcessPoolExecutor, ThreadPoolExecutor, as_completed
from pathlib import Path
from typing import (
Iterator, List, Optional, Dict, Any, Union, Tuple, Iterable,
)
logger = logging.getLogger(__name__)
# === Special tokens (alinhados com HuggingFace conventions) ===
BOS_TOKEN = "<s>"
PAD_TOKEN = "<pad>"
EOS_TOKEN = "</s>"
UNK_TOKEN = "<unk>"
SPECIAL_TOKENS = [BOS_TOKEN, PAD_TOKEN, EOS_TOKEN, UNK_TOKEN]
# IDs canônicos (precedência sobre merges)
BOS_ID = 0
PAD_ID = 1
EOS_ID = 2
UNK_ID = 3
# Vocabulário padrão (reduzido de 250K para 16K — adequado ao escopo BiGRU_T)
DEFAULT_VOCAB_SIZE = 16_384
# ---------------------------------------------------------------------------
# Mapeamento Byte-Level (GPT-2 style: 0-255 -> Ġ-style unicode strings)
# ---------------------------------------------------------------------------
def bytes_to_unicode() -> Dict[int, str]:
"""Retorna mapeamento byte (0-255) -> símbolo unicode (Ġ-style do GPT-2).
Bytes correspondentes a caracteres printable (33-126, 161-172, 174-255)
mapeiam para si mesmos. Os demais (0-32, 127-160, 173) mapeiam para
codepoints a partir de 256 (Ġ=256+32=288 → 'Ġ', etc.).
"""
bs = (
list(range(ord("!"), ord("~") + 1))
+ list(range(ord("¡"), ord("¬") + 1))
+ list(range(ord("®"), ord("ÿ") + 1))
)
cs = bs[:]
n = 0
for b in range(256):
if b not in bs:
bs.append(b)
cs.append(256 + n)
n += 1
cs = [chr(c) for c in cs]
return dict(zip(bs, cs))
# Tabelas globais (construídas uma vez no import)
BYTE_TO_SYMBOL: Dict[int, str] = bytes_to_unicode()
SYMBOL_TO_BYTE: Dict[str, int] = {v: k for k, v in BYTE_TO_SYMBOL.items()}
# Alfabeto byte-level (256 símbolos)
ALPHABET: List[str] = [BYTE_TO_SYMBOL[i] for i in range(256)]
# ---------------------------------------------------------------------------
# FUNÇÕES AUXILIARES DO ALGORITMO PARALELO (top-level para picklability)
# ---------------------------------------------------------------------------
def distribute_texts(
text_iterator: Iterable[str],
num_workers: int,
chunk_size: int = 500,
) -> List[List[str]]:
"""Converte um iterador de textos em partições balanceadas (shards).
Cada shard é uma lista de strings que será processada por um worker.
A distribuição é round-robin sobre chunks para balancear carga.
Args:
text_iterator: iterador yielding strings de texto
num_workers: número de shards a produzir
chunk_size: textos acumulados antes de formar um chunk
Returns:
Lista de `num_workers` shards (cada shard é List[str]).
"""
if num_workers < 1:
num_workers = 1
shards: List[List[str]] = [[] for _ in range(num_workers)]
current_chunk: List[str] = []
chunk_idx = 0
total = 0
for text in text_iterator:
if not text or len(str(text).strip()) < 10:
continue
current_chunk.append(str(text))
total += 1
if len(current_chunk) >= chunk_size:
# Round-robin assignment
shards[chunk_idx % num_workers].extend(current_chunk)
current_chunk = []
chunk_idx += 1
# Flush final
if current_chunk:
shards[chunk_idx % num_workers].extend(current_chunk)
# Remove shards vazios (pode acontecer se num_workers > chunks)
shards = [s for s in shards if s]
if not shards:
raise RuntimeError(
"distribute_texts: nenhum texto válido encontrado no iterador"
)
logger.info(
"distribute_texts: %d textos em %d shards (chunk_size=%d)",
total, len(shards), chunk_size,
)
return shards
def pre_tokenize_shard(texts: List[str]) -> List[List[str]]:
"""Converte uma lista de textos em uma lista de listas de símbolos byte-level.
Cada símbolo é uma string representando um byte (Ġ-style).
Tokens especiais (se presentes no texto como substrings) são tratados
como símbolos únicos — mas nesta implementação simples, expandimos tudo
para bytes (a detecção de especiais é feita no encode, não no treino).
Args:
texts: lista de strings (um shard)
Returns:
Lista de listas de símbolos (uma lista por documento).
"""
shard_syms: List[List[str]] = []
for text in texts:
doc_syms: List[str] = []
# Pré-tokenização ByteLevel: cada caractere -> UTF-8 bytes -> símbolos
for ch in text:
utf8_bytes = ch.encode("utf-8")
for b in utf8_bytes:
doc_syms.append(BYTE_TO_SYMBOL[b])
shard_syms.append(doc_syms)
return shard_syms
def count_pairs_in_shard(
shard: List[List[str]],
min_freq: int = 2,
) -> Dict[Tuple[str, str], int]:
"""Conta pares adjacentes dentro de cada documento, sem cruzar fronteiras.
Args:
shard: lista de documentos (cada doc é lista de símbolos)
min_freq: frequência mínima para manter o par (poda local)
Returns:
Dicionário {(esq, dir): contagem_local}.
"""
counts: Dict[Tuple[str, str], int] = defaultdict(int)
for doc in shard:
if len(doc) < 2:
continue
for i in range(len(doc) - 1):
pair = (doc[i], doc[i + 1])
counts[pair] += 1
# Poda local: descarta pares com contagem < min_freq
if min_freq > 1:
return {p: c for p, c in counts.items() if c >= min_freq}
return dict(counts)
def apply_merge_in_shard(
shard: List[List[str]],
left: str,
right: str,
replacement: str,
) -> List[List[str]]:
"""Substitui toda ocorrência adjacente de (left, right) por `replacement`.
Args:
shard: lista de documentos (cada doc é lista de símbolos)
left: símbolo esquerdo do merge
right: símbolo direito do merge
replacement: novo símbolo que substitui o par
Returns:
Novo shard com o merge aplicado em todos os documentos.
"""
new_shard: List[List[str]] = []
for doc in shard:
if len(doc) < 2:
new_shard.append(doc)
continue
new_doc: List[str] = []
i = 0
n = len(doc)
while i < n:
if i < n - 1 and doc[i] == left and doc[i + 1] == right:
new_doc.append(replacement)
i += 2
else:
new_doc.append(doc[i])
i += 1
new_shard.append(new_doc)
return new_shard
def build_bpe_from_merges(
merges: List[Tuple[str, str, str]],
token_to_id: Dict[str, int],
unk_token: str = UNK_TOKEN,
add_prefix_space: bool = False,
):
"""Constrói um tokenizers.Tokenizer a partir dos merges e vocab calculados.
Converte o formato interno (lista de tuplas (left, right, new)) para o
formato esperado pelo tokenizers.models.BPE (lista de strings "left right").
Args:
merges: lista de (esq, dir, novo_token)
token_to_id: mapeamento token -> id
unk_token: token de desconhecido
add_prefix_space: se True, adiciona espaço prefixo no pre-tokenizer
Returns:
tokenizers.Tokenizer configurado com BPE + ByteLevel.
"""
from tokenizers import Tokenizer
from tokenizers.models import BPE
from tokenizers.pre_tokenizers import ByteLevel
from tokenizers.processors import ByteLevel as ByteLevelProcessor
from tokenizers.decoders import ByteLevel as ByteLevelDecoder
# Converte merges para formato HF: lista de tuplas (left, right)
hf_merges = [(left, right) for (left, right, _new) in merges]
# BPE model com vocab e merges
bpe = BPE(
vocab=token_to_id,
merges=hf_merges,
unk_token=unk_token,
)
tok = Tokenizer(bpe)
tok.pre_tokenizer = ByteLevel(add_prefix_space=add_prefix_space)
tok.post_processor = ByteLevelProcessor(trim_offsets=False)
# CRÍTICO: ByteLevel decoder inverte o mapeamento byte-level de volta para UTF-8.
tok.decoder = ByteLevelDecoder()
return tok
# ---------------------------------------------------------------------------
# Classe principal
# ---------------------------------------------------------------------------# ============================================================================
# V6.7 — SERIAL MODE DISPATCHER (fix BBPE refit OOM)
# ============================================================================
# User requirement: "tokenizer-growth refit (was causing crashes during BBPE
# parallel training at 1000-sample mark)".
#
# Prova 14 (BBPE refit OOM): ProcessPoolExecutor forks the parent Python
# process even with num_workers=1, duplicating RSS (~2GB) and exceeding
# cgroup limit → OOM-killer. Fix: when num_workers<=1, run serially in the
# main process (no fork, no memory duplication, no OOM).
#
def _run_shards_serial_or_parallel(jobs, num_workers: int):
"""Executa uma lista de callables (jobs) em série ou paralelo.
- Se num_workers <= 1: executa serialmente no processo atual (NO FORK).
Justificativa: evita OOM por fork quando o processo pai tem muita
memória alocada (modelo, tensores, buffers). Amdahl: para 1 worker,
paralelismo não traz ganho, só overhead.
- Se num_workers >= 2: usa ProcessPoolExecutor (paralelismo real).
Args:
jobs: lista de callables (sem args) — use functools.partial ou lambda.
num_workers: número de processos paralelos (<=1 = serial).
Returns:
Lista de resultados na mesma ordem dos jobs.
"""
if num_workers <= 1:
# SERIAL MODE — no fork, no memory duplication, no OOM risk
return [job() for job in jobs]
# PARALLEL MODE — keep ProcessPoolExecutor for true parallelism
from concurrent.futures import ProcessPoolExecutor
with ProcessPoolExecutor(max_workers=num_workers) as executor:
futures = [executor.submit(job) for job in jobs]
return [f.result() for f in futures]
class BBPETokenizer:
"""BBPE Tokenizer com treinamento paralelo Map-Reduce.
API pública (compatível com versão anterior):
- encode(text) -> List[int]
- decode(ids) -> str
- encode_batch(texts) -> List[List[int]]
- encode_batch_parallel(texts) -> List[List[int]] [NOVO]
- save(path) / load(path)
- train_parallel_from_stream(iter, ...) [NOVO — algoritmo Map-Reduce]
- train_from_stream(iter, ...) [delega para paralelo]
- train_from_files(files, ...)
- encode_tensor(texts, max_length)
- validate_roundtrip(test_texts)
"""
def __init__(
self,
vocab_size: int = DEFAULT_VOCAB_SIZE,
bos_token: str = BOS_TOKEN,
pad_token: str = PAD_TOKEN,
eos_token: str = EOS_TOKEN,
unk_token: str = UNK_TOKEN,
add_prefix_space: bool = False,
num_workers: int = 4,
):
self.vocab_size = vocab_size
self.bos_token = bos_token
self.pad_token = pad_token
self.eos_token = eos_token
self.unk_token = unk_token
self.add_prefix_space = add_prefix_space
self.num_workers = max(1, num_workers)
self._tokenizer = None # Lazy init
self._vocab: Optional[Dict[str, int]] = None
self._id_to_token: Optional[Dict[int, str]] = None
# Merges aprendidos (para inspeção / re-build)
self._merges: List[Tuple[str, str, str]] = []
# ------------------------------------------------------------------
# TREINAMENTO PARALELO (Map-Reduce) — NOVO
# ------------------------------------------------------------------
def train_parallel_from_stream(
self,
text_iterator: Iterator[str],
save_path: Optional[Path] = None,
min_frequency: int = 2,
num_workers: int = 4,
show_progress: bool = True,
chunk_size: int = 500,
) -> None:
"""Treina o tokenizer BBPE em paralelo (Map-Reduce com ProcessPoolExecutor).
Algoritmo:
ETAPA 0: distribute_texts + pre_tokenize_shard (inicialização)
LAÇO:
FASE 1: MAP — count_pairs_in_shard em paralelo
FASE 2: REDUCE — agrega contagens
FASE 3: CHOICE — melhor par (freq máx, desempate lexicográfico)
FASE 4: APPLY — apply_merge_in_shard em paralelo
FASE 5: UPDATE — atualiza vocab + merges
ETAPA FINAL: build_bpe_from_merges
Args:
text_iterator: iterador yielding strings de texto
save_path: caminho para salvar o tokenizer JSON
min_frequency: frequência mínima de um par para ser mergeado
num_workers: número de processos paralelos
show_progress: exibir progresso por iteração
chunk_size: textos por chunk na distribuição
"""
logger.info(
"BBPE PARALLEL train: vocab_size=%d, min_freq=%d, workers=%d",
self.vocab_size, min_frequency, num_workers,
)
# --- ETAPA 0: Inicialização ---
# 0.1 Distribui textos em shards balanceados
shards = distribute_texts(text_iterator, num_workers, chunk_size)
# 0.2 Pré-tokeniza cada shard em sequências de símbolos byte-level
# (paralelo, pois pre_tokenize é CPU-bound)
shard_symbols = _run_shards_serial_or_parallel(
[lambda s=shard: pre_tokenize_shard(s) for shard in shards],
num_workers,
)
logger.info(
"BBPE PARALLEL: pré-tokenização concluída (%d shards)",
len(shard_symbols),
)
# 0.3 Estruturas globais
current_vocab: set = set(ALPHABET + SPECIAL_TOKENS)
token_to_id: Dict[str, int] = {}
# IDs canônicos para especiais primeiro
for i, tok in enumerate(SPECIAL_TOKENS):
token_to_id[tok] = i
# Depois o alfabeto byte-level
next_id = len(SPECIAL_TOKENS)
for sym in ALPHABET:
if sym not in token_to_id:
token_to_id[sym] = next_id
next_id += 1
merges: List[Tuple[str, str, str]] = []
# --- LAÇO PRINCIPAL DE MERGES ---
iteration = 0
target_merges = self.vocab_size - len(token_to_id)
while len(current_vocab) < self.vocab_size:
iteration += 1
# --- FASE 1: MAP (contagem local de pares) ---
local_counts = _run_shards_serial_or_parallel(
[lambda s=sym_shard: count_pairs_in_shard(s, min_frequency)
for sym_shard in shard_symbols],
num_workers,
)
# --- FASE 2: REDUCE (agregação) ---
global_counts: Dict[Tuple[str, str], int] = defaultdict(int)
for lc in local_counts:
for pair, cnt in lc.items():
global_counts[pair] += cnt
# Poda global (garante min_frequency)
global_counts = {
p: c for p, c in global_counts.items() if c >= min_frequency
}
if not global_counts:
logger.info(
"BBPE PARALLEL: nenhum par com freq >= %d restante. "
"Vocab final: %d (target %d)",
min_frequency, len(current_vocab), self.vocab_size,
)
break
# --- FASE 3: ESCOLHA DO MELHOR PAR ---
# (frequência máxima, desempate lexicográfico)
best_pair = max(
global_counts.items(), key=lambda x: (x[1], x[0])
)
(esq, dir_), freq = best_pair
# Gera novo token (concatenação byte-level)
new_token_str = esq + dir_
# Se já existe (raro), gera nome único
while new_token_str in current_vocab:
new_token_str += "_"
new_id = next_id
next_id += 1
# --- FASE 4: APPLY (aplicação do merge nos shards) ---
shard_symbols = _run_shards_serial_or_parallel(
[lambda s=sym_shard: apply_merge_in_shard(s, esq, dir_, new_token_str)
for sym_shard in shard_symbols],
num_workers,
)
# --- FASE 5: ATUALIZAÇÃO DO VOCABULÁRIO ---
current_vocab.add(new_token_str)
token_to_id[new_token_str] = new_id
merges.append((esq, dir_, new_token_str))
if show_progress and (
iteration <= 20
or iteration % 50 == 0
or len(current_vocab) >= self.vocab_size - 5
):
logger.info(
"BBPE PARALLEL iter %d: '%s' + '%s' -> '%s' "
"(freq=%d) | Vocab=%d/%d",
iteration, esq, dir_, new_token_str, freq,
len(current_vocab), self.vocab_size,
)
# Libera memória periodicamente (estilo Xavante)
if iteration % 100 == 0:
gc.collect()
# --- ETAPA FINAL: Construção do tokenizer interno ---
logger.info(
"BBPE PARALLEL: concluído. %d merges, vocab=%d. Construindo tokenizer HF...",
len(merges), len(token_to_id),
)
self._merges = merges
self._tokenizer = build_bpe_from_merges(
merges=merges,
token_to_id=token_to_id,
unk_token=self.unk_token,
add_prefix_space=self.add_prefix_space,
)
self._build_vocab_cache()
# Atualiza vocab_size com tamanho real
self.vocab_size = len(self._vocab)
logger.info(
"BBPE PARALLEL: tokenizer construído. Vocab real: %d",
len(self._vocab),
)
if save_path is not None:
self.save(save_path)
# ------------------------------------------------------------------
# TREINAMENTO (compatibilidade — delega para paralelo)
# ------------------------------------------------------------------
def train_from_stream(
self,
text_iterator: Iterator[str],
save_path: Optional[Union[str, Path]] = None,
min_frequency: int = 2,
show_progress: bool = True,
chunk_size: int = 500,
num_workers: Optional[int] = None,
) -> None:
"""Treina o tokenizer BBPE a partir de um iterador de textos.
REFACTORED: agora delega para train_parallel_from_stream (Map-Reduce).
Mantém a assinatura para compatibilidade com código existente.
Args:
text_iterator: iterador yielding strings de texto
save_path: caminho para salvar o tokenizer JSON
min_frequency: frequência mínima de um par para ser mergeado
show_progress: exibir progresso
chunk_size: textos por chunk na distribuição
num_workers: número de processos paralelos (default: self.num_workers)
"""
workers = num_workers if num_workers is not None else self.num_workers
self.train_parallel_from_stream(
text_iterator=text_iterator,
save_path=Path(save_path) if save_path else None,
min_frequency=min_frequency,
num_workers=workers,
show_progress=show_progress,
chunk_size=chunk_size,
)
def train_from_files(
self,
file_paths: List[Union[str, Path]],
save_path: Optional[Union[str, Path]] = None,
min_frequency: int = 2,
show_progress: bool = True,
num_workers: Optional[int] = None,
) -> None:
"""Treina o tokenizer BBPE a partir de arquivos de texto.
Lê os arquivos e cria um iterador de linhas, delegando para
train_parallel_from_stream.
"""
def _file_line_iterator(paths):
for p in paths:
p = Path(p)
if not p.exists():
logger.warning("Arquivo não encontrado: %s", p)
continue
with open(p, "r", encoding="utf-8", errors="replace") as f:
for line in f:
line = line.strip()
if line:
yield line
workers = num_workers if num_workers is not None else self.num_workers
self.train_parallel_from_stream(
text_iterator=_file_line_iterator(file_paths),
save_path=Path(save_path) if save_path else None,
min_frequency=min_frequency,
num_workers=workers,
show_progress=show_progress,
chunk_size=500,
)
# ------------------------------------------------------------------
# Save / Load
# ------------------------------------------------------------------
def fit(
self,
texts: List[str],
vocab_size: Optional[int] = None,
min_frequency: int = 2,
) -> None:
"""V6.6 — Compatibilidade com a API SimpleBBPETokenizer.fit().
Treina o tokenizer BBPE a partir de uma lista de textos. Esta método
existe para que o KohonenLearningSystem possa chamar
`kls.tokenizer.fit(corpus_inicial)` sem saber se o tokenizer é
SimpleBBPETokenizer (word-level) ou BBPETokenizer (byte-level).
Args:
texts: lista de strings para treinar o tokenizer.
vocab_size: tamanho do vocabulário alvo (default: self.vocab_size).
min_frequency: frequência mínima de um par para ser mergeado.
"""
if vocab_size is not None:
self.vocab_size = int(vocab_size)
# Limita o corpus para velocidade (primeiros 1000 textos)
# — o BBPE byte-level não precisa de corpus massivo para treinar
# merges úteis; 1000 textos já dá diversidade suficiente.
corpus = list(texts[:1000]) if not isinstance(texts, Iterator) else texts
try:
self.train_parallel_from_stream(
text_iterator=iter(corpus) if isinstance(corpus, list) else corpus,
save_path=None, # in-memory only
min_frequency=min_frequency,
num_workers=1, # single worker for fit() (faster for small corpus)
show_progress=False,
chunk_size=500,
)
except Exception as e:
logger.warning(f"BBPETokenizer.fit() train_parallel_from_stream failed: {e}")
# Fallback: cria um tokenizer byte-level mínimo (256 bytes + specials)
# sem merges, permitindo que o KLS continue funcionando.
try:
from tokenizers import Tokenizer
from tokenizers.models import BPE
from tokenizers.pre_tokenizers import ByteLevel
# Vocab inicial: 256 bytes base mapeados para [4, 260)
# (deixando 0-3 para speciais BOS/PAD/EOS/UNK)
byte_vocab = {chr(b): b + 4 for b in range(256)}
# Adiciona speciais
for tok, tid in [(BOS_TOKEN, BOS_ID), (PAD_TOKEN, PAD_ID),
(EOS_TOKEN, EOS_ID), (UNK_TOKEN, UNK_ID)]:
byte_vocab[tok] = tid
self._tokenizer = Tokenizer(BPE(byte_vocab, {}, unk_token=UNK_TOKEN))
self._tokenizer.pre_tokenizer = ByteLevel(add_prefix_space=False)
self._build_vocab_cache()
self.vocab_size = len(self._vocab)
logger.info(f"BBPETokenizer.fit() fallback byte-level: vocab_size={self.vocab_size}")
except Exception as e2:
logger.error(f"BBPETokenizer.fit() fallback failed: {e2}")
raise
def save(self, path: Union[str, Path]) -> None:
"""Salva o tokenizer em arquivo JSON (formato HuggingFace)."""
if self._tokenizer is None:
raise RuntimeError("Tokenizer não treinado. Chame train_*() primeiro.")
path = Path(path)
path.parent.mkdir(parents=True, exist_ok=True)
self._tokenizer.save(str(path))
logger.info("BBPE tokenizer saved: %s", path)
@classmethod
def load(cls, path: Union[str, Path]) -> "BBPETokenizer":
"""Carrega um tokenizer salvo."""
from tokenizers import Tokenizer
path = Path(path)
if not path.exists():
raise FileNotFoundError(f"Tokenizer file not found: {path}")
instance = cls() # default vocab_size
instance._tokenizer = Tokenizer.from_file(str(path))
instance._build_vocab_cache()
instance.vocab_size = len(instance._vocab)
logger.info(
"BBPE tokenizer loaded: %s (vocab_size=%d)",
path, instance.vocab_size,
)
return instance
def _build_vocab_cache(self) -> None:
"""Constrói caches internos de vocab/id_to_token."""
if self._tokenizer is None:
return
self._vocab = self._tokenizer.get_vocab()
self._id_to_token = {v: k for k, v in self._vocab.items()}
# ------------------------------------------------------------------
# Encode / Decode
# ------------------------------------------------------------------
def encode(
self,
text: str,
add_special_tokens: bool = False,
max_length: Optional[int] = None,
truncation: bool = True,
) -> List[int]:
"""Codifica um texto em IDs de tokens.
Args:
text: string de entrada
add_special_tokens: adicionar <s>...</s>
max_length: truncar para este tamanho (após special tokens)
truncation: se False e max_length excedido, raise em vez de truncar
"""
if self._tokenizer is None:
raise RuntimeError("Tokenizer não carregado.")
enc = self._tokenizer.encode(text, add_special_tokens=add_special_tokens)
ids = enc.ids
if max_length is not None:
if len(ids) > max_length:
if not truncation:
raise ValueError(
f"Input too long: {len(ids)} > {max_length} and truncation=False"
)
if add_special_tokens:
ids = ids[:max_length - 1] + [EOS_ID] if max_length >= 1 else [EOS_ID]
else:
ids = ids[:max_length]
return ids
def encode_batch(
self,
texts: List[str],
add_special_tokens: bool = False,
max_length: Optional[int] = None,
) -> List[List[int]]:
"""Codifica um batch de textos (sequencial)."""
return [self.encode(t, add_special_tokens, max_length) for t in texts]
def encode_batch_parallel(
self,
texts: List[str],
add_special_tokens: bool = False,
max_length: Optional[int] = None,
max_workers: Optional[int] = None,
) -> List[List[int]]:
"""Codifica um batch de textos em paralelo (ThreadPoolExecutor).
A codificação é embaraçosamente paralelizável: cada texto é codificado
independentemente. Usa threads (não processos) porque o tokenizers
library libera o GIL durante a codificação C++.
Args:
texts: lista de strings
add_special_tokens: adicionar <s>...</s>
max_length: truncar para este tamanho
max_workers: número de threads (default: min(32, len(texts)))
Returns:
Lista de listas de IDs.
"""
if not texts:
return []
workers = max_workers or min(32, max(1, len(texts)))
with ThreadPoolExecutor(max_workers=workers) as executor:
results = list(executor.map(
lambda t: self.encode(t, add_special_tokens, max_length),
texts,
))
return results
def decode(
self,
ids: List[int],
skip_special_tokens: bool = True,
) -> str:
"""Decodifica IDs de volta para texto."""
if self._tokenizer is None:
raise RuntimeError("Tokenizer não carregado.")
return self._tokenizer.decode(ids, skip_special_tokens=skip_special_tokens)
def decode_batch(
self,
ids_list: List[List[int]],
skip_special_tokens: bool = True,
) -> List[str]:
"""Decodifica um batch de listas de IDs."""
return [self.decode(ids, skip_special_tokens) for ids in ids_list]
# ------------------------------------------------------------------
# PyTorch integration
# ------------------------------------------------------------------
def encode_tensor(
self,
texts: List[str],
max_length: int,
pad_to_max_length: bool = True,
add_special_tokens: bool = False,
):
"""Codifica textos para tensor PyTorch (B, T) com padding.
Returns:
input_ids: LongTensor (B, T)
attention_mask: LongTensor (B, T) — 1 para tokens reais, 0 para pad
"""
import torch
batch_ids = [
self.encode(t, add_special_tokens=add_special_tokens, max_length=max_length)
for t in texts
]
if pad_to_max_length:
padded = []
masks = []
for ids in batch_ids:
if len(ids) >= max_length:
padded.append(ids[:max_length])
masks.append([1] * max_length)
else:
pad_len = max_length - len(ids)
padded.append(ids + [PAD_ID] * pad_len)
masks.append([1] * len(ids) + [0] * pad_len)
input_ids = torch.tensor(padded, dtype=torch.long)
attention_mask = torch.tensor(masks, dtype=torch.long)
else:
input_ids = [torch.tensor(ids, dtype=torch.long) for ids in batch_ids]
attention_mask = [torch.ones(len(ids), dtype=torch.long) for ids in batch_ids]
return input_ids, attention_mask
# ------------------------------------------------------------------
# Properties
# ------------------------------------------------------------------
@property
def vocab(self) -> Dict[str, int]:
if self._vocab is None:
self._build_vocab_cache()
return self._vocab or {}
@property
def id_to_token(self) -> Dict[int, str]:
if self._id_to_token is None:
self._build_vocab_cache()
return self._id_to_token or {}
@property
def actual_vocab_size(self) -> int:
"""Tamanho real do vocabulário carregado/treinado."""
return len(self.vocab)
@property
def merges(self) -> List[Tuple[str, str, str]]:
"""Lista de merges aprendidos (para inspeção)."""
return self._merges
def __len__(self) -> int:
return self.actual_vocab_size
# ------------------------------------------------------------------
# Validation
# ------------------------------------------------------------------
def validate_roundtrip(self, test_texts: List[str]) -> Dict[str, Any]:
"""Valida que encode→decode preserva o texto (Teorema 20).
Returns:
Dict com: success_rate, avg_compression, failed_examples
"""
results = {
"success_rate": 0.0,
"avg_compression": 0.0,
"avg_tokens_per_text": 0.0,
"total_texts": len(test_texts),
"failed_examples": [],
}
if not test_texts:
return results
successes = 0
total_compression = 0.0
total_tokens = 0
for text in test_texts:
try:
ids = self.encode(text, add_special_tokens=False)
decoded = self.decode(ids, skip_special_tokens=True)
expected = text
got = decoded
if expected == got:
successes += 1
else:
if self.add_prefix_space and got.startswith(" "):
got = got[1:]
if expected == got:
successes += 1
else:
results["failed_examples"].append({
"input": text[:100],
"decoded": decoded[:100],
"ids_count": len(ids),
})
n_bytes = len(text.encode("utf-8"))
n_tokens = len(ids)
if n_tokens > 0:
total_compression += n_bytes / n_tokens
total_tokens += n_tokens
except Exception as e:
results["failed_examples"].append({
"input": text[:100],
"error": str(e),
})
results["success_rate"] = successes / len(test_texts)
results["avg_compression"] = total_compression / max(successes, 1)
results["avg_tokens_per_text"] = total_tokens / max(successes, 1)
return results
# ---------------------------------------------------------------------------
# Byte-level alphabet (for BBPE initial alphabet) — compatibilidade
# ---------------------------------------------------------------------------
class ByteLevel:
"""Wrapper para o alfabeto byte-level (256 bytes)."""
@staticmethod
def alphabet() -> List[str]:
"""Retorna os 256 caracteres byte-level (Ġ-style do GPT-2)."""
return list(ALPHABET)
# ---------------------------------------------------------------------------
# Convenience factory
# ---------------------------------------------------------------------------
def create_or_load_tokenizer(
path: Union[str, Path],
text_iterator: Optional[Iterator[str]] = None,
vocab_size: int = DEFAULT_VOCAB_SIZE,
min_frequency: int = 2,
num_workers: int = 4,
) -> BBPETokenizer:
"""Carrega um tokenizer existente ou treina um novo (paralelo).
Args:
path: caminho do arquivo JSON
text_iterator: iterador de textos para treinar (se arquivo não existe)
vocab_size: tamanho do vocabulário (default 16.384)
min_frequency: frequência mínima para merges
num_workers: número de processos paralelos no treino
Returns:
BBPETokenizer carregado/treinado
"""
path = Path(path)
if path.exists():
logger.info("Loading existing BBPE tokenizer: %s", path)
return BBPETokenizer.load(path)
if text_iterator is None:
raise ValueError(
f"Tokenizer file {path} does not exist and no text_iterator provided"
)
logger.info("Training new BBPE tokenizer (parallel): %s", path)
tok = BBPETokenizer(vocab_size=vocab_size, num_workers=num_workers)
tok.train_parallel_from_stream(
text_iterator,
save_path=path,
min_frequency=min_frequency,
num_workers=num_workers,
)
return tok
__all__ = [
"BBPETokenizer",
"ByteLevel",
"create_or_load_tokenizer",
"BOS_TOKEN",
"PAD_TOKEN",
"EOS_TOKEN",
"UNK_TOKEN",
"BOS_ID",
"PAD_ID",
"EOS_ID",
"UNK_ID",
"SPECIAL_TOKENS",
"DEFAULT_VOCAB_SIZE",
"ALPHABET",
"BYTE_TO_SYMBOL",
"SYMBOL_TO_BYTE",
"bytes_to_unicode",
"distribute_texts",
"pre_tokenize_shard",
"count_pairs_in_shard",
"apply_merge_in_shard",
"build_bpe_from_merges",
]