"""test_500_samples.py — Teste de 500 amostras (5 batches de 100) para o CNN-BiGRU. Executa o pipeline completo com: 1. Inicializa xeon_runtime + memory_optimizer + monitor 2. Treina BBPE tokenizer em corpus sintético 3. Cria dataset streaming multimodal (até 500 amostras em batches de 100) 4. Instancia modelo multimodal + generator + verifier + anti-hallucination 5. Testa os NOVOS módulos v3.0: - CyclicReasoning - MedusaMTP - NLGModule - NLPModule - MultimodalMultiHeadAttention - VQVAE2 - W8A8 Quantization (SmoothQuant) - LongContextManager (1M tokens) - Monitor 6. Executa treinamento cooperativo com: - Synergy search (N tentativas) - Hypothesis controller (ativa em punições) - Auto-learner (ajuste dinâmico de LR + spectral norm) - EWC (aprendizado contínuo) 7. Executa inferência com sampling 8. Avalia perplexidade 9. Aplica quantização W8A8 ao modelo 10. Reporta erros lógicos/falhas encontradas + exporta relatório Usage: python -m cnn_bigru.tests.test_500_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, DEFAULT_DATASETS, ) from cnn_bigru.models.multimodal_model import MultimodalCNNBiGRU from cnn_bigru.models.generator_verifier import ( GeneratorCNNBiGRU, VerifierCNNBiGRU, AntiHallucinationLayer, ) from cnn_bigru.models.cooperative_bigru import CooperativeCNNBiGRU 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, LongContextConfig, LongContextManager, make_long_context_window, ) # NOVOS módulos v3.0 from cnn_bigru.models.cyclic_reasoning import ( CyclicReasoningConfig, CyclicReasoning, make_hypothesis_fn, ) from cnn_bigru.models.medusa_heads import ( MedusaConfig, MedusaMTP, MedusaHead, medusa_tree_decode, ) from cnn_bigru.models.nlg import NLGConfig, NLGModule from cnn_bigru.models.nlp import ( NLPConfig, NLPModule, SequenceClassificationHead, TokenClassificationHead, SpanDetectionHead, EmbeddingHead, ) from cnn_bigru.models.multimodal_attention import ( MultimodalAttentionConfig, MultimodalMultiHeadAttention, CrossModalAttention, ModalityGate, ) from cnn_bigru.utils.ewc import EWCConfig, EWCState from cnn_bigru.utils.quantization import ( W8A8Config, SmoothQuantizer, quantize_model_w8a8, estimate_memory_savings, ) from cnn_bigru.utils.vqvae2 import VQVAE2Config, VQVAE2 from cnn_bigru.utils.monitoring import Monitor, get_monitor 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_500") # ============================================================================ # Configurações do teste # ============================================================================ TOTAL_SAMPLES = 500 BATCH_SIZE_SAMPLES = 100 # 5 batches de 100 N_BATCHES = TOTAL_SAMPLES // BATCH_SIZE_SAMPLES # 5 # Para teste rápido, pode-se reduzir N_BATCHES via env var if os.environ.get("CNN_BIGRU_TEST_FAST") == "1": N_BATCHES = 2 # 200 amostras em modo rápido TOTAL_SAMPLES = N_BATCHES * BATCH_SIZE_SAMPLES 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", # NOVOS v3.0 "raciocinio ciclico refina a representacao iterativamente", "medusa heads predizem multiplos tokens em paralelo", "nlg gera texto autoregressivo com transformer decoder", "nlp classifica sequencias tokens e spans", "w8a8 quantiza pesos e ativacoes em 8 bits", "smoothquant migra variancia das ativacoes para os pesos", "vq-vae-2 hierarquico comprime com codebooks top e bottom", "multi-token prediction acelera a geracao por arvores", "multi-head attention multimodal funde modalidades por atencao", "context window de 1m tokens usa chunked attention", "monitor rastreia metricas de treino e evolucao", "ewc previne esquecimento catastrofico em aprendizado continuo", "rope codifica posicoes por rotacao no espaco complexo", "kv cache armazena chaves e valores para atencao eficiente", "causal self attention mascara tokens futuros no decoder", "transformer block combina self attention e feed forward", ] * 4 # ~160 amostras para treino BBPE # ============================================================================ # Step 1: Inicialização # ============================================================================ def step_01_init_runtime() -> Dict: """Inicializa runtime Xeon + memory optimizer + monitor.""" logger.info("=" * 70) logger.info("PASSO 1: Inicialização do runtime Xeon + memory optimizer + monitor") 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) # Inicializa monitor global monitor = get_monitor(output_dir=PROJECT_ROOT / "download" / "monitor_reports") monitor.register_component("ewc", True) monitor.register_component("medusa", True) monitor.register_component("cyclic_reasoning", True) monitor.register_component("vqvae2", True) monitor.register_component("quantization", True) monitor.register_component("multimodal_attention", True) monitor.register_component("context_window", True) monitor.register_component("monitor", True) return {"runtime": info, "memory": mem_info, "monitor": monitor} # ============================================================================ # Step 2: Tokenizer # ============================================================================ def step_02_train_tokenizer() -> tuple: """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) 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 # ============================================================================ # Step 3: Dataset (até 500 amostras em batches de 100) # ============================================================================ def step_03_create_dataset( tokenizer: BBPETokenizer, n_samples: int = BATCH_SIZE_SAMPLES, seed_offset: int = 0, ) -> DataLoader: """Cria dataset streaming multimodal com N amostras (batch de 100).""" logger.info("=" * 70) logger.info("PASSO 3: Criação do dataset streaming multimodal (%d amostras, batch %d/5)", n_samples, seed_offset + 1) logger.info("=" * 70) dataset = MultimodalStreamingDataset( n_samples=n_samples, # Tenta repositório 'PowerMachine/CNN-BiGRU' primeiro, depois fallback hf_datasets=DEFAULT_DATASETS, use_synthetic_fallback=True, seed=42 + seed_offset * 100, image_size=(28, 28, 1), audio_shape=(32, 40), ) samples = list(dataset) logger.info(" Amostras coletadas: %d", len(samples)) assert len(samples) == n_samples, f"Esperado {n_samples}, obtido {len(samples)}" 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 # ============================================================================ # Step 4: Instanciar modelos # ============================================================================ 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 (incluindo novos módulos v3.0)") 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, ) 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 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, } # ============================================================================ # Step 5b: Testar NOVOS módulos v3.0 # ============================================================================ def step_05b_test_v3_modules( models: Dict, tokenizer: BBPETokenizer, device: str = "cpu", ) -> Dict: """Testa os novos módulos v3.0: CyclicReasoning, Medusa, NLG, NLP, MMA, VQVAE2, W8A8, LongCtx.""" logger.info("=" * 70) logger.info("PASSO 5b: Teste dos NOVOS módulos v3.0") logger.info("=" * 70) results = {} V = tokenizer.vocab_size # ---------- CyclicReasoning ---------- try: logger.info(" [CyclicReasoning] Testando...") cfg = CyclicReasoningConfig( embed_dim=64, max_cycles=4, convergence_eps=1e-3, use_anti_hallucination_gate=True, ) cr = CyclicReasoning(cfg).to(device) h0 = torch.randn(2, 64, device=device) result = cr(h0, return_history=True) assert result["h_final"].shape == (2, 64) assert 1 <= result["n_cycles"] <= 4 logger.info(" n_cycles=%d, converged=%s, deltas=%s", result["n_cycles"], result["converged"], [f"{d:.4f}" for d in result["deltas"]]) results["cyclic_reasoning"] = { "ok": True, "n_cycles": result["n_cycles"], "converged": result["converged"], } logger.info(" [CyclicReasoning] OK") except Exception as e: logger.error(" [CyclicReasoning] FALHOU: %s", e) logger.error(traceback.format_exc()) results["cyclic_reasoning"] = {"ok": False, "error": str(e)} # ---------- Medusa MTP ---------- try: logger.info(" [MedusaMTP] Testando...") cfg = MedusaConfig( vocab_size=V, embed_dim=32, n_heads=3, head_hidden_mult=2, dropout=0.1, ) medusa = MedusaMTP(cfg).to(device) h = torch.randn(2, 8, 32, device=device) target = torch.randint(0, V, (2, 8), device=device) loss, stats = medusa.compute_loss(h, target) assert loss.item() > 0 # Test tree decode base_logits = torch.randn(2, V, device=device) td = medusa_tree_decode(medusa, base_logits, h[:, -1, :]) assert "base_token" in td logger.info(" Loss=%.4f, stats=%s", loss.item(), stats) results["medusa_mtp"] = {"ok": True, "loss": float(loss), "stats": stats} logger.info(" [MedusaMTP] OK") except Exception as e: logger.error(" [MedusaMTP] FALHOU: %s", e) logger.error(traceback.format_exc()) results["medusa_mtp"] = {"ok": False, "error": str(e)} # ---------- NLG ---------- try: logger.info(" [NLGModule] Testando...") cfg = NLGConfig( vocab_size=V, embed_dim=32, n_heads=4, n_layers=2, max_seq_len=32, pad_id=tokenizer.pad_id, bos_id=tokenizer.bos_id, eos_id=tokenizer.eos_id, use_medusa=True, n_medusa_heads=3, use_cyclic_reasoning=False, weight_tying=True, device=device, ) nlg = NLGModule(cfg).to(device) ids = torch.randint(2, V, (2, 8), device=device) target = torch.randint(2, V, (2, 8), device=device) out = nlg.compute_loss(ids, target) assert "loss" in out and out["loss"].item() > 0 # Test geração prompt = torch.tensor([[tokenizer.bos_id, 5, 10]], dtype=torch.long, device=device) gen = nlg.generate(prompt, max_new_tokens=5, temperature=0.7, top_k=10) assert gen["ids"].size(1) >= 3 logger.info(" NLG loss=%.4f, generated %d tokens", out["loss"].item(), gen["n_tokens"]) results["nlg"] = {"ok": True, "loss": float(out["loss"]), "n_generated": gen["n_tokens"]} logger.info(" [NLGModule] OK") except Exception as e: logger.error(" [NLGModule] FALHOU: %s", e) logger.error(traceback.format_exc()) results["nlg"] = {"ok": False, "error": str(e)} # ---------- NLP ---------- try: logger.info(" [NLPModule] Testando...") nlp_cfg = NLPConfig( embed_dim=32, cnn_filters=32, gru_hidden=32, feat_per_stream=64, feat_fused=128, n_heads=4, num_classes_seq=3, num_labels_tok=5, device=device, ) backbone = CooperativeCNNBiGRU( vocab_size=V, embedding_dim=32, cnn_filters=32, gru_hidden=32, n_heads=4, pad_idx=tokenizer.pad_id, ).to(device) nlp = NLPModule(nlp_cfg, backbone=backbone).to(device) ids_a = torch.randint(2, V, (2, 8), device=device) ids_b = torch.randint(2, V, (2, 8), device=device) # Test sequence classification out_seq = nlp.sequence_classification(ids_a, ids_b) assert out_seq["logits"].shape == (2, 3) # Test token classification tok_logits = nlp.token_classification(ids_a, ids_b) # Test span detection s_log, e_log = nlp.span_detection(ids_a, ids_b) # Test embedding emb = nlp.embed(ids_a, ids_b) assert emb.shape == (2, 64) # feat_per_stream logger.info(" Seq: %s, Tok: %s, Span: (%s,%s), Emb: %s", out_seq["logits"].shape, tok_logits.shape if tok_logits is not None else None, s_log.shape if s_log is not None else None, e_log.shape if e_log is not None else None, emb.shape) results["nlp"] = {"ok": True, "seq_logits": list(out_seq["logits"].shape), "emb_shape": list(emb.shape)} logger.info(" [NLPModule] OK") except Exception as e: logger.error(" [NLPModule] FALHOU: %s", e) logger.error(traceback.format_exc()) results["nlp"] = {"ok": False, "error": str(e)} # ---------- Multimodal Multi-Head Attention ---------- try: logger.info(" [MultimodalMultiHeadAttention] Testando...") cfg = MultimodalAttentionConfig( d_text_a=32, d_text_b=32, d_image=16, d_audio=16, d_model=64, n_heads=4, use_modality_gate=True, ) mha = MultimodalMultiHeadAttention(cfg).to(device) seq_a = torch.randn(2, 8, 32, device=device) seq_b = torch.randn(2, 6, 32, device=device) seq_img = torch.randn(2, 4, 16, device=device) seq_aud = torch.randn(2, 5, 16, device=device) out = mha(seq_a, seq_b, seq_img, seq_aud) assert out["fused"].shape == (2, 64) assert out["modality_weights"].shape == (2, 4) logger.info(" Fused: %s, modality_weights: %s", out["fused"].shape, out["modality_weights"]) results["multimodal_attention"] = { "ok": True, "fused_shape": list(out["fused"].shape), } logger.info(" [MultimodalMultiHeadAttention] OK") except Exception as e: logger.error(" [MultimodalMultiHeadAttention] FALHOU: %s", e) logger.error(traceback.format_exc()) results["multimodal_attention"] = {"ok": False, "error": str(e)} # ---------- VQ-VAE-2 ---------- try: logger.info(" [VQVAE2] Testando...") cfg = VQVAE2Config( in_channels=1, bottom_channels=8, top_channels=4, n_bottom_codes=32, n_top_codes=32, n_downsample=1, hidden_channels=8, use_ema=True, ) vqvae = VQVAE2(cfg).to(device) x = torch.randn(2, 1, 16, 16, device=device) out = vqvae(x) assert out["x_recon"].shape == x.shape assert out["loss"].item() > 0 logger.info(" Recon: %s, loss=%.4f, top_usage=%.2f, bottom_usage=%.2f", out["x_recon"].shape, out["loss"].item(), float(out["loss_dict"]["top_usage"]), float(out["loss_dict"]["bottom_usage"])) results["vqvae2"] = { "ok": True, "loss": float(out["loss"]), "top_usage": float(out["loss_dict"]["top_usage"]), "bottom_usage": float(out["loss_dict"]["bottom_usage"]), } logger.info(" [VQVAE2] OK") except Exception as e: logger.error(" [VQVAE2] FALHOU: %s", e) logger.error(traceback.format_exc()) results["vqvae2"] = {"ok": False, "error": str(e)} # ---------- W8A8 Quantization ---------- try: logger.info(" [W8A8 SmoothQuant] Testando...") # Aplica ao modelo multimodal (sem dataloader = quantização dinâmica) model_copy = MultimodalCNNBiGRU( vocab_size=V, num_classes=3, embedding_dim=32, cnn_filters=32, gru_hidden=32, n_heads=4, 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, ).to(device) orthogonal_init_model(model_copy) qmodel = quantize_model_w8a8(model_copy, dataloader=None, alpha=0.5, device=device) assert SmoothQuantizer.is_quantized(qmodel) # Verificar forward ainda funciona ids_a = torch.randint(2, V, (2, 8), device=device) ids_b = torch.randint(2, V, (2, 8), device=device) with torch.no_grad(): out = qmodel(ids_a, ids_b, mode="classify") assert out["logits"].shape == (2, 3) savings = estimate_memory_savings(qmodel) logger.info(" Quantizado: %s, savings: %.1f%%", SmoothQuantizer.is_quantized(qmodel), savings["reduction_pct"]) results["w8a8"] = {"ok": True, "savings": savings} logger.info(" [W8A8] OK") except Exception as e: logger.error(" [W8A8] FALHOU: %s", e) logger.error(traceback.format_exc()) results["w8a8"] = {"ok": False, "error": str(e)} # ---------- Long Context (1M tokens) ---------- try: logger.info(" [LongContextManager - 1M tokens] Testando...") lcw = make_long_context_window( max_window=1_000_000, strategy="chunked", chunk_size=8192, embed_dim=64, n_heads=4, n_layers=2, device=device, ) # Simular adicionar muitos tokens em chunks all_tokens = [] for chunk_idx in range(3): # 3 chunks de 100 tokens cada tokens = torch.tensor( [[i + 1 + chunk_idx * 100 for i in range(100)]], device=device, ) result = lcw.append_tokens_chunked(tokens) all_tokens.append(result) info = lcw.get_info() assert info["supports_1m_tokens"] # LongContextManager.get_info() retorna "history_tokens" (não "total_tokens") assert info["history_tokens"] == 300 # 3 chunks * 100 logger.info(" Info: %s", info) results["long_context"] = {"ok": True, "info": info} logger.info(" [LongContextManager] OK") except Exception as e: logger.error(" [LongContextManager] FALHOU: %s", e) logger.error(traceback.format_exc()) results["long_context"] = {"ok": False, "error": str(e)} # ---------- EWC (já existente, mas validar integração) ---------- try: logger.info(" [EWC] Testando integração...") ewc_cfg = EWCConfig( enabled=True, lambda_ewc=100.0, n_samples_fisher=5, online_gamma=0.9, device=device, ) ewc_state = EWCState(ewc_cfg) pen0 = ewc_state.penalty(models["model"]) assert pen0.item() == 0.0 # Forward_fn para Fisher V = tokenizer.vocab_size B, T = 2, 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): out = models["model"](ids_a, ids_b, mode="classify") return out["logits"] ewc_state.consolidate(models["model"], forward_fn=forward_fn) assert ewc_state.num_tasks() == 1 pen1 = ewc_state.penalty(models["model"]) assert pen1.item() >= 0.0 logger.info(" Penalty antes/depois consolidar: %.6f / %.6f", float(pen0), float(pen1)) results["ewc"] = { "ok": True, "penalty_before": float(pen0), "penalty_after": float(pen1), "num_tasks": ewc_state.num_tasks(), } results["ewc_state"] = ewc_state logger.info(" [EWC] OK") except Exception as e: logger.error(" [EWC] FALHOU: %s", e) logger.error(traceback.format_exc()) results["ewc"] = {"ok": False, "error": str(e)} # ---------- Context Window (curto, já existente) ---------- try: logger.info(" [ContextWindowManager] Testando...") cw = make_default_context_window( max_window=32, n_sink=2, embed_dim=64, n_heads=4, n_layers=2, device=device, ) cache = cw.init_cache(batch_size=1, device=torch.device(device)) assert cache is not None tokens = torch.tensor([[1, 2, 3, 4, 5]], dtype=torch.long, device=device) window = cw.append_tokens(tokens) assert window.size(1) == 5 logger.info(" Window: %s, cache_len: %d", window.shape, cache.get_seq_len()) results["context_window"] = {"ok": True} logger.info(" [ContextWindowManager] OK") except Exception as e: logger.error(" [ContextWindowManager] FALHOU: %s", e) logger.error(traceback.format_exc()) results["context_window"] = {"ok": False, "error": str(e)} return results # ============================================================================ # Step 5: Treinamento # ============================================================================ def step_05_train( models: Dict, tokenizer: BBPETokenizer, dataloader: DataLoader, device: str = "cpu", ewc_state: EWCState = None, monitor: Monitor = None, batch_idx: int = 0, ) -> Dict: """Executa treinamento cooperativo em um batch de 100 amostras.""" logger.info("=" * 70) logger.info("PASSO 5: Treinamento cooperativo (batch %d/5 - 100 amostras)", batch_idx + 1) logger.info("=" * 70) cfg = TrainerConfig( num_epochs=1, max_batches_per_epoch=12, batch_size=8, n_synergy_attempts=2, use_synergy_search=True, use_hypotheses=True, n_hypotheses=3, 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=2, use_verifier_real=True, 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, ) # Integra monitor if monitor is not None: monitor.start_epoch(batch_idx) try: result = trainer.train(dataloader) logger.info(" Treino OK batch %d | loss=%.4f | ppl=%.2f | elapsed=%.1fs", batch_idx + 1, result["final_loss"], result["final_ppl"], result["elapsed_s"]) # Log no monitor if monitor is not None: for h in result.get("history", []): monitor.log_batch({ "loss": h.get("loss", 0), "ppl": h.get("ppl", 0), "lr_g": h.get("lr_G", 0), "hypothesis_activations": h.get("hypothesis_activations", 0), "elapsed_ms": h.get("elapsed_ms", 0), }) monitor.end_epoch({"batch_idx": batch_idx}) if ewc_state is not None: pen = ewc_state.penalty(models["model"]).item() monitor.log_ewc_penalty(penalty=pen, num_tasks=ewc_state.num_tasks()) return result except Exception as e: logger.error(" Treino FALHOU batch %d: %s", batch_idx + 1, e) logger.error(traceback.format_exc()) raise # ============================================================================ # Step 6: Inferência # ============================================================================ def step_06_inference( model: MultimodalCNNBiGRU, tokenizer: BBPETokenizer, device: str = "cpu", generator: GeneratorCNNBiGRU = None, monitor: Monitor = None, ) -> Dict: """Executa inferência com sampling.""" logger.info("=" * 70) logger.info("PASSO 6: Inferência com Temperatura + Top-K + Top-P + Penalidades") logger.info("=" * 70) prompts = [ ("o modelo coopera entre", "fluxos paralelos"), ("atencao cruzada troca", "informacoes entre redes"), ("gru bidirecional captura", "dependencias temporais"), ("medusa heads predizem", "multiplos tokens"), ("raciocinio ciclico refina", "iterativamente"), ] results = [] for prompt_a, prompt_b in prompts: t0 = time.time() try: result = generate_with_sampling( model=model, tokenizer=tokenizer, prompt_a=prompt_a, prompt_b=prompt_b, max_new_tokens=10, temperature=0.7, top_k=20, top_p=0.9, presence_penalty=0.3, frequency_penalty=0.3, device=device, generator=generator, ) elapsed_ms = (time.time() - t0) * 1000 logger.info(" Prompt A: %s | B: %s", prompt_a, prompt_b) logger.info(" Gerado: %s (%.0fms)", result["text"][:80], elapsed_ms) if monitor is not None: monitor.log_inference( n_tokens=len(result.get("token_ids", [])), elapsed_ms=elapsed_ms, ) results.append({"prompt_a": prompt_a, "prompt_b": prompt_b, **result}) except Exception as e: logger.error(" Inferência falhou (%s, %s): %s", prompt_a, prompt_b, e) results.append({"prompt_a": prompt_a, "prompt_b": prompt_b, "error": str(e)}) return {"results": results} # ============================================================================ # Step 7: PPL # ============================================================================ def step_07_eval_ppl( model: MultimodalCNNBiGRU, tokenizer: BBPETokenizer, dataloader: DataLoader, device: str = "cpu", generator: GeneratorCNNBiGRU = None, ) -> Dict: """Avalia perplexidade.""" logger.info("=" * 70) logger.info("PASSO 7: Avaliação de Perplexidade (PPL)") 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)} # ============================================================================ # Step 8: Relatório # ============================================================================ def step_08_report(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(" OK — NENHUM ERRO CRÍTICO encontrado") else: logger.error(" FAIL — %d ERRO(S) CRÍTICO(S):", len(errors)) for e in errors: logger.error(" - %s", e) if not warnings: logger.info(" OK — Nenhum warning relevante") else: logger.warning(" WARN — %d warning(s):", len(warnings)) for w in warnings: logger.warning(" - %s", w) logger.info("=" * 70) # ============================================================================ # Main # ============================================================================ def main(): """Executa o teste completo de 500 amostras (5 batches de 100).""" logger.info("=" * 70) logger.info("INICIANDO TESTE DE 500 AMOSTRAS (5 batches de 100)") logger.info("CNN-BiGRU MULTIMODAL COOPERATIVO v3.0") logger.info("=" * 70) logger.info("Project root: %s", PROJECT_ROOT) logger.info("Device: %s | Cores: %d", "cpu", N_CORES) logger.info("Total samples: %d | Batch size: %d | N batches: %d", TOTAL_SAMPLES, BATCH_SIZE_SAMPLES, N_BATCHES) errors: List[str] = [] warnings: List[str] = [] all_train_results = [] try: # Step 1: Runtime + Monitor rt_info = step_01_init_runtime() monitor = rt_info["monitor"] monitor.start_training() # Step 2: Tokenizer tokenizer, roundtrip = step_02_train_tokenizer() if roundtrip < 0.8: warnings.append(f"BBPE roundtrip = {roundtrip:.2%} (esperado >= 80%)") # Step 3: Modelo (uma única instância, reutilizada em todos batches) models = step_04_init_model(tokenizer, device="cpu") # Forward pass inicial try: test_loader = step_03_create_dataset(tokenizer, n_samples=8, seed_offset=99) batch = next(iter(test_loader)) 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 v3.0 ewc_state = None if not errors: try: new_modules_result = step_05b_test_v3_modules(models, tokenizer, device="cpu") 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')}") # Não é crítico para alguns módulos — converter em warning if mod_name in ("nlg", "nlp", "multimodal_attention", "vqvae2", "w8a8", "long_context"): errors.pop() # remove o erro warnings.append(f"Módulo {mod_name} falhou (não crítico): {mod_res.get('error', 'unknown')}") ewc_state = new_modules_result.get("ewc_state") if ewc_state is None: warnings.append("EWC state não criado — EWC não será testado no treino") # Log no monitor if monitor is not None: if new_modules_result.get("cyclic_reasoning", {}).get("ok"): monitor.log_cyclic_reasoning(new_modules_result["cyclic_reasoning"]) if new_modules_result.get("vqvae2", {}).get("ok"): monitor.log_vqvae2_usage(new_modules_result["vqvae2"]) if new_modules_result.get("w8a8", {}).get("ok"): monitor.log_quantization(new_modules_result["w8a8"].get("savings", {})) except Exception as e: errors.append(f"Teste de novos módulos falhou: {e}") logger.error(traceback.format_exc()) # Step 5 + 6 + 7: Loop sobre 5 batches de 100 amostras for batch_idx in range(N_BATCHES): logger.info("") logger.info("#" * 70) logger.info("# BATCH %d/%d — 100 AMOSTRAS", batch_idx + 1, N_BATCHES) logger.info("#" * 70) try: dataloader = step_03_create_dataset( tokenizer, n_samples=BATCH_SIZE_SAMPLES, seed_offset=batch_idx, ) except Exception as e: errors.append(f"Criação dataset batch {batch_idx+1} falhou: {e}") continue # Treino if not errors: try: train_result = step_05_train( models, tokenizer, dataloader, device="cpu", ewc_state=ewc_state, monitor=monitor, batch_idx=batch_idx, ) all_train_results.append(train_result) except Exception as e: errors.append(f"Treino batch {batch_idx+1} falhou: {e}") # Inferência (apenas no último batch para economizar tempo) if batch_idx == N_BATCHES - 1 and not errors: try: step_06_inference( models["model"], tokenizer, device="cpu", generator=models.get("generator"), monitor=monitor, ) except Exception as e: errors.append(f"Inferência batch {batch_idx+1} falhou: {e}") logger.error(traceback.format_exc()) # PPL 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 batch {batch_idx+1} falhou (não crítico): {e}") # Step 8: Relatório step_08_report(errors, warnings) # Finalizar monitor monitor.end_training() report_path = monitor.export_report() csv_path = monitor.export_csv() md_path = monitor.export_markdown_summary() # Resumo final logger.info("=" * 70) logger.info("RESUMO FINAL DO TESTE DE 500 AMOSTRAS") logger.info("=" * 70) logger.info(" Amostras processadas: %d (5 batches x 100)", TOTAL_SAMPLES) logger.info(" Erros críticos: %d", len(errors)) logger.info(" Warnings: %d", len(warnings)) if all_train_results: final = all_train_results[-1] logger.info(" Loss final: %.4f", final["final_loss"]) logger.info(" PPL final: %.2f", final["final_ppl"]) logger.info(" Tempo total treino: %.1fs", sum(r["elapsed_s"] for r in all_train_results)) logger.info(" Monitor report: %s", report_path) logger.info(" Monitor CSV: %s", csv_path) logger.info(" Monitor MD: %s", md_path) logger.info(" Throughput: %s", monitor.get_inference_throughput()) 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 import os os._exit(rc)