CNN-BiGRU / cnn_bigru /tests /test_500_samples.py
PowerMachine's picture
v3.0: reorganiza arquivos sob cnn_bigru/ (preserva árvore de pastas)
49b8205 verified
Raw History Blame Contribute Delete
41.4 kB
"""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)