CNN-BiGRU / cnn_bigru /tests /test_50_samples.py
PowerMachine's picture
v3.0: reorganiza arquivos sob cnn_bigru/ (preserva árvore de pastas)
49b8205 verified
Raw History Blame Contribute Delete
30.3 kB
"""test_50_samples.py — Teste de 50 amostras para erros lógicos ou falhas.
Executa o pipeline completo:
1. Inicializa xeon_runtime + memory_optimizer
2. Treina BBPE tokenizer em corpus sintético
3. Cria dataset streaming multimodal (50 amostras)
4. Instancia modelo multimodal + generator + verifier + anti-hallucination
5. Executa treinamento cooperativo com:
- Synergy search (N tentativas)
- Hypothesis controller (ativa em punições)
- Auto-learner (ajuste dinâmico de LR + spectral norm)
6. Executa inferência com sampling
7. Avalia perplexidade
8. Reporta erros lógicos/falhas encontradas
Usage:
python -m cnn_bigru.tests.test_50_samples
"""
from __future__ import annotations
import logging
import os
import sys
import time
import traceback
from pathlib import Path
from typing import Dict, List
# Setup paths (must come before torch import for xeon_runtime)
PROJECT_ROOT = Path(__file__).resolve().parent.parent.parent
sys.path.insert(0, str(PROJECT_ROOT))
# Ativa Xeon runtime ANTES de importar torch
from cnn_bigru.utils.xeon_runtime import optimize_xeon_environment
N_CORES = optimize_xeon_environment()
import numpy as np
import torch
from torch.utils.data import DataLoader
from cnn_bigru.tokenizer.bbpe_tokenizer import BBPETokenizer
from cnn_bigru.data.streaming_dataset import (
MultimodalStreamingDataset,
collate_multimodal,
)
from cnn_bigru.models.multimodal_model import MultimodalCNNBiGRU
from cnn_bigru.models.generator_verifier import (
GeneratorCNNBiGRU,
VerifierCNNBiGRU,
AntiHallucinationLayer,
)
from cnn_bigru.models.rope import RotaryPositionEmbedding
from cnn_bigru.models.transformer_block import (
TransformerBlockConfig,
CausalSelfAttention,
TransformerBlock,
TransformerDecoderStack,
)
from cnn_bigru.models.context_window import (
ContextWindowConfig,
ContextWindowManager,
KVCache,
)
from cnn_bigru.utils.ewc import EWCConfig, EWCState
from cnn_bigru.losses.losses import LossConfig, MultiLoss
from cnn_bigru.training.trainer import TrainerConfig, CooperativeTrainer
from cnn_bigru.training.auto_learner import (
AutoLearnConfig,
orthogonal_init_model,
)
from cnn_bigru.training.hypothesis_controller import (
HypothesisConfig,
HypothesisController,
)
from cnn_bigru.inference.inference import (
generate_with_sampling,
evaluate_perplexity,
make_default_context_window,
)
from cnn_bigru.utils.memory_optimizer import MemoryOptimizer
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s | %(levelname)s | %(name)s | %(message)s",
datefmt="%H:%M:%S",
)
logger = logging.getLogger("test_50")
# ============================================================================
# Corpus para treino do tokenizer
# ============================================================================
CORPUS_TEXTS = [
"o modelo cooperativo cnn-bigru combina fluxos paralelos",
"a atencao cruzada troca informacoes entre redes a e b",
"gru bidirecional captura dependencias temporais em ambas direcoes",
"a porta de atenuacao controla o fluxo de informacao entre celulas",
"a camada anti-alucinacao usa logica fuzzy de lukasiewicz",
"o verificador classifica passos com sigmoid binaria",
"penalidades linear e exponencial reduzem o erro grave",
"ajuste dinamico de lr estabiliza o treinamento cooperativo",
"normalizacao espectral limita a norma dos pesos",
"inicializacao ortogonal estabiliza matrizes recorrentes",
"o gerador decodificador usa atencao bahdanau sobre o encoder",
"a fusao multimodal combina texto imagem e audio",
"hipoteses sao ativadas quando o verificador pune o passo",
"synergy search tenta n configuracoes e escolhe a melhor",
"a perplexidade mede a confusao do modelo na previsao",
"top-k e top-p filtram a distribuicao de probabilidade",
"temperature ajusta a entropia das previsoes do modelo",
"presence penalty pune tokens ja aparecidos na geracao",
"frequency penalty pune proporcionalmente a frequencia do token",
"streaming dataset carrega amostras sem materializacao completa",
"byte bpe tokeniza qualquer string utf-8 sem unk",
"o otimizador adamw combina momentum e weight decay",
"gradient clipping previne explosao de gradiente",
"amp reduz vram com precisao mista bfloat16",
"memory optimizer limpa cache entre batches",
] * 4 # ~100 amostras para treino BBPE
def step_01_init_runtime() -> Dict:
"""Inicializa runtime e reporta configurações."""
logger.info("=" * 70)
logger.info("PASSO 1: Inicialização do runtime Xeon + memory optimizer")
logger.info("=" * 70)
from cnn_bigru.utils.xeon_runtime import get_runtime_info
info = get_runtime_info()
for k, v in info.items():
logger.info(" %s: %s", k, v)
mem = MemoryOptimizer(enable_amp=False)
mem.configure()
mem_info = mem.get_memory_mb()
logger.info(" Memória: %s", mem_info)
return {"runtime": info, "memory": mem_info}
def step_02_train_tokenizer() -> BBPETokenizer:
"""Treina BBPE tokenizer em corpus sintético."""
logger.info("=" * 70)
logger.info("PASSO 2: Treino do BBPE tokenizer")
logger.info("=" * 70)
t0 = time.time()
tok = BBPETokenizer.train_from_texts(
CORPUS_TEXTS,
vocab_size=2000,
min_frequency=1,
)
elapsed = time.time() - t0
logger.info(" Vocab size: %d", tok.vocab_size)
logger.info(" BOS/PAD/EOS IDs: %d/%d/%d", tok.bos_id, tok.pad_id, tok.eos_id)
logger.info(" Tempo: %.2fs", elapsed)
# Validação roundtrip
test_strs = CORPUS_TEXTS[:5]
roundtrip = tok.validate_roundtrip(test_strs)
logger.info(" Roundtrip accuracy: %.2f%%", roundtrip * 100)
if roundtrip < 0.8:
logger.warning(" Roundtrip baixo — possível problema no tokenizer")
return tok, roundtrip
def step_03_create_dataset(tokenizer: BBPETokenizer, n_samples: int = 50) -> DataLoader:
"""Cria dataset streaming multimodal com N amostras."""
logger.info("=" * 70)
logger.info("PASSO 3: Criação do dataset streaming multimodal (%d amostras)", n_samples)
logger.info("=" * 70)
# Para teste determinístico, usamos fallback sintético
dataset = MultimodalStreamingDataset(
n_samples=n_samples,
hf_datasets=[], # skip HF forçando sintético
use_synthetic_fallback=True,
seed=42,
image_size=(28, 28, 1),
audio_shape=(32, 40),
)
# Coleta todas as amostras em uma lista (para DataLoader iterável)
samples = list(dataset)
logger.info(" Amostras coletadas: %d", len(samples))
assert len(samples) == n_samples, f"Esperado {n_samples}, obtido {len(samples)}"
# Cria DataLoader com collate_fn
from torch.utils.data import DataLoader
loader = DataLoader(
samples,
batch_size=8,
shuffle=False,
collate_fn=lambda b: collate_multimodal(b, tokenizer, max_len=32),
)
logger.info(" DataLoader criado: batch_size=8, max_len=32")
return loader
def step_04_init_model(tokenizer: BBPETokenizer, device: str = "cpu") -> Dict:
"""Instancia todos os componentes do modelo."""
logger.info("=" * 70)
logger.info("PASSO 4: Instanciação dos modelos")
logger.info("=" * 70)
V = tokenizer.vocab_size
# Modelo multimodal principal
model = MultimodalCNNBiGRU(
vocab_size=V,
num_classes=3,
embedding_dim=32,
cnn_filters=32,
gru_hidden=32,
n_heads=4,
dropout=0.1,
pad_idx=tokenizer.pad_id,
img_channels=1,
img_hidden=16,
img_out_dim=32,
audio_freq=40,
audio_hidden=16,
audio_out_dim=32,
fusion_dim=64,
use_spectral_norm=False,
)
# Aplica inicialização ortogonal
n_init = orthogonal_init_model(model)
logger.info(" Modelo multimodal: %d params, %d camadas ortogonalizadas",
sum(p.numel() for p in model.parameters()), n_init)
# Generator (encoder + decoder)
generator = GeneratorCNNBiGRU(
vocab_size=V,
embedding_dim=32,
cnn_filters=32,
gru_hidden=32,
n_heads=4,
dropout=0.1,
pad_idx=tokenizer.pad_id,
max_proof_len=16,
)
orthogonal_init_model(generator)
# Verifier
verifier = VerifierCNNBiGRU(
vocab_size=V,
embedding_dim=32,
cnn_filters=32,
gru_hidden=32,
n_heads=4,
dropout=0.1,
pad_idx=tokenizer.pad_id,
)
orthogonal_init_model(verifier)
# Anti-hallucination
anti_hall = AntiHallucinationLayer(vocab_size=V, embed_dim=16)
logger.info(" Generator params: %d", sum(p.numel() for p in generator.parameters()))
logger.info(" Verifier params: %d", sum(p.numel() for p in verifier.parameters()))
logger.info(" Anti-hall params: %d", sum(p.numel() for p in anti_hall.parameters()))
return {
"model": model,
"generator": generator,
"verifier": verifier,
"anti_hall": anti_hall,
}
def step_05_train(
models: Dict,
tokenizer: BBPETokenizer,
dataloader: DataLoader,
device: str = "cpu",
ewc_state: EWCState = None,
) -> Dict:
"""Executa treinamento cooperativo.
NOVO v2.0: aceita ewc_state opcional para testar EWC.
"""
logger.info("=" * 70)
logger.info("PASSO 5: Treinamento cooperativo (synergy + hipóteses + auto-learn)")
logger.info("=" * 70)
cfg = TrainerConfig(
num_epochs=2,
max_batches_per_epoch=6,
batch_size=8,
n_synergy_attempts=3,
use_synergy_search=True,
use_hypotheses=True,
n_hypotheses=4,
loss_config=LossConfig(
alpha=1.0, beta=0.5, gamma_loss=0.3, delta=0.01,
lambda_penal=0.1, mu_exp_penal=0.05,
gamma_exp=1.0, threshold_err=0.5,
l2_reg=1e-5, use_curvature=True, curvature_eps=1e-3,
),
auto_config=AutoLearnConfig(
kappa_curv=0.01, grad_clip=1.0, spectral_radius=1.0,
lr_min=1e-6, lr_max=1e-2,
initial_lr_G=1e-3, initial_lr_V=1e-3,
use_spectral_norm=True, apply_after_step=True,
l2_reg=1e-5,
),
ewc_config=ewc_state.config if ewc_state else None,
device=device,
log_every=1,
use_verifier_real=True, # NOVO v2.0: usar verificador real
use_generator=True,
use_hypothesis_output=True,
)
trainer = CooperativeTrainer(
model=models["model"],
tokenizer=tokenizer,
config=cfg,
generator=models["generator"],
verifier=models["verifier"],
anti_hallucination=models["anti_hall"],
ewc_state=ewc_state,
)
try:
result = trainer.train(dataloader)
logger.info(" Treinamento OK | final_loss=%.4f | final_ppl=%.2f | elapsed=%.1fs",
result["final_loss"], result["final_ppl"], result["elapsed_s"])
return result
except Exception as e:
logger.error(" Treinamento FALHOU: %s", e)
logger.error(traceback.format_exc())
raise
def step_05b_test_new_modules(
models: Dict,
tokenizer: BBPETokenizer,
device: str = "cpu",
) -> Dict:
"""Testa os novos módulos v2.0: EWC, Context Window, RoPE, TransformerBlock.
NOVO v2.0: esta função exercita isoladamente cada novo módulo para garantir
que estão funcionando e integráveis ao pipeline.
"""
logger.info("=" * 70)
logger.info("PASSO 5b: Teste dos novos módulos v2.0 (EWC, ContextWindow, RoPE, TransformerBlock)")
logger.info("=" * 70)
results = {"ewc": None, "context_window": None, "rope": None, "transformer": None}
# ---------- EWC ----------
try:
logger.info(" [EWC] Testando EWCState...")
ewc_cfg = EWCConfig(
enabled=True,
lambda_ewc=100.0,
n_samples_fisher=5, # poucas amostras para teste rápido
online_gamma=0.9,
device=device,
)
ewc_state = EWCState(ewc_cfg)
# Antes de consolidar: penalty deve ser 0
pen0 = ewc_state.penalty(models["model"])
logger.info(" Penalty antes de consolidar: %.6f", float(pen0))
assert pen0.item() == 0.0, "EWC penalty deveria ser 0 antes de consolidar"
# Simular forward_fn para cálculo de Fisher
# IMPORTANTE: NÃO usar torch.no_grad() aqui — a EWC precisa de gradientes
# para calcular a matriz de Fisher (grad^2 da log-verossimilhança)
V = tokenizer.vocab_size
B = 2
T = 8
ids_a = torch.randint(0, V, (B, T), device=device)
ids_b = torch.randint(0, V, (B, T), device=device)
def forward_fn(_idx=None):
# Sem torch.no_grad() — a EWC interna ativa enable_grad()
out = models["model"](ids_a, ids_b, mode="classify")
return out["logits"]
# Consolidar (calcula Fisher e armazena theta_star)
ewc_state.consolidate(models["model"], forward_fn=forward_fn)
assert ewc_state.num_tasks() == 1, f"Esperado 1 tarefa, obtido {ewc_state.num_tasks()}"
# Penalty deve ser > 0 agora (parâmetros não mudaram desde consolidate, mas
# a penalidade ainda assim é computada e deve ser >= 0)
pen1 = ewc_state.penalty(models["model"])
logger.info(" Penalty após consolidar: %.6f", float(pen1.detach()))
assert pen1.item() >= 0.0, "EWC penalty deve ser >= 0"
results["ewc"] = {
"ok": True,
"penalty_before": float(pen0),
"penalty_after": float(pen1),
"num_tasks": ewc_state.num_tasks(),
}
logger.info(" [EWC] ✓ OK")
# Guardar para uso posterior no trainer
results["ewc_state"] = ewc_state
except Exception as e:
logger.error(" [EWC] ✗ FALHOU: %s", e)
logger.error(traceback.format_exc())
results["ewc"] = {"ok": False, "error": str(e)}
# ---------- Context Window ----------
try:
logger.info(" [ContextWindow] Testando ContextWindowManager + KVCache...")
cw = make_default_context_window(
max_window=32,
n_sink=2,
embed_dim=64,
n_heads=4,
n_layers=2,
device=device,
)
# Inicializar cache
cache = cw.init_cache(batch_size=1, device=torch.device(device))
assert cache is not None
assert cache.get_seq_len() == 0
# Append tokens
tokens1 = torch.tensor([[1, 2, 3, 4, 5]], dtype=torch.long, device=device)
window = cw.append_tokens(tokens1)
assert window.size(1) == 5
assert cw.kv_cache.get_seq_len() == 0 # cache só é populado quando chamamos KVCache.update
# Simular update do cache KV
dummy_k = torch.randn(1, 4, 5, 16, device=device) # [B, n_heads, T, head_dim]
dummy_v = torch.randn(1, 4, 5, 16, device=device)
new_k, new_v = cache.update(0, dummy_k, dummy_v)
assert cache.get_seq_len() == 5
# Adicionar mais tokens para forçar eviction
tokens2 = torch.tensor([[6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20]],
dtype=torch.long, device=device)
cw.append_tokens(tokens2)
# Total histórico: 5 + 15 = 20, mas max_window=32, então sem eviction ainda
assert cw.token_history and sum(t.size(1) for t in cw.token_history) == 20
# Forçar eviction adicionando mais tokens
tokens3 = torch.tensor([[21, 22, 23, 24, 25, 26, 27, 28, 29, 30, 31, 32, 33, 34]],
dtype=torch.long, device=device)
cw.append_tokens(tokens3) # total = 34 > 32 = max_window
cw.evict_cache()
info = cw.get_info()
logger.info(" Context window info: %s", info)
results["context_window"] = {"ok": True, "info": info}
logger.info(" [ContextWindow] ✓ OK")
except Exception as e:
logger.error(" [ContextWindow] ✗ FALHOU: %s", e)
logger.error(traceback.format_exc())
results["context_window"] = {"ok": False, "error": str(e)}
# ---------- RoPE ----------
try:
logger.info(" [RoPE] Testando RotaryPositionEmbedding...")
head_dim = 32
max_seq = 16
rope = RotaryPositionEmbedding(head_dim=head_dim, max_seq_len=max_seq).to(device)
# x: [B, n_heads, T, head_dim]
x = torch.randn(2, 4, 8, head_dim, device=device)
x_rot = rope(x)
assert x_rot.shape == x.shape, f"Shape mismatch: {x_rot.shape} vs {x.shape}"
# Verificar que RoPE preserva a norma (é uma rotação)
norm_before = x.norm(dim=-1)
norm_after = x_rot.norm(dim=-1)
diff = (norm_before - norm_after).abs().max().item()
logger.info(" Norma preservada (diff=%.6f)", diff)
assert diff < 1e-4, f"RoPE não preservou norma (diff={diff})"
# Testar com offset (para geração autoregressiva)
x_rot2 = rope(x, offset=10)
assert x_rot2.shape == x.shape
results["rope"] = {"ok": True, "norm_diff": diff}
logger.info(" [RoPE] ✓ OK")
except Exception as e:
logger.error(" [RoPE] ✗ FALHOU: %s", e)
logger.error(traceback.format_exc())
results["rope"] = {"ok": False, "error": str(e)}
# ---------- TransformerBlock ----------
try:
logger.info(" [TransformerBlock] Testando CausalSelfAttention + TransformerBlock...")
embed_dim = 64
n_heads = 4
n_layers = 2
V = tokenizer.vocab_size
block_config = TransformerBlockConfig(
embed_dim=embed_dim,
n_heads=n_heads,
ff_dim=4 * embed_dim,
dropout=0.1,
use_rope=True,
max_seq_len=32,
)
block = TransformerBlock(block_config).to(device)
# Forward sem cache
x = torch.randn(2, 8, embed_dim, device=device)
out, k, v = block(x)
assert out.shape == x.shape, f"Output shape: {out.shape} vs {x.shape}"
assert k.shape == (2, n_heads, 8, embed_dim // n_heads)
assert v.shape == (2, n_heads, 8, embed_dim // n_heads)
# Forward com padding mask
mask = torch.tensor([[1, 1, 1, 1, 0, 0, 0, 0], [1, 1, 1, 1, 1, 1, 1, 0]],
dtype=torch.float, device=device)
out2, k2, v2 = block(x, padding_mask=mask)
assert out2.shape == x.shape
# Forward com cache KV
cache_k = None
cache_v = None
for t in range(3):
x_t = torch.randn(2, 1, embed_dim, device=device)
out_t, cache_k, cache_v = block(
x_t, kv_cache_k=cache_k, kv_cache_v=cache_v, position_offset=t
)
assert out_t.shape == (2, 1, embed_dim)
assert cache_k.size(2) == t + 1, f"Cache len: {cache_k.size(2)}, esperado {t+1}"
logger.info(" Cache KV após 3 steps: %d tokens", cache_k.size(2))
# Testar TransformerDecoderStack completo
decoder = TransformerDecoderStack(
vocab_size=V,
embed_dim=embed_dim,
n_heads=n_heads,
n_layers=n_layers,
max_seq_len=32,
use_rope=True,
pad_id=tokenizer.pad_id,
weight_tying=True,
).to(device)
idx = torch.randint(0, V, (2, 8), device=device)
logits, new_caches = decoder(idx)
assert logits.shape == (2, 8, V), f"Logits shape: {logits.shape}"
assert len(new_caches) == n_layers
# Verificar weight tying
assert decoder.lm_head.weight is decoder.token_embedding.weight
results["transformer"] = {
"ok": True,
"block_output_shape": list(out.shape),
"decoder_logits_shape": list(logits.shape),
"weight_tying": True,
}
logger.info(" [TransformerBlock] ✓ OK")
except Exception as e:
logger.error(" [TransformerBlock] ✗ FALHOU: %s", e)
logger.error(traceback.format_exc())
results["transformer"] = {"ok": False, "error": str(e)}
return results
def step_06_inference(
model: MultimodalCNNBiGRU,
tokenizer: BBPETokenizer,
device: str = "cpu",
generator: GeneratorCNNBiGRU = None,
use_context_window: bool = True,
) -> Dict:
"""Executa inferência com sampling.
NOVO v2.0: usa generator (se disponível) para geração autoregressiva REAL,
e integra context_window para suporte a sequências longas.
"""
logger.info("=" * 70)
logger.info("PASSO 6: Inferência com Temperatura + Top-K + Top-P + Penalidades")
if generator is not None:
logger.info(" (usando GeneratorCNNBiGRU para geração autoregressiva REAL)")
else:
logger.info(" (sem generator — usando fallback single-step)")
logger.info("=" * 70)
prompts = [
("o modelo coopera entre", "fluxos paralelos"),
("atencao cruzada troca", "informacoes entre redes"),
("gru bidirecional captura", "dependencias temporais"),
]
# Inicializar context window se solicitado
cw = None
if use_context_window:
try:
cw = make_default_context_window(
max_window=64,
n_sink=2,
embed_dim=32,
n_heads=4,
n_layers=2,
device=device,
)
logger.info(" Context window ativo (max_window=64, sink=2)")
except Exception as e:
logger.warning(" Context window falhou ao inicializar: %s — desativando", e)
cw = None
results = []
for prompt_a, prompt_b in prompts:
try:
result = generate_with_sampling(
model=model,
tokenizer=tokenizer,
prompt_a=prompt_a,
prompt_b=prompt_b,
max_new_tokens=16,
temperature=0.7,
top_k=20,
top_p=0.9,
presence_penalty=0.3,
frequency_penalty=0.3,
device=device,
generator=generator,
context_window=cw,
)
logger.info(" Prompt A: %s", prompt_a)
logger.info(" Prompt B: %s", prompt_b)
logger.info(" Gerado: %s", result["text"][:100])
logger.info(" Metrics: %s", result["metrics"])
results.append({"prompt_a": prompt_a, "prompt_b": prompt_b, **result})
except Exception as e:
logger.error(" Inferência falhou para (%s, %s): %s", prompt_a, prompt_b, e)
logger.error(traceback.format_exc())
results.append({"prompt_a": prompt_a, "prompt_b": prompt_b, "error": str(e)})
return {"results": results}
def step_07_eval_ppl(
model: MultimodalCNNBiGRU,
tokenizer: BBPETokenizer,
dataloader: DataLoader,
device: str = "cpu",
generator: GeneratorCNNBiGRU = None,
) -> Dict:
"""Avalia perplexidade.
NOVO v2.0: usa generator (se disponível) para PPL real baseado em
geração autoregressiva com teacher forcing.
"""
logger.info("=" * 70)
logger.info("PASSO 7: Avaliação de Perplexidade (PPL)")
if generator is not None:
logger.info(" (usando GeneratorCNNBiGRU com teacher forcing)")
else:
logger.info(" (sem generator — usando fallback single-step)")
logger.info("=" * 70)
try:
result = evaluate_perplexity(
model=model,
dataloader=dataloader,
tokenizer=tokenizer,
device=device,
max_batches=5,
generator=generator,
)
logger.info(" Loss: %.4f | PPL: %.2f | batches: %d | used_generator: %s",
result["loss"], result["ppl"], result["n_batches"],
result.get("used_generator", False))
return result
except Exception as e:
logger.error(" PPL falhou: %s", e)
logger.error(traceback.format_exc())
return {"error": str(e)}
def step_08_report_errors(errors: List[str], warnings: List[str]) -> None:
"""Reporta erros lógicos ou falhas encontradas."""
logger.info("=" * 70)
logger.info("PASSO 8: Relatório de erros lógicos ou falhas")
logger.info("=" * 70)
if not errors:
logger.info(" ✓ NENHUM ERRO CRÍTICO encontrado")
else:
logger.error(" ✗ %d ERRO(S) CRÍTICO(S) encontrado(s):", len(errors))
for e in errors:
logger.error(" - %s", e)
if not warnings:
logger.info(" ✓ Nenhum warning relevante")
else:
logger.warning(" ⚠ %d warning(s):", len(warnings))
for w in warnings:
logger.warning(" - %s", w)
logger.info("=" * 70)
def main():
"""Executa o teste completo de 50 amostras."""
logger.info("INICIANDO TESTE DE 50 AMOSTRAS — CNN-BiGRU MULTIMODAL COOPERATIVO")
logger.info("Project root: %s", PROJECT_ROOT)
logger.info("Device: %s", "cpu")
logger.info("Cores: %d", N_CORES)
errors: List[str] = []
warnings: List[str] = []
try:
# Step 1
step_01_init_runtime()
# Step 2
tokenizer, roundtrip = step_02_train_tokenizer()
if roundtrip < 0.8:
warnings.append(f"BBPE roundtrip accuracy = {roundtrip:.2%} (esperado >= 80%)")
# Step 3
dataloader = step_03_create_dataset(tokenizer, n_samples=50)
# Step 4
models = step_04_init_model(tokenizer, device="cpu")
# Verificação: forward pass simples antes do treino
try:
batch = next(iter(dataloader))
with torch.no_grad():
out = models["model"](
batch["input_ids_a"], batch["input_ids_b"],
images=batch["images"], audios=batch["audios"],
mode="classify",
)
logger.info(" Forward pass OK | logits shape: %s", out["logits"].shape)
assert out["logits"].shape == (batch["input_ids_a"].size(0), 3), \
f"Shape inesperado: {out['logits'].shape}"
except Exception as e:
errors.append(f"Forward pass inicial falhou: {e}")
logger.error(traceback.format_exc())
# Step 5b: Testar novos módulos v2.0 (EWC, Context Window, RoPE, TransformerBlock)
ewc_state = None
if not errors:
try:
new_modules_result = step_05b_test_new_modules(models, tokenizer, device="cpu")
# Verificar se algum módulo falhou
for mod_name, mod_res in new_modules_result.items():
if mod_name == "ewc_state":
continue
if isinstance(mod_res, dict) and not mod_res.get("ok", True):
errors.append(f"Módulo {mod_name} falhou: {mod_res.get('error', 'unknown')}")
# Guardar EWC state para usar no treino
ewc_state = new_modules_result.get("ewc_state")
if ewc_state is None:
warnings.append("EWC state não foi criado — EWC não será testado no treino")
except Exception as e:
errors.append(f"Teste de novos módulos falhou: {e}")
logger.error(traceback.format_exc())
# Step 5: Treinamento (com EWC se disponível)
if not errors:
try:
train_result = step_05_train(models, tokenizer, dataloader, device="cpu",
ewc_state=ewc_state)
except Exception as e:
errors.append(f"Treinamento falhou: {e}")
train_result = None
else:
train_result = None
# Step 6: Inferência com generator + context window
if not errors:
try:
step_06_inference(
models["model"], tokenizer, device="cpu",
generator=models.get("generator"),
use_context_window=True,
)
except Exception as e:
errors.append(f"Inferência falhou: {e}")
logger.error(traceback.format_exc())
# Step 7: PPL com generator
if not errors:
try:
step_07_eval_ppl(
models["model"], tokenizer, dataloader, device="cpu",
generator=models.get("generator"),
)
except Exception as e:
warnings.append(f"PPL falhou (não crítico): {e}")
# Step 8
step_08_report_errors(errors, warnings)
# Resumo final
logger.info("=" * 70)
logger.info("RESUMO FINAL DO TESTE DE 50 AMOSTRAS")
logger.info("=" * 70)
logger.info(" Amostras processadas: 50")
logger.info(" Erros críticos: %d", len(errors))
logger.info(" Warnings: %d", len(warnings))
if train_result:
logger.info(" Loss final: %.4f", train_result["final_loss"])
logger.info(" PPL final: %.2f", train_result["final_ppl"])
logger.info(" Tempo: %.1fs", train_result["elapsed_s"])
logger.info(" Synergy attempts: %d", len(train_result.get("synergy_history", [])))
logger.info(" Hypothesis activations: %d",
sum(h.get("hypothesis_activations", 0) for h in train_result.get("history", [])))
if errors:
logger.error(" STATUS: FALHA — %d erro(s)", len(errors))
return 1
else:
logger.info(" STATUS: SUCESSO")
return 0
except Exception as e:
logger.error("ERRO FATAL: %s", e)
logger.error(traceback.format_exc())
return 2
if __name__ == "__main__":
try:
rc = main()
except SystemExit:
raise
except Exception as e:
logger.error("Unhandled exception: %s", e)
rc = 2
# Evita o "Fatal Python error: PyGILState_Release" no shutdown causado
# por threads do tokenizers/datasets que ainda estão ativas.
import os
os._exit(rc)