"""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 ). """ 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 = "" PAD_TOKEN = "" EOS_TOKEN = "" UNK_TOKEN = "" 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", ]