CNN-BiGRU / cnn_bigru /tokenizer /bbpe_tokenizer.py
PowerMachine's picture
v3.0: reorganiza arquivos sob cnn_bigru/ (preserva árvore de pastas)
49b8205 verified
Raw History Blame Contribute Delete
5.69 kB
"""bbpe_tokenizer.py — BBPE (Byte-Level BPE) Tokenizer simplificado.
Inspirado em PowerMachine/BiGRU_T_version/src/bigru_t/tokenizer/bbpe_tokenizer.py,
mas simplificado para o projeto CNN-BiGRU. Usa a biblioteca `tokenizers` da
HuggingFace com modelo BPE byte-level.
Cobertura universal UTF-8 (qualquer string é tokenizável sem <unk>).
"""
from __future__ import annotations
import json
import logging
import os
from pathlib import Path
from typing import List, Optional, Sequence
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
from tokenizers.trainers import BpeTrainer
logger = logging.getLogger(__name__)
BOS_TOKEN = "<s>"
PAD_TOKEN = "<pad>"
EOS_TOKEN = "</s>"
UNK_TOKEN = "<unk>"
SPECIAL_TOKENS = [BOS_TOKEN, PAD_TOKEN, EOS_TOKEN, UNK_TOKEN]
BOS_ID = 0
PAD_ID = 1
EOS_ID = 2
UNK_ID = 3
class BBPETokenizer:
"""Wrapper sobre Tokenizer (HuggingFace) com API conveniente.
Garante: dec(enc(s)) = s para todo s UTF-8.
"""
def __init__(self, tokenizer: Optional[Tokenizer] = None):
self._tok = tokenizer
@property
def tokenizer(self) -> Tokenizer:
if self._tok is None:
raise RuntimeError("Tokenizer não inicializado. Use train_from_* ou load.")
return self._tok
@property
def vocab_size(self) -> int:
return self.tokenizer.get_vocab_size()
@property
def bos_id(self) -> int:
return self.tokenizer.token_to_id(BOS_TOKEN) or BOS_ID
@property
def pad_id(self) -> int:
return self.tokenizer.token_to_id(PAD_TOKEN) or PAD_ID
@property
def eos_id(self) -> int:
return self.tokenizer.token_to_id(EOS_TOKEN) or EOS_ID
@property
def unk_id(self) -> int:
return self.tokenizer.token_to_id(UNK_TOKEN) or UNK_ID
def encode(self, text: str, add_special: bool = True) -> List[int]:
if add_special:
text = f"{BOS_TOKEN}{text}{EOS_TOKEN}"
return self.tokenizer.encode(text).ids
def encode_batch(self, texts: Sequence[str], add_special: bool = True) -> List[List[int]]:
if add_special:
texts = [f"{BOS_TOKEN}{t}{EOS_TOKEN}" for t in texts]
return [e.ids for e in self.tokenizer.encode_batch(list(texts))]
def decode(self, ids: List[int], skip_special: bool = True) -> str:
if skip_special:
ids = [i for i in ids if i not in (self.bos_id, self.pad_id, self.eos_id, self.unk_id)]
return self.tokenizer.decode(ids)
def save(self, path: str) -> None:
Path(path).parent.mkdir(parents=True, exist_ok=True)
self.tokenizer.save(path)
logger.info("Tokenizer salvo em %s", path)
@classmethod
def load(cls, path: str) -> "BBPETokenizer":
tok = Tokenizer.from_file(path)
return cls(tok)
@classmethod
def train_from_files(
cls,
files: Sequence[str],
vocab_size: int = 8000,
min_frequency: int = 2,
) -> "BBPETokenizer":
# Desativa paralelismo de tokenizers para evitar problemas de GIL no shutdown
os.environ["TOKENIZERS_PARALLELISM"] = "false"
tokenizer = Tokenizer(BPE(unk_token=UNK_TOKEN))
tokenizer.pre_tokenizer = ByteLevel(add_prefix_space=False)
trainer = BpeTrainer(
vocab_size=vocab_size,
min_frequency=min_frequency,
special_tokens=SPECIAL_TOKENS,
show_progress=False,
)
valid_files = [str(Path(p).resolve()) for p in files if Path(p).exists()]
if not valid_files:
raise FileNotFoundError(f"Nenhum arquivo válido em: {files}")
tokenizer.train(valid_files, trainer)
tokenizer.post_processor = ByteLevelProcessor(trim_offsets=False)
# Decoder byte-level: reverte a codificação byte-level para UTF-8
tokenizer.decoder = ByteLevelDecoder()
logger.info("BBPE treinado com vocab_size=%d", tokenizer.get_vocab_size())
return cls(tokenizer)
@classmethod
def train_from_texts(
cls,
texts: Sequence[str],
vocab_size: int = 8000,
min_frequency: int = 1,
) -> "BBPETokenizer":
"""Treina a partir de uma lista de textos em memória."""
# Salva temporariamente
import tempfile
with tempfile.NamedTemporaryFile(
mode="w", suffix=".txt", delete=False, encoding="utf-8"
) as f:
for line in texts:
line = line.strip()
if line:
f.write(line + "\n")
tmp_path = f.name
try:
return cls.train_from_files([tmp_path], vocab_size=vocab_size,
min_frequency=min_frequency)
finally:
try:
os.unlink(tmp_path)
except OSError:
pass
def validate_roundtrip(self, test_texts: Sequence[str]) -> float:
"""Valida que dec(enc(s)) == s. Retorna fração de sucessos."""
if not test_texts:
return 1.0
ok = 0
for s in test_texts:
try:
ids = self.encode(s, add_special=False)
decoded = self.decode(ids, skip_special=False)
if decoded == s:
ok += 1
except Exception:
pass
return ok / len(test_texts)
__all__ = [
"BBPETokenizer",
"BOS_TOKEN", "PAD_TOKEN", "EOS_TOKEN", "UNK_TOKEN",
"BOS_ID", "PAD_ID", "EOS_ID", "UNK_ID",
"SPECIAL_TOKENS",
]