v3.0: 9 novos módulos (CyclicReasoning, Medusa, NLG, NLP, MMA, VQVAE2, W8A8, LongContext 1M, Monitor) + bug fixes (attempt 1)
Browse files- __init__.py +151 -3
- data/streaming_dataset.py +13 -5
- docs/MATH_ANALYSIS.md +373 -0
- inference/inference.py +19 -19
- models/context_window.py +250 -0
- models/cyclic_reasoning.py +411 -0
- models/generator_verifier.py +3 -1
- models/medusa_heads.py +409 -0
- models/multimodal_attention.py +470 -0
- models/nlg.py +457 -0
- models/nlp.py +654 -0
- scripts/push_to_hf.py +94 -13
- tests/test_500_samples.py +1035 -0
- utils/monitoring.py +729 -0
- utils/quantization.py +585 -0
- utils/vqvae2.py +570 -0
__init__.py
CHANGED
|
@@ -1,5 +1,5 @@
|
|
| 1 |
"""
|
| 2 |
-
CNN-BiGRU — Modelo Multimodal Cooperativo Autoaprendível (
|
| 3 |
|
| 4 |
Implementação completa de um LLM multimodal com arquitetura CNN-BiGRU cooperativa:
|
| 5 |
- Núcleo dual-stream CNN-BiGRU com 3 níveis de pontes (cross-attention, gated cells, fusão)
|
|
@@ -10,7 +10,7 @@ Implementação completa de um LLM multimodal com arquitetura CNN-BiGRU cooperat
|
|
| 10 |
- Synergy search (N tentativas de combinação de hiperparâmetros)
|
| 11 |
- Múltiplas funções de perda: L_G, L_V, L_AH, L_reg, L_EWC
|
| 12 |
|
| 13 |
-
|
| 14 |
- EWC (Elastic Weight Consolidation) para aprendizado contínuo
|
| 15 |
- Context Window com cache KV (sliding window + sink)
|
| 16 |
- RoPE (Rotary Position Embeddings)
|
|
@@ -18,10 +18,21 @@ NOVO v2.0:
|
|
| 18 |
- Self-attention final layer verdadeira (CLS-token + multi-head)
|
| 19 |
- Weight tying opcional
|
| 20 |
- Inferência autoregressiva REAL via GeneratorCNNBiGRU
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 21 |
"""
|
| 22 |
from __future__ import annotations
|
| 23 |
|
| 24 |
-
__version__ = "
|
| 25 |
__author__ = "CNN-BiGRU Project"
|
| 26 |
|
| 27 |
# Versão e metadados
|
|
@@ -35,6 +46,14 @@ __all__ = [
|
|
| 35 |
"SemanticEmbedder",
|
| 36 |
"EWCConfig",
|
| 37 |
"EWCState",
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 38 |
# Runtime
|
| 39 |
"optimize_xeon_environment",
|
| 40 |
"get_runtime_info",
|
|
@@ -57,6 +76,27 @@ __all__ = [
|
|
| 57 |
"KVCache",
|
| 58 |
"ContextWindowManager",
|
| 59 |
"ContextWindowConfig",
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 60 |
# Losses
|
| 61 |
"LossConfig",
|
| 62 |
"MultiLoss",
|
|
@@ -90,16 +130,43 @@ def __getattr__(name: str):
|
|
| 90 |
if name == "EWCConfig":
|
| 91 |
return EWCConfig
|
| 92 |
return EWCState
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 93 |
if name in ("optimize_xeon_environment", "get_runtime_info"):
|
| 94 |
from .utils import xeon_runtime
|
| 95 |
if name == "optimize_xeon_environment":
|
| 96 |
return xeon_runtime.optimize_xeon_environment
|
| 97 |
return xeon_runtime.get_runtime_info
|
|
|
|
| 98 |
if name in ("MultimodalStreamingDataset", "collate_multimodal"):
|
| 99 |
from .data.streaming_dataset import MultimodalStreamingDataset, collate_multimodal
|
| 100 |
if name == "MultimodalStreamingDataset":
|
| 101 |
return MultimodalStreamingDataset
|
| 102 |
return collate_multimodal
|
|
|
|
| 103 |
if name == "CooperativeCNNBiGRU":
|
| 104 |
from .models.cooperative_bigru import CooperativeCNNBiGRU
|
| 105 |
return CooperativeCNNBiGRU
|
|
@@ -138,6 +205,7 @@ def __getattr__(name: str):
|
|
| 138 |
if name == "TransformerBlock":
|
| 139 |
return TransformerBlock
|
| 140 |
return TransformerDecoderStack
|
|
|
|
| 141 |
if name in ("KVCache", "ContextWindowManager", "ContextWindowConfig"):
|
| 142 |
from .models.context_window import KVCache, ContextWindowManager, ContextWindowConfig
|
| 143 |
if name == "KVCache":
|
|
@@ -145,11 +213,90 @@ def __getattr__(name: str):
|
|
| 145 |
if name == "ContextWindowManager":
|
| 146 |
return ContextWindowManager
|
| 147 |
return ContextWindowConfig
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 148 |
if name in ("LossConfig", "MultiLoss"):
|
| 149 |
from .losses.losses import LossConfig, MultiLoss
|
| 150 |
if name == "LossConfig":
|
| 151 |
return LossConfig
|
| 152 |
return MultiLoss
|
|
|
|
| 153 |
if name in ("TrainerConfig", "CooperativeTrainer"):
|
| 154 |
from .training.trainer import TrainerConfig, CooperativeTrainer
|
| 155 |
if name == "TrainerConfig":
|
|
@@ -171,6 +318,7 @@ def __getattr__(name: str):
|
|
| 171 |
if name == "HypothesisController":
|
| 172 |
return HypothesisController
|
| 173 |
return SynergySearcher
|
|
|
|
| 174 |
if name in ("generate_with_sampling", "evaluate_perplexity"):
|
| 175 |
from .inference.inference import generate_with_sampling, evaluate_perplexity
|
| 176 |
if name == "generate_with_sampling":
|
|
|
|
| 1 |
"""
|
| 2 |
+
CNN-BiGRU — Modelo Multimodal Cooperativo Autoaprendível (v3.0)
|
| 3 |
|
| 4 |
Implementação completa de um LLM multimodal com arquitetura CNN-BiGRU cooperativa:
|
| 5 |
- Núcleo dual-stream CNN-BiGRU com 3 níveis de pontes (cross-attention, gated cells, fusão)
|
|
|
|
| 10 |
- Synergy search (N tentativas de combinação de hiperparâmetros)
|
| 11 |
- Múltiplas funções de perda: L_G, L_V, L_AH, L_reg, L_EWC
|
| 12 |
|
| 13 |
+
v2.0:
|
| 14 |
- EWC (Elastic Weight Consolidation) para aprendizado contínuo
|
| 15 |
- Context Window com cache KV (sliding window + sink)
|
| 16 |
- RoPE (Rotary Position Embeddings)
|
|
|
|
| 18 |
- Self-attention final layer verdadeira (CLS-token + multi-head)
|
| 19 |
- Weight tying opcional
|
| 20 |
- Inferência autoregressiva REAL via GeneratorCNNBiGRU
|
| 21 |
+
|
| 22 |
+
NOVO v3.0:
|
| 23 |
+
- Cyclic Reasoning (Raciocínio Cíclico) com convergência antecipada
|
| 24 |
+
- NLG (Natural Language Generation) com Transformer Decoder + Medusa MTP
|
| 25 |
+
- NLP (Natural Language Processing) com 4 tarefas: SeqCls, TokCls, Span, Embed
|
| 26 |
+
- W8A8 Quantization via SmoothQuant (alpha=0.5, per-channel)
|
| 27 |
+
- VQ-VAE-2 Hierárquico (top + bottom codebooks, EMA update)
|
| 28 |
+
- Multi-Token Prediction (MTP) com Medusa Heads
|
| 29 |
+
- Multimodal Multi-Head Attention (cross-modal + modality gate)
|
| 30 |
+
- Long Context Window (1M tokens) via chunked / ring attention
|
| 31 |
+
- Monitor completo (treino, inferência, evolução, sistema)
|
| 32 |
"""
|
| 33 |
from __future__ import annotations
|
| 34 |
|
| 35 |
+
__version__ = "3.0.0"
|
| 36 |
__author__ = "CNN-BiGRU Project"
|
| 37 |
|
| 38 |
# Versão e metadados
|
|
|
|
| 46 |
"SemanticEmbedder",
|
| 47 |
"EWCConfig",
|
| 48 |
"EWCState",
|
| 49 |
+
"W8A8Config",
|
| 50 |
+
"SmoothQuantizer",
|
| 51 |
+
"quantize_model_w8a8",
|
| 52 |
+
"estimate_memory_savings",
|
| 53 |
+
"VQVAE2Config",
|
| 54 |
+
"VQVAE2",
|
| 55 |
+
"Monitor",
|
| 56 |
+
"get_monitor",
|
| 57 |
# Runtime
|
| 58 |
"optimize_xeon_environment",
|
| 59 |
"get_runtime_info",
|
|
|
|
| 76 |
"KVCache",
|
| 77 |
"ContextWindowManager",
|
| 78 |
"ContextWindowConfig",
|
| 79 |
+
"LongContextConfig",
|
| 80 |
+
"LongContextManager",
|
| 81 |
+
"make_long_context_window",
|
| 82 |
+
"CyclicReasoningConfig",
|
| 83 |
+
"CyclicReasoning",
|
| 84 |
+
"MedusaConfig",
|
| 85 |
+
"MedusaMTP",
|
| 86 |
+
"MedusaHead",
|
| 87 |
+
"medusa_tree_decode",
|
| 88 |
+
"NLGConfig",
|
| 89 |
+
"NLGModule",
|
| 90 |
+
"NLPConfig",
|
| 91 |
+
"NLPModule",
|
| 92 |
+
"SequenceClassificationHead",
|
| 93 |
+
"TokenClassificationHead",
|
| 94 |
+
"SpanDetectionHead",
|
| 95 |
+
"EmbeddingHead",
|
| 96 |
+
"MultimodalAttentionConfig",
|
| 97 |
+
"MultimodalMultiHeadAttention",
|
| 98 |
+
"CrossModalAttention",
|
| 99 |
+
"ModalityGate",
|
| 100 |
# Losses
|
| 101 |
"LossConfig",
|
| 102 |
"MultiLoss",
|
|
|
|
| 130 |
if name == "EWCConfig":
|
| 131 |
return EWCConfig
|
| 132 |
return EWCState
|
| 133 |
+
# W8A8 Quantization
|
| 134 |
+
if name in ("W8A8Config", "SmoothQuantizer"):
|
| 135 |
+
from .utils.quantization import W8A8Config, SmoothQuantizer
|
| 136 |
+
if name == "W8A8Config":
|
| 137 |
+
return W8A8Config
|
| 138 |
+
return SmoothQuantizer
|
| 139 |
+
if name == "quantize_model_w8a8":
|
| 140 |
+
from .utils.quantization import quantize_model_w8a8
|
| 141 |
+
return quantize_model_w8a8
|
| 142 |
+
if name == "estimate_memory_savings":
|
| 143 |
+
from .utils.quantization import estimate_memory_savings
|
| 144 |
+
return estimate_memory_savings
|
| 145 |
+
# VQ-VAE-2
|
| 146 |
+
if name in ("VQVAE2Config", "VQVAE2"):
|
| 147 |
+
from .utils.vqvae2 import VQVAE2Config, VQVAE2
|
| 148 |
+
if name == "VQVAE2Config":
|
| 149 |
+
return VQVAE2Config
|
| 150 |
+
return VQVAE2
|
| 151 |
+
# Monitor
|
| 152 |
+
if name in ("Monitor", "get_monitor"):
|
| 153 |
+
from .utils.monitoring import Monitor, get_monitor
|
| 154 |
+
if name == "Monitor":
|
| 155 |
+
return Monitor
|
| 156 |
+
return get_monitor
|
| 157 |
+
# Runtime
|
| 158 |
if name in ("optimize_xeon_environment", "get_runtime_info"):
|
| 159 |
from .utils import xeon_runtime
|
| 160 |
if name == "optimize_xeon_environment":
|
| 161 |
return xeon_runtime.optimize_xeon_environment
|
| 162 |
return xeon_runtime.get_runtime_info
|
| 163 |
+
# Data
|
| 164 |
if name in ("MultimodalStreamingDataset", "collate_multimodal"):
|
| 165 |
from .data.streaming_dataset import MultimodalStreamingDataset, collate_multimodal
|
| 166 |
if name == "MultimodalStreamingDataset":
|
| 167 |
return MultimodalStreamingDataset
|
| 168 |
return collate_multimodal
|
| 169 |
+
# Models
|
| 170 |
if name == "CooperativeCNNBiGRU":
|
| 171 |
from .models.cooperative_bigru import CooperativeCNNBiGRU
|
| 172 |
return CooperativeCNNBiGRU
|
|
|
|
| 205 |
if name == "TransformerBlock":
|
| 206 |
return TransformerBlock
|
| 207 |
return TransformerDecoderStack
|
| 208 |
+
# Context window (incl. 1M tokens)
|
| 209 |
if name in ("KVCache", "ContextWindowManager", "ContextWindowConfig"):
|
| 210 |
from .models.context_window import KVCache, ContextWindowManager, ContextWindowConfig
|
| 211 |
if name == "KVCache":
|
|
|
|
| 213 |
if name == "ContextWindowManager":
|
| 214 |
return ContextWindowManager
|
| 215 |
return ContextWindowConfig
|
| 216 |
+
if name in ("LongContextConfig", "LongContextManager", "make_long_context_window"):
|
| 217 |
+
from .models.context_window import (
|
| 218 |
+
LongContextConfig,
|
| 219 |
+
LongContextManager,
|
| 220 |
+
make_long_context_window,
|
| 221 |
+
)
|
| 222 |
+
if name == "LongContextConfig":
|
| 223 |
+
return LongContextConfig
|
| 224 |
+
if name == "LongContextManager":
|
| 225 |
+
return LongContextManager
|
| 226 |
+
return make_long_context_window
|
| 227 |
+
# Cyclic Reasoning
|
| 228 |
+
if name in ("CyclicReasoningConfig", "CyclicReasoning"):
|
| 229 |
+
from .models.cyclic_reasoning import CyclicReasoningConfig, CyclicReasoning
|
| 230 |
+
if name == "CyclicReasoningConfig":
|
| 231 |
+
return CyclicReasoningConfig
|
| 232 |
+
return CyclicReasoning
|
| 233 |
+
# Medusa MTP
|
| 234 |
+
if name in ("MedusaConfig", "MedusaMTP", "MedusaHead", "medusa_tree_decode"):
|
| 235 |
+
from .models.medusa_heads import (
|
| 236 |
+
MedusaConfig,
|
| 237 |
+
MedusaMTP,
|
| 238 |
+
MedusaHead,
|
| 239 |
+
medusa_tree_decode,
|
| 240 |
+
)
|
| 241 |
+
if name == "MedusaConfig":
|
| 242 |
+
return MedusaConfig
|
| 243 |
+
if name == "MedusaMTP":
|
| 244 |
+
return MedusaMTP
|
| 245 |
+
if name == "MedusaHead":
|
| 246 |
+
return MedusaHead
|
| 247 |
+
return medusa_tree_decode
|
| 248 |
+
# NLG
|
| 249 |
+
if name in ("NLGConfig", "NLGModule"):
|
| 250 |
+
from .models.nlg import NLGConfig, NLGModule
|
| 251 |
+
if name == "NLGConfig":
|
| 252 |
+
return NLGConfig
|
| 253 |
+
return NLGModule
|
| 254 |
+
# NLP
|
| 255 |
+
if name in ("NLPConfig", "NLPModule",
|
| 256 |
+
"SequenceClassificationHead", "TokenClassificationHead",
|
| 257 |
+
"SpanDetectionHead", "EmbeddingHead"):
|
| 258 |
+
from .models.nlp import (
|
| 259 |
+
NLPConfig,
|
| 260 |
+
NLPModule,
|
| 261 |
+
SequenceClassificationHead,
|
| 262 |
+
TokenClassificationHead,
|
| 263 |
+
SpanDetectionHead,
|
| 264 |
+
EmbeddingHead,
|
| 265 |
+
)
|
| 266 |
+
if name == "NLPConfig":
|
| 267 |
+
return NLPConfig
|
| 268 |
+
if name == "NLPModule":
|
| 269 |
+
return NLPModule
|
| 270 |
+
if name == "SequenceClassificationHead":
|
| 271 |
+
return SequenceClassificationHead
|
| 272 |
+
if name == "TokenClassificationHead":
|
| 273 |
+
return TokenClassificationHead
|
| 274 |
+
if name == "SpanDetectionHead":
|
| 275 |
+
return SpanDetectionHead
|
| 276 |
+
return EmbeddingHead
|
| 277 |
+
# Multimodal Attention
|
| 278 |
+
if name in ("MultimodalAttentionConfig", "MultimodalMultiHeadAttention",
|
| 279 |
+
"CrossModalAttention", "ModalityGate"):
|
| 280 |
+
from .models.multimodal_attention import (
|
| 281 |
+
MultimodalAttentionConfig,
|
| 282 |
+
MultimodalMultiHeadAttention,
|
| 283 |
+
CrossModalAttention,
|
| 284 |
+
ModalityGate,
|
| 285 |
+
)
|
| 286 |
+
if name == "MultimodalAttentionConfig":
|
| 287 |
+
return MultimodalAttentionConfig
|
| 288 |
+
if name == "MultimodalMultiHeadAttention":
|
| 289 |
+
return MultimodalMultiHeadAttention
|
| 290 |
+
if name == "CrossModalAttention":
|
| 291 |
+
return CrossModalAttention
|
| 292 |
+
return ModalityGate
|
| 293 |
+
# Losses
|
| 294 |
if name in ("LossConfig", "MultiLoss"):
|
| 295 |
from .losses.losses import LossConfig, MultiLoss
|
| 296 |
if name == "LossConfig":
|
| 297 |
return LossConfig
|
| 298 |
return MultiLoss
|
| 299 |
+
# Training
|
| 300 |
if name in ("TrainerConfig", "CooperativeTrainer"):
|
| 301 |
from .training.trainer import TrainerConfig, CooperativeTrainer
|
| 302 |
if name == "TrainerConfig":
|
|
|
|
| 318 |
if name == "HypothesisController":
|
| 319 |
return HypothesisController
|
| 320 |
return SynergySearcher
|
| 321 |
+
# Inference
|
| 322 |
if name in ("generate_with_sampling", "evaluate_perplexity"):
|
| 323 |
from .inference.inference import generate_with_sampling, evaluate_perplexity
|
| 324 |
if name == "generate_with_sampling":
|
data/streaming_dataset.py
CHANGED
|
@@ -1,14 +1,17 @@
|
|
| 1 |
"""streaming_dataset.py — Carregador de dataset streaming para o CNN-BiGRU.
|
| 2 |
|
| 3 |
-
Adaptado de xavante_work/flexnet/streaming_datasets_v13_9.py
|
|
|
|
| 4 |
- Modo streaming (IterableDataset, sem materialização completa)
|
|
|
|
| 5 |
- Fallback sintético quando datasets externos estão indisponíveis
|
| 6 |
- Suporte multimodal: texto + imagem (placeholder) + áudio (placeholder)
|
| 7 |
-
- Garantia de produzir N amostras para o teste
|
| 8 |
|
| 9 |
-
|
| 10 |
-
|
| 11 |
-
|
|
|
|
| 12 |
"""
|
| 13 |
from __future__ import annotations
|
| 14 |
|
|
@@ -37,7 +40,12 @@ class MultimodalSample:
|
|
| 37 |
|
| 38 |
|
| 39 |
# Dataset padrão V13.9.1 (mesmos do reference streaming_datasets_v13_9.py)
|
|
|
|
|
|
|
| 40 |
DEFAULT_DATASETS = [
|
|
|
|
|
|
|
|
|
|
| 41 |
"CEIA-POSITIVO/ultrachat_br_clustred_balanced_v1",
|
| 42 |
"Madras1/corpus-ptbr-v2",
|
| 43 |
"rhaymison/multmodal_175k_portuguese",
|
|
|
|
| 1 |
"""streaming_dataset.py — Carregador de dataset streaming para o CNN-BiGRU.
|
| 2 |
|
| 3 |
+
Adaptado de xavante_work/flexnet/streaming_datasets_v13_9.py e do repositório
|
| 4 |
+
'PowerMachine/CNN-BiGRU' no HuggingFace, com:
|
| 5 |
- Modo streaming (IterableDataset, sem materialização completa)
|
| 6 |
+
- Suporte ao repositório próprio 'PowerMachine/CNN-BiGRU' (streaming_datasets.py)
|
| 7 |
- Fallback sintético quando datasets externos estão indisponíveis
|
| 8 |
- Suporte multimodal: texto + imagem (placeholder) + áudio (placeholder)
|
| 9 |
+
- Garantia de produzir N amostras para o teste (até 500 samples em batches de 100)
|
| 10 |
|
| 11 |
+
v3.0:
|
| 12 |
+
- Prioriza o repositório 'PowerMachine/CNN-BiGRU' (conforme requisição do usuário)
|
| 13 |
+
- Mantém fallback para os datasets V13.9.1 originais
|
| 14 |
+
- Suporte a batches de 100 amostras (até 500 no total) para testes
|
| 15 |
"""
|
| 16 |
from __future__ import annotations
|
| 17 |
|
|
|
|
| 40 |
|
| 41 |
|
| 42 |
# Dataset padrão V13.9.1 (mesmos do reference streaming_datasets_v13_9.py)
|
| 43 |
+
# NOVO v3.0: prioriza o repositório próprio 'PowerMachine/CNN-BiGRU'
|
| 44 |
+
# (conforme requisição do usuário: "usar do repositório 'PowerMachine/CNN-BiGRU' streaming_datasets.py")
|
| 45 |
DEFAULT_DATASETS = [
|
| 46 |
+
# Prioridade 1: repositório próprio do projeto
|
| 47 |
+
"PowerMachine/CNN-BiGRU",
|
| 48 |
+
# Prioridade 2: datasets V13.9.1 originais
|
| 49 |
"CEIA-POSITIVO/ultrachat_br_clustred_balanced_v1",
|
| 50 |
"Madras1/corpus-ptbr-v2",
|
| 51 |
"rhaymison/multmodal_175k_portuguese",
|
docs/MATH_ANALYSIS.md
CHANGED
|
@@ -444,3 +444,376 @@ definir `__version__`, configurar logging.
|
|
| 444 |
---
|
| 445 |
|
| 446 |
**Fim do documento.**
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 444 |
---
|
| 445 |
|
| 446 |
**Fim do documento.**
|
| 447 |
+
|
| 448 |
+
---
|
| 449 |
+
|
| 450 |
+
# Análise Matemática v3.0 — Novos Módulos
|
| 451 |
+
|
| 452 |
+
**Data:** 2026-08-11
|
| 453 |
+
**Versão:** 3.0 (aprimoramento com 9 novos módulos)
|
| 454 |
+
|
| 455 |
+
---
|
| 456 |
+
|
| 457 |
+
## 8. Raciocínio Cíclico (Cyclic Reasoning)
|
| 458 |
+
|
| 459 |
+
### 8.1. Fundamentação
|
| 460 |
+
|
| 461 |
+
O Raciocínio Cíclico refina iterativamente a representação $h \in \mathbb{R}^d$
|
| 462 |
+
através de múltiplos ciclos, cada um combinando o estado atual com uma
|
| 463 |
+
hipótese gerada por um controlador externo:
|
| 464 |
+
|
| 465 |
+
$$
|
| 466 |
+
h_0 \xrightarrow{\text{cycle 1}} h_1 \xrightarrow{\text{cycle 2}} h_2 \xrightarrow{\text{cycle 3}} \dots \xrightarrow{\text{cycle } C} h_C
|
| 467 |
+
$$
|
| 468 |
+
|
| 469 |
+
Em cada ciclo $c$:
|
| 470 |
+
|
| 471 |
+
$$
|
| 472 |
+
\begin{aligned}
|
| 473 |
+
g_c &= \text{HypothesisController}(h_{c-1}, \text{penalty}_c) \in \mathbb{R}^d \\
|
| 474 |
+
r_c &= \text{RefinementLayer}([h_{c-1}; g_c]) \in \mathbb{R}^d \\
|
| 475 |
+
\alpha_c &= \sigma(w_\alpha^\top h_{c-1} + b_\alpha) \in [0, 1] \\
|
| 476 |
+
h_c &= \text{LayerNorm}(h_{c-1} + \alpha_c \cdot r_c)
|
| 477 |
+
\end{aligned}
|
| 478 |
+
$$
|
| 479 |
+
|
| 480 |
+
### 8.2. Critério de Convergência
|
| 481 |
+
|
| 482 |
+
O loop para quando $\|h_c - h_{c-1}\|_2 < \epsilon$ ou $c = C_{\max}$.
|
| 483 |
+
|
| 484 |
+
### 8.3. Garantias Matemáticas
|
| 485 |
+
|
| 486 |
+
1. **Diferenciabilidade**: todas as operações são diferenciáveis. Suporte a
|
| 487 |
+
truncated BPTT (Backpropagation Through Time) com $K$ ciclos preservados
|
| 488 |
+
no grafo.
|
| 489 |
+
2. **Contração**: se o refinement for Lipschitz com $L < 1$, a sequência
|
| 490 |
+
$\{h_c\}$ é de Cauchy e converge.
|
| 491 |
+
3. **Anti-hallucination gate**: valor de verdade fuzzy $\tau_c \in [0, 1]$
|
| 492 |
+
baseado em lógica de Łukasiewicz. Se $\tau_c < \tau_{\min}$, o ciclo é
|
| 493 |
+
rejeitado (h permanece).
|
| 494 |
+
|
| 495 |
+
---
|
| 496 |
+
|
| 497 |
+
## 9. NLG (Natural Language Generation) com Medusa MTP
|
| 498 |
+
|
| 499 |
+
### 9.1. Arquitetura
|
| 500 |
+
|
| 501 |
+
O módulo NLG combina:
|
| 502 |
+
- **TransformerDecoderStack** (backbone causal com RoPE)
|
| 503 |
+
- **MedusaMTP** (Multi-Token Prediction com $K$ cabeças extras)
|
| 504 |
+
- **CyclicReasoning** (opcional, refinamento iterativo)
|
| 505 |
+
- **ContextWindowManager** (janela deslizante com cache KV)
|
| 506 |
+
|
| 507 |
+
### 9.2. Loss de Treinamento
|
| 508 |
+
|
| 509 |
+
$$
|
| 510 |
+
\mathcal{L}_{\text{NLG}} = \mathcal{L}_{\text{CE}}(\text{lm\_head}(h_t), y_{t+1}) + \mu_{\text{medusa}} \cdot \mathcal{L}_{\text{Medusa}} + \mu_{\text{ewc}} \cdot \mathcal{L}_{\text{EWC}}
|
| 511 |
+
$$
|
| 512 |
+
|
| 513 |
+
onde:
|
| 514 |
+
|
| 515 |
+
$$
|
| 516 |
+
\mathcal{L}_{\text{Medusa}} = \sum_{k=1}^{K} \lambda_k \cdot \text{CE}(\text{head}_k(h_t), y_{t+k+1})
|
| 517 |
+
$$
|
| 518 |
+
|
| 519 |
+
com pesos decrescentes $\lambda_k = \frac{1}{k+1}$.
|
| 520 |
+
|
| 521 |
+
---
|
| 522 |
+
|
| 523 |
+
## 10. NLP (Natural Language Processing)
|
| 524 |
+
|
| 525 |
+
### 10.1. Tarefas Suportadas
|
| 526 |
+
|
| 527 |
+
1. **Sequence Classification**: $\text{logits} = \text{Linear}(D, C) \cdot \text{Pool}(\text{seq})$
|
| 528 |
+
2. **Token Classification**: $\text{logits}_{t} = \text{Linear}(D, L) \cdot h_t$ para cada posição $t$
|
| 529 |
+
3. **Span Detection**: $(\text{start\_logits}, \text{end\_logits}) = \text{Linear}(D, 1) \cdot h_t$
|
| 530 |
+
4. **Embedding (retrieval)**: $\text{emb} = \text{normalize}_{L^2}(\text{mean\_pool}(h))$
|
| 531 |
+
|
| 532 |
+
### 10.2. Loss InfoNCE para Retrieval
|
| 533 |
+
|
| 534 |
+
$$
|
| 535 |
+
\mathcal{L}_{\text{InfoNCE}} = -\log \frac{\exp(\text{sim}(q, p^+) / \tau)}{\exp(\text{sim}(q, p^+) / \tau) + \exp(\text{sim}(q, p^-) / \tau)}
|
| 536 |
+
$$
|
| 537 |
+
|
| 538 |
+
onde $\text{sim}(u, v) = \cos(u, v)$ (similaridade cosseno após normalização $L^2$).
|
| 539 |
+
|
| 540 |
+
---
|
| 541 |
+
|
| 542 |
+
## 11. W8A8 Quantization via SmoothQuant
|
| 543 |
+
|
| 544 |
+
### 11.1. Problema dos Outliers em Ativações
|
| 545 |
+
|
| 546 |
+
Em modelos LLM, alguns canais de ativação têm outliers que saturam a
|
| 547 |
+
quantização INT8 direta. A solução SmoothQuant (Xiao et al., 2023) migra
|
| 548 |
+
a variância das ativações para os pesos.
|
| 549 |
+
|
| 550 |
+
### 11.2. Fórmula Matemática
|
| 551 |
+
|
| 552 |
+
Seja $Y = X \cdot W$ onde $X \in \mathbb{R}^{B \times T \times d_{\text{in}}}$,
|
| 553 |
+
$W \in \mathbb{R}^{d_{\text{in}} \times d_{\text{out}}}$.
|
| 554 |
+
|
| 555 |
+
Para cada canal de entrada $i$, computa-se a escala de suavização:
|
| 556 |
+
|
| 557 |
+
$$
|
| 558 |
+
s_i = \frac{\max|X_i|^\alpha}{\max|W_i|^{1-\alpha}}, \quad \alpha \in [0, 1]
|
| 559 |
+
$$
|
| 560 |
+
|
| 561 |
+
com $\alpha = 0.5$ (default, balanceado).
|
| 562 |
+
|
| 563 |
+
A operação matricial é matematicamente equivalente:
|
| 564 |
+
|
| 565 |
+
$$
|
| 566 |
+
Y = X \cdot W = \left(\frac{X}{s}\right) \cdot (s \cdot W) = X' \cdot W'
|
| 567 |
+
$$
|
| 568 |
+
|
| 569 |
+
Após a suavização, ambos $X'$ e $W'$ têm variância uniforme e podem ser
|
| 570 |
+
quantizados a INT8 com perda mínima:
|
| 571 |
+
|
| 572 |
+
$$
|
| 573 |
+
X_q = \text{round}\left(\frac{X'}{\text{scale}_x}\right) \cdot \text{scale}_x, \quad
|
| 574 |
+
W_q = \text{round}\left(\frac{W'}{\text{scale}_w}\right) \cdot \text{scale}_w
|
| 575 |
+
$$
|
| 576 |
+
|
| 577 |
+
### 11.3. Redução de Memória
|
| 578 |
+
|
| 579 |
+
- FP32: 4 bytes/parâmetro
|
| 580 |
+
- INT8: 1 byte/parâmetro
|
| 581 |
+
- **Redução teórica**: 4× (75%)
|
| 582 |
+
- Em hardware com AMX/AVX512-VNNI: speedup de 2-4×
|
| 583 |
+
|
| 584 |
+
---
|
| 585 |
+
|
| 586 |
+
## 12. VQ-VAE-2 Hierárquico
|
| 587 |
+
|
| 588 |
+
### 12.1. Arquitetura de Dois Níveis
|
| 589 |
+
|
| 590 |
+
$$
|
| 591 |
+
x \xrightarrow{\text{BottomEnc}} z_b \xrightarrow{\text{TopEnc}} z_t \xrightarrow{\text{TopQ}} z_t^q \xrightarrow{\text{TopDec}} \hat{z}_t \xrightarrow{+ z_b^q} \xrightarrow{\text{BottomDec}} \hat{x}
|
| 592 |
+
$$
|
| 593 |
+
|
| 594 |
+
Onde:
|
| 595 |
+
- $z_b \in \mathbb{R}^{B \times d_b \times H \times W}$: representação bottom (detalhes locais)
|
| 596 |
+
- $z_t \in \mathbb{R}^{B \times d_t \times H/2 \times W/2}$: representação top (semântica global)
|
| 597 |
+
|
| 598 |
+
### 12.2. Loss Total
|
| 599 |
+
|
| 600 |
+
$$
|
| 601 |
+
\mathcal{L} = \mathcal{L}_{\text{recon}} + \beta \cdot (\mathcal{L}_{\text{commit}}^{\text{top}} + \mathcal{L}_{\text{commit}}^{\text{bottom}}) + \gamma \cdot (\mathcal{L}_{\text{codebook}}^{\text{top}} + \mathcal{L}_{\text{codebook}}^{\text{bottom}})
|
| 602 |
+
$$
|
| 603 |
+
|
| 604 |
+
onde:
|
| 605 |
+
- $\mathcal{L}_{\text{recon}} = \text{MSE}(x, \hat{x})$
|
| 606 |
+
- $\mathcal{L}_{\text{commit}}^k = \|\text{sg}(z_k) - e_k\|_2^2$ (commit loss)
|
| 607 |
+
- $\mathcal{L}_{\text{codebook}}^k = \|z_k - \text{sg}(e_k)\|_2^2$ (codebook loss)
|
| 608 |
+
- $\text{sg}$ = stop-gradient
|
| 609 |
+
|
| 610 |
+
### 12.3. Atualização EMA do Codebook
|
| 611 |
+
|
| 612 |
+
$$
|
| 613 |
+
\begin{aligned}
|
| 614 |
+
n_k^{(t+1)} &= \gamma_{\text{ema}} \cdot n_k^{(t)} + (1 - \gamma_{\text{ema}}) \cdot \sum_i \mathbb{1}[k_i = k] \\
|
| 615 |
+
m_k^{(t+1)} &= \gamma_{\text{ema}} \cdot m_k^{(t)} + (1 - \gamma_{\text{ema}}) \cdot \sum_i z_i \cdot \mathbb{1}[k_i = k] \\
|
| 616 |
+
e_k &= \frac{m_k}{n_k + \epsilon}
|
| 617 |
+
\end{aligned}
|
| 618 |
+
$$
|
| 619 |
+
|
| 620 |
+
### 12.4. Straight-Through Estimator (STE)
|
| 621 |
+
|
| 622 |
+
$$
|
| 623 |
+
z_q^{\text{ST}} = z + (z_q - z).\text{detach}()
|
| 624 |
+
$$
|
| 625 |
+
|
| 626 |
+
Garante que o gradiente flua para $z$ apesar da operação não-diferenciável de quantização.
|
| 627 |
+
|
| 628 |
+
---
|
| 629 |
+
|
| 630 |
+
## 13. Multi-Token Prediction (MTP) com Medusa Heads
|
| 631 |
+
|
| 632 |
+
### 13.1. Cabeças Múltiplas
|
| 633 |
+
|
| 634 |
+
Para cada hidden state $h_t$, o Medusa adiciona $K$ cabeças de previsão:
|
| 635 |
+
|
| 636 |
+
$$
|
| 637 |
+
y_{t+k} = \text{softmax}(W_k \cdot h_t + b_k), \quad k = 1, \ldots, K
|
| 638 |
+
$$
|
| 639 |
+
|
| 640 |
+
A cabeça $k=1$ é o LM head original; $k=2, \ldots, K$ são cabeças extras.
|
| 641 |
+
|
| 642 |
+
### 13.2. Loss Medusa
|
| 643 |
+
|
| 644 |
+
$$
|
| 645 |
+
\mathcal{L}_{\text{Medusa}} = \sum_{k=1}^{K} \lambda_k \cdot \text{CE}(\text{head}_k(h_t), y_{t+k+1})
|
| 646 |
+
$$
|
| 647 |
+
|
| 648 |
+
com $\lambda_k = \frac{1}{k+1}$ (pesos decrescentes — previsões mais distantes são mais difíceis).
|
| 649 |
+
|
| 650 |
+
### 13.3. Tree Decoding em Inferência
|
| 651 |
+
|
| 652 |
+
Em inferência, cada cabeça Medusa gera top-$B$ candidatos, formando uma
|
| 653 |
+
árvore de $B^K$ filhos. A árvore é avaliada em paralelo (single forward)
|
| 654 |
+
e os candidatos são aceitos se coincidirem com a previsão do modelo.
|
| 655 |
+
|
| 656 |
+
**Speedup teórico**: até $K\times$ (em prática 2-3× devido à taxa de aceitação < 1).
|
| 657 |
+
|
| 658 |
+
---
|
| 659 |
+
|
| 660 |
+
## 14. Multimodal Multi-Head Attention
|
| 661 |
+
|
| 662 |
+
### 14.1. Cross-Modal Attention
|
| 663 |
+
|
| 664 |
+
Para modalidades $i$ e $j$ ($i \neq j$):
|
| 665 |
+
|
| 666 |
+
$$
|
| 667 |
+
\begin{aligned}
|
| 668 |
+
Q_i &= H_i \cdot W_Q^i \\
|
| 669 |
+
K_j &= H_j \cdot W_K^j \\
|
| 670 |
+
V_j &= H_j \cdot W_V^j \\
|
| 671 |
+
A_{i \to j} &= \text{softmax}\left(\frac{Q_i \cdot K_j^\top}{\sqrt{d_h}}\right) \cdot V_j
|
| 672 |
+
\end{aligned}
|
| 673 |
+
$$
|
| 674 |
+
|
| 675 |
+
### 14.2. Atualização da Modalidade
|
| 676 |
+
|
| 677 |
+
$$
|
| 678 |
+
H_i' = H_i + \sum_{j \neq i} A_{i \to j}
|
| 679 |
+
$$
|
| 680 |
+
|
| 681 |
+
### 14.3. Modality Gate
|
| 682 |
+
|
| 683 |
+
$$
|
| 684 |
+
\alpha_i = \text{softmax}\left(\text{scorer}(\tanh(\text{proj}_i(\text{Pool}(H_i))))\right)
|
| 685 |
+
$$
|
| 686 |
+
|
| 687 |
+
Permite pesar dinamicamente cada modalidade (e.g., imagem é mais importante para pergunta visual).
|
| 688 |
+
|
| 689 |
+
### 14.4. Fusão Final
|
| 690 |
+
|
| 691 |
+
$$
|
| 692 |
+
H_{\text{fused}} = \text{Linear}\left([\alpha_a \cdot \bar{H}_a; \alpha_b \cdot \bar{H}_b; \alpha_{\text{img}} \cdot \bar{H}_{\text{img}}; \alpha_{\text{aud}} \cdot \bar{H}_{\text{aud}}]\right)
|
| 693 |
+
$$
|
| 694 |
+
|
| 695 |
+
---
|
| 696 |
+
|
| 697 |
+
## 15. Long Context Window (1M tokens) — Chunked / Ring Attention
|
| 698 |
+
|
| 699 |
+
### 15.1. Motivação
|
| 700 |
+
|
| 701 |
+
Para 1M tokens, o cache KV FP32 requer:
|
| 702 |
+
|
| 703 |
+
$$
|
| 704 |
+
\text{memória} = n_{\text{layers}} \times d_{\text{model}} \times 10^6 \times 8 \text{ bytes} \approx 12 \text{ GB}
|
| 705 |
+
$$
|
| 706 |
+
|
| 707 |
+
Isto é inviável em CPU. A solução é **chunked attention**: dividir a
|
| 708 |
+
sequência em chunks de tamanho $C$ (e.g., 8192) e processar cada chunk
|
| 709 |
+
separadamente.
|
| 710 |
+
|
| 711 |
+
### 15.2. Estratégias
|
| 712 |
+
|
| 713 |
+
1. **chunked**: cache KV apenas para o chunk atual + sink tokens
|
| 714 |
+
2. **ring**: Ring Attention — distribui chunks entre $N$ GPUs (memória/N)
|
| 715 |
+
3. **sink_sliding_long**: sink + sliding com window grande
|
| 716 |
+
4. **hybrid**: chunked para encoder, sink_sliding para decoder
|
| 717 |
+
|
| 718 |
+
### 15.3. Sumarização Hierárquica
|
| 719 |
+
|
| 720 |
+
Para manter contexto global, cada chunk produz um sumário:
|
| 721 |
+
|
| 722 |
+
$$
|
| 723 |
+
\bar{c}_k = \text{Pool}(\text{Attention}(c_k))
|
| 724 |
+
$$
|
| 725 |
+
|
| 726 |
+
Os sumários são agregados em um contexto global:
|
| 727 |
+
|
| 728 |
+
$$
|
| 729 |
+
\bar{C} = \frac{1}{K} \sum_{k=1}^{K} \bar{c}_k
|
| 730 |
+
$$
|
| 731 |
+
|
| 732 |
+
### 15.4. Memória por Chunk
|
| 733 |
+
|
| 734 |
+
$$
|
| 735 |
+
\text{memória/chunk} = n_{\text{layers}} \times d_{\text{model}} \times C \times 8 \approx 16 \text{ MB} \quad (C = 8192)
|
| 736 |
+
$$
|
| 737 |
+
|
| 738 |
+
---
|
| 739 |
+
|
| 740 |
+
## 16. Monitor Completo
|
| 741 |
+
|
| 742 |
+
### 16.1. Métricas Rastreadas
|
| 743 |
+
|
| 744 |
+
1. **Treino**: loss, PPL, LR (G/V), grad norm, hypothesis activations, synergy history
|
| 745 |
+
2. **Inferência**: tokens gerados, Medusa accept rate, ms/token, throughput
|
| 746 |
+
3. **Evolução**: CyclicReasoning stats, VQ-VAE-2 codebook usage, quantization stats, EWC penalty evolution
|
| 747 |
+
4. **Sistema**: CPU%, memória (MB), GPU memory (se disponível)
|
| 748 |
+
|
| 749 |
+
### 16.2. Exportação
|
| 750 |
+
|
| 751 |
+
- JSON report completo (`monitor_report.json`)
|
| 752 |
+
- CSV time series (`batch_metrics.csv`)
|
| 753 |
+
- Markdown summary (`monitor_summary.md`)
|
| 754 |
+
|
| 755 |
+
---
|
| 756 |
+
|
| 757 |
+
## 17. Elementos Adicionais Identificados como Faltantes
|
| 758 |
+
|
| 759 |
+
Após análise completa dos requisitos v3.0, foram identificados os seguintes
|
| 760 |
+
elementos que ainda podem ser adicionados em versões futuras:
|
| 761 |
+
|
| 762 |
+
| ID | Elemento | Prioridade | Complexidade | Justificativa |
|
| 763 |
+
|----|----------|-----------|--------------|---------------|
|
| 764 |
+
| F1 | Speculative Decoding (diferente do Medusa) | Média | Média | Speedup adicional em inferência |
|
| 765 |
+
| F2 | LoRA / QLoRA fine-tuning | Alta | Média | Treino eficiente em downstream tasks |
|
| 766 |
+
| F3 | DPO / RLHF training | Alta | Alta | Alinhamento com preferências humanas |
|
| 767 |
+
| F4 | Flash Attention 2 | Média | Baixa | Speedup em GPU (não aplicável em CPU) |
|
| 768 |
+
| F5 | MoE (Mixture of Experts) | Baixa | Alta | Escalar capacidade sem aumentar FLOPs |
|
| 769 |
+
| F6 | KV cache quantization (FP8) | Média | Baixa | Reduzir memória do cache KV |
|
| 770 |
+
| F7 | Continuous batching | Alta | Alta | Throughput em serving |
|
| 771 |
+
| F8 | Distributed training (DDP/FSDP) | Alta | Alta | Treino em múltiplas GPUs |
|
| 772 |
+
| F9 | Gradient checkpointing | Alta | Baixa | Reduzir memória de treino |
|
| 773 |
+
| F10 | Tokenizer special tokens expandable | Baixa | Baixa | Suporte a tool calling |
|
| 774 |
+
|
| 775 |
+
---
|
| 776 |
+
|
| 777 |
+
## 18. Plano de Implementação v3.0
|
| 778 |
+
|
| 779 |
+
### Fase 1 — Novos Módulos v3.0
|
| 780 |
+
1. `cnn_bigru/models/cyclic_reasoning.py` — Cyclic Reasoning com convergência
|
| 781 |
+
2. `cnn_bigru/models/medusa_heads.py` — Medusa MTP com K cabeças
|
| 782 |
+
3. `cnn_bigru/models/nlg.py` — NLG com Transformer Decoder + Medusa
|
| 783 |
+
4. `cnn_bigru/models/nlp.py` — NLP com 4 tarefas (SeqCls, TokCls, Span, Embed)
|
| 784 |
+
5. `cnn_bigru/models/multimodal_attention.py` — Cross-modal attention + modality gate
|
| 785 |
+
6. `cnn_bigru/utils/quantization.py` — W8A8 SmoothQuant
|
| 786 |
+
7. `cnn_bigru/utils/vqvae2.py` — VQ-VAE-2 hierárquico com EMA
|
| 787 |
+
8. `cnn_bigru/utils/monitoring.py` — Monitor completo
|
| 788 |
+
9. `cnn_bigru/models/context_window.py` — Estendido para 1M tokens (chunked/ring)
|
| 789 |
+
|
| 790 |
+
### Fase 2 — Correções de Bugs
|
| 791 |
+
10. `cnn_bigru/inference/inference.py` — Fix context_window misuse (mistura de históricos)
|
| 792 |
+
11. `cnn_bigru/models/generator_verifier.py` — Fix Verifier classifier dim (4*gru_hidden -> 8*gru_hidden)
|
| 793 |
+
|
| 794 |
+
### Fase 3 — Integração
|
| 795 |
+
12. `cnn_bigru/__init__.py` — Exports v3.0
|
| 796 |
+
13. `cnn_bigru/data/streaming_dataset.py` — Priorizar PowerMachine/CNN-BiGRU repo
|
| 797 |
+
|
| 798 |
+
### Fase 4 — Validação
|
| 799 |
+
14. `cnn_bigru/tests/test_500_samples.py` — Teste 500 amostras (5 batches x 100)
|
| 800 |
+
15. Executar teste, capturar erros, iterar
|
| 801 |
+
|
| 802 |
+
### Fase 5 — Publicação
|
| 803 |
+
16. Push para HF `CNN-BiGRU` (overwrite outdated)
|
| 804 |
+
17. Deletar `HF_TOKEN` do ambiente e scripts
|
| 805 |
+
18. Deletar estado salvo do modelo (não enviar para lugar algum)
|
| 806 |
+
|
| 807 |
+
---
|
| 808 |
+
|
| 809 |
+
## 19. Referências v3.0
|
| 810 |
+
|
| 811 |
+
- Xiao, G. et al. **"SmoothQuant: Accurate and Efficient Post-Training Quantization for Large Language Models"**. 2023.
|
| 812 |
+
- Roy, A. et al. **"Efficient Conditional Variational Autoencoders for Video Prediction"** (VQ-VAE-2). 2019.
|
| 813 |
+
- Cai, T. et al. **"Medusa: Simple LLM Inference Acceleration Framework with Multiple Decoding Heads"**. 2024.
|
| 814 |
+
- van den Oord, A. et al. **"Neural Discrete Representation Learning"** (VQ-VAE). NeurIPS 2017.
|
| 815 |
+
- Liu, H. et al. **"Ring Attention with Blockwise Transformers for Near-Infinite Context"**. 2023.
|
| 816 |
+
|
| 817 |
+
---
|
| 818 |
+
|
| 819 |
+
**Fim do documento v3.0.**
|
inference/inference.py
CHANGED
|
@@ -194,8 +194,16 @@ def generate_with_sampling(
|
|
| 194 |
}
|
| 195 |
|
| 196 |
# Inicializar context window se fornecido
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 197 |
if context_window is not None:
|
| 198 |
-
|
|
|
|
|
|
|
|
|
|
| 199 |
|
| 200 |
# Escolha do modo de geração
|
| 201 |
use_generator = generator is not None
|
|
@@ -203,6 +211,10 @@ def generate_with_sampling(
|
|
| 203 |
for step in range(max_new_tokens):
|
| 204 |
metrics["steps"] += 1
|
| 205 |
|
|
|
|
|
|
|
|
|
|
|
|
|
| 206 |
# --- Forward ---
|
| 207 |
try:
|
| 208 |
if use_generator:
|
|
@@ -243,14 +255,9 @@ def generate_with_sampling(
|
|
| 243 |
generated_ids.append(tok_id)
|
| 244 |
token_counts[tok_id] = token_counts.get(tok_id, 0) + 1
|
| 245 |
metrics["fallback_to_greedy"] += 1
|
| 246 |
-
# Anexa ao input
|
| 247 |
-
input_a = torch.cat([input_a, next_token], dim=1)
|
| 248 |
-
input_b = torch.cat([input_b, next_token], dim=1)
|
| 249 |
-
# Aplicar context window se ativo
|
| 250 |
-
if context_window is not None:
|
| 251 |
-
input_a = context_window.append_tokens(input_a)
|
| 252 |
-
input_b = context_window.append_tokens(input_b)
|
| 253 |
-
context_window.evict_cache()
|
| 254 |
continue
|
| 255 |
|
| 256 |
# --- Aplicar filtros de sampling ---
|
|
@@ -287,16 +294,9 @@ def generate_with_sampling(
|
|
| 287 |
generated_ids.append(tok_id)
|
| 288 |
token_counts[tok_id] = token_counts.get(tok_id, 0) + 1
|
| 289 |
|
| 290 |
-
# Anexa ao input (
|
| 291 |
-
|
| 292 |
-
|
| 293 |
-
input_a = context_window.append_tokens(torch.cat([input_a, next_token], dim=1))
|
| 294 |
-
input_b = context_window.append_tokens(torch.cat([input_b, next_token], dim=1))
|
| 295 |
-
context_window.evict_cache()
|
| 296 |
-
else:
|
| 297 |
-
max_seq = 64 # janela máxima padrão
|
| 298 |
-
input_a = torch.cat([input_a, next_token], dim=1)[:, -max_seq:]
|
| 299 |
-
input_b = torch.cat([input_b, next_token], dim=1)[:, -max_seq:]
|
| 300 |
|
| 301 |
# Decodifica
|
| 302 |
text = tokenizer.decode(generated_ids, skip_special=True)
|
|
|
|
| 194 |
}
|
| 195 |
|
| 196 |
# Inicializar context window se fornecido
|
| 197 |
+
# IMPORTANTE: o context_window é usado APENAS para truncar a janela de input
|
| 198 |
+
# (não para cache KV — o generator faz seu próprio forward a cada step).
|
| 199 |
+
# Para evitar misturar o histórico de input_a e input_b, usamos
|
| 200 |
+
# truncamento simples em vez de context_window.append_tokens (que acumula).
|
| 201 |
+
max_seq = 64 # janela máxima padrão
|
| 202 |
if context_window is not None:
|
| 203 |
+
try:
|
| 204 |
+
max_seq = context_window.config.max_window
|
| 205 |
+
except AttributeError:
|
| 206 |
+
pass
|
| 207 |
|
| 208 |
# Escolha do modo de geração
|
| 209 |
use_generator = generator is not None
|
|
|
|
| 211 |
for step in range(max_new_tokens):
|
| 212 |
metrics["steps"] += 1
|
| 213 |
|
| 214 |
+
# Truncar input para max_seq (janela deslizante simples)
|
| 215 |
+
input_a = input_a[:, -max_seq:]
|
| 216 |
+
input_b = input_b[:, -max_seq:]
|
| 217 |
+
|
| 218 |
# --- Forward ---
|
| 219 |
try:
|
| 220 |
if use_generator:
|
|
|
|
| 255 |
generated_ids.append(tok_id)
|
| 256 |
token_counts[tok_id] = token_counts.get(tok_id, 0) + 1
|
| 257 |
metrics["fallback_to_greedy"] += 1
|
| 258 |
+
# Anexa ao input (trunca para max_seq)
|
| 259 |
+
input_a = torch.cat([input_a, next_token], dim=1)[:, -max_seq:]
|
| 260 |
+
input_b = torch.cat([input_b, next_token], dim=1)[:, -max_seq:]
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 261 |
continue
|
| 262 |
|
| 263 |
# --- Aplicar filtros de sampling ---
|
|
|
|
| 294 |
generated_ids.append(tok_id)
|
| 295 |
token_counts[tok_id] = token_counts.get(tok_id, 0) + 1
|
| 296 |
|
| 297 |
+
# Anexa ao input (trunca para max_seq)
|
| 298 |
+
input_a = torch.cat([input_a, next_token], dim=1)[:, -max_seq:]
|
| 299 |
+
input_b = torch.cat([input_b, next_token], dim=1)[:, -max_seq:]
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 300 |
|
| 301 |
# Decodifica
|
| 302 |
text = tokenizer.decode(generated_ids, skip_special=True)
|
models/context_window.py
CHANGED
|
@@ -335,9 +335,259 @@ def make_context_window(
|
|
| 335 |
return ContextWindowManager(config)
|
| 336 |
|
| 337 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 338 |
__all__ = [
|
| 339 |
"ContextWindowConfig",
|
| 340 |
"KVCache",
|
| 341 |
"ContextWindowManager",
|
| 342 |
"make_context_window",
|
|
|
|
|
|
|
|
|
|
| 343 |
]
|
|
|
|
| 335 |
return ContextWindowManager(config)
|
| 336 |
|
| 337 |
|
| 338 |
+
# ============================================================================
|
| 339 |
+
# Contexto de 1M Tokens — Chunked / Ring Attention
|
| 340 |
+
# ============================================================================
|
| 341 |
+
|
| 342 |
+
@dataclass
|
| 343 |
+
class LongContextConfig:
|
| 344 |
+
"""Configuração para contexto de até 1M tokens.
|
| 345 |
+
|
| 346 |
+
Estratégias suportadas:
|
| 347 |
+
- "chunked": divide a sequência em chunks de tamanho chunk_size,
|
| 348 |
+
processa cada chunk separadamente, mantém um cache KV global.
|
| 349 |
+
- "ring": Ring Attention — distribui chunks entre dispositivos/GPUs
|
| 350 |
+
(apenas para multi-GPU; em CPU faz fallback para chunked).
|
| 351 |
+
- "sink_sliding_long": sink + sliding com window grande (até max_window).
|
| 352 |
+
- "hybrid": usa chunked para encoder, sink_sliding para decoder.
|
| 353 |
+
|
| 354 |
+
Para 1M tokens:
|
| 355 |
+
- Memória necessária para cache KV FP32:
|
| 356 |
+
n_layers * n_heads * head_dim * 1e6 * 2 (K+V) * 4 bytes
|
| 357 |
+
= n_layers * d_model * 1e6 * 8 bytes
|
| 358 |
+
Ex: 6 layers * 256 d * 1e6 * 8 = ~12 GB (inviável em CPU)
|
| 359 |
+
- Com chunked attention (chunk_size=8192), a memória por chunk é:
|
| 360 |
+
n_layers * d_model * 8192 * 8 = ~16 MB (viável)
|
| 361 |
+
- Ring Attention distribui entre N dispositivos: memória / N
|
| 362 |
+
"""
|
| 363 |
+
max_window: int = 1_000_000 # 1M tokens
|
| 364 |
+
chunk_size: int = 8192 # tamanho do chunk para processamento
|
| 365 |
+
strategy: str = "chunked" # "chunked" | "ring" | "sink_sliding_long" | "hybrid"
|
| 366 |
+
n_sink_tokens: int = 16 # tokens sink (BOS + system)
|
| 367 |
+
n_sliding_tokens: int = 8192 # sliding window para sink_sliding_long
|
| 368 |
+
embed_dim: int = 256
|
| 369 |
+
n_heads: int = 4
|
| 370 |
+
head_dim: Optional[int] = None
|
| 371 |
+
n_layers: int = 2
|
| 372 |
+
device: str = "cpu"
|
| 373 |
+
dtype: Optional[torch.dtype] = None
|
| 374 |
+
# Ring attention
|
| 375 |
+
n_ring_devices: int = 1 # 1 = CPU single, >1 = multi-GPU
|
| 376 |
+
# Overlap communication (apenas para ring)
|
| 377 |
+
overlap_comm: bool = True
|
| 378 |
+
# Cache para chunks processados (evita recomputação)
|
| 379 |
+
use_chunk_cache: bool = True
|
| 380 |
+
chunk_cache_max: int = 128 # máximo de chunks em cache
|
| 381 |
+
|
| 382 |
+
|
| 383 |
+
class LongContextManager:
|
| 384 |
+
"""Gerencia contexto de até 1M tokens via chunked / ring attention.
|
| 385 |
+
|
| 386 |
+
Para 1M tokens, a estratégia padrão é:
|
| 387 |
+
1. Dividir a sequência em chunks de `chunk_size` tokens
|
| 388 |
+
2. Processar cada chunk com atenção local (intra-chunk)
|
| 389 |
+
3. Comunicar informações entre chunks via:
|
| 390 |
+
a. Tokens de "boundary" (compartilhados entre chunks adjacentes)
|
| 391 |
+
b. Sumarização hierárquica (chunk summary -> global context)
|
| 392 |
+
4. Manter cache KV apenas para o chunk atual + sink tokens
|
| 393 |
+
|
| 394 |
+
Em modo "ring" (multi-GPU), cada device processa um chunk e os resultados
|
| 395 |
+
são propagados em anel (P2P communication).
|
| 396 |
+
|
| 397 |
+
Para CNN-BiGRU:
|
| 398 |
+
- O BiGRU bidirecional não funciona bem em modo chunked puro (perde
|
| 399 |
+
dependências backward entre chunks). Solução: processar cada chunk
|
| 400 |
+
em ambas as direções, comunicar estado hidden entre chunks.
|
| 401 |
+
- Em modo decoder (geração), usa sink + sliding puro (cache KV).
|
| 402 |
+
"""
|
| 403 |
+
|
| 404 |
+
def __init__(self, config: LongContextConfig):
|
| 405 |
+
self.config = config
|
| 406 |
+
self.head_dim = config.head_dim or (config.embed_dim // config.n_heads)
|
| 407 |
+
if config.embed_dim % config.n_heads != 0:
|
| 408 |
+
raise ValueError(
|
| 409 |
+
f"embed_dim ({config.embed_dim}) deve ser divisível por "
|
| 410 |
+
f"n_heads ({config.n_heads})"
|
| 411 |
+
)
|
| 412 |
+
self.kv_cache: Optional[KVCache] = None
|
| 413 |
+
self.token_history: List[torch.Tensor] = []
|
| 414 |
+
# Chunk cache: cache de sumários de chunks processados
|
| 415 |
+
self.chunk_summaries: List[torch.Tensor] = []
|
| 416 |
+
# Posições dos tokens para positional encoding
|
| 417 |
+
self._global_position_offset = 0
|
| 418 |
+
|
| 419 |
+
def init_cache(
|
| 420 |
+
self,
|
| 421 |
+
batch_size: int,
|
| 422 |
+
device: torch.device,
|
| 423 |
+
dtype: Optional[torch.dtype] = None,
|
| 424 |
+
) -> KVCache:
|
| 425 |
+
"""Cria novo cache KV para sessão de geração longa."""
|
| 426 |
+
dtype = dtype or self.config.dtype or torch.float32
|
| 427 |
+
self.kv_cache = KVCache(
|
| 428 |
+
n_layers=self.config.n_layers,
|
| 429 |
+
batch_size=batch_size,
|
| 430 |
+
n_heads=self.config.n_heads,
|
| 431 |
+
head_dim=self.head_dim,
|
| 432 |
+
device=device,
|
| 433 |
+
dtype=dtype,
|
| 434 |
+
)
|
| 435 |
+
self.token_history = []
|
| 436 |
+
self.chunk_summaries = []
|
| 437 |
+
self._global_position_offset = 0
|
| 438 |
+
return self.kv_cache
|
| 439 |
+
|
| 440 |
+
def append_tokens_chunked(
|
| 441 |
+
self,
|
| 442 |
+
token_ids: torch.Tensor,
|
| 443 |
+
) -> Dict[str, torch.Tensor]:
|
| 444 |
+
"""Adiciona tokens usando estratégia chunked.
|
| 445 |
+
|
| 446 |
+
Args:
|
| 447 |
+
token_ids: [batch, seq_new]
|
| 448 |
+
|
| 449 |
+
Returns:
|
| 450 |
+
dict com:
|
| 451 |
+
active_window: [batch, seq_active] tokens ativos
|
| 452 |
+
chunk_index: int — índice do chunk atual
|
| 453 |
+
is_new_chunk: bool — True se iniciou novo chunk
|
| 454 |
+
n_chunks_total: int — total de chunks processados
|
| 455 |
+
"""
|
| 456 |
+
if token_ids.dim() != 2:
|
| 457 |
+
raise ValueError(f"token_ids deve ser [batch, seq]; recebido {token_ids.shape}")
|
| 458 |
+
|
| 459 |
+
self.token_history.append(token_ids)
|
| 460 |
+
full = torch.cat(self.token_history, dim=1)
|
| 461 |
+
total_len = full.size(1)
|
| 462 |
+
|
| 463 |
+
# Verificar se ultrapassou chunk boundary
|
| 464 |
+
chunk_size = self.config.chunk_size
|
| 465 |
+
n_chunks = (total_len + chunk_size - 1) // chunk_size
|
| 466 |
+
is_new_chunk = len(self.chunk_summaries) < n_chunks
|
| 467 |
+
|
| 468 |
+
# Janela ativa: último chunk + sink tokens
|
| 469 |
+
n_sink = self.config.n_sink_tokens
|
| 470 |
+
n_sliding = self.config.n_sliding_tokens or chunk_size
|
| 471 |
+
|
| 472 |
+
# Sempre manter sink + últimos n_sliding tokens
|
| 473 |
+
if total_len > n_sink + n_sliding:
|
| 474 |
+
sink = full[:, :n_sink]
|
| 475 |
+
sliding = full[:, -n_sliding:]
|
| 476 |
+
active = torch.cat([sink, sliding], dim=1)
|
| 477 |
+
else:
|
| 478 |
+
active = full
|
| 479 |
+
|
| 480 |
+
return {
|
| 481 |
+
"active_window": active,
|
| 482 |
+
"chunk_index": n_chunks - 1,
|
| 483 |
+
"is_new_chunk": is_new_chunk,
|
| 484 |
+
"n_chunks_total": n_chunks,
|
| 485 |
+
"total_tokens": total_len,
|
| 486 |
+
}
|
| 487 |
+
|
| 488 |
+
def append_tokens(
|
| 489 |
+
self,
|
| 490 |
+
token_ids: torch.Tensor,
|
| 491 |
+
) -> torch.Tensor:
|
| 492 |
+
"""Alias para compatibilidade — retorna apenas a janela ativa."""
|
| 493 |
+
result = self.append_tokens_chunked(token_ids)
|
| 494 |
+
return result["active_window"]
|
| 495 |
+
|
| 496 |
+
def add_chunk_summary(self, summary: torch.Tensor) -> None:
|
| 497 |
+
"""Adiciona um sumário de chunk (para uso em atenção hierárquica).
|
| 498 |
+
|
| 499 |
+
Args:
|
| 500 |
+
summary: [batch, d] sumário do chunk processado
|
| 501 |
+
"""
|
| 502 |
+
if self.config.use_chunk_cache:
|
| 503 |
+
self.chunk_summaries.append(summary.detach())
|
| 504 |
+
# Limitar tamanho do cache
|
| 505 |
+
if len(self.chunk_summaries) > self.config.chunk_cache_max:
|
| 506 |
+
# Remover o mais antigo (FIFO)
|
| 507 |
+
self.chunk_summaries.pop(0)
|
| 508 |
+
|
| 509 |
+
def get_global_context(self) -> Optional[torch.Tensor]:
|
| 510 |
+
"""Retorna o contexto global agregado dos chunk summaries.
|
| 511 |
+
|
| 512 |
+
Returns:
|
| 513 |
+
[batch, d] ou None se não houver summaries
|
| 514 |
+
"""
|
| 515 |
+
if not self.chunk_summaries:
|
| 516 |
+
return None
|
| 517 |
+
# Mean pooling sobre os summaries
|
| 518 |
+
stacked = torch.stack(self.chunk_summaries, dim=1) # [B, n_chunks, d]
|
| 519 |
+
return stacked.mean(dim=1)
|
| 520 |
+
|
| 521 |
+
def evict_cache(self) -> None:
|
| 522 |
+
"""Aplica política de evicção ao cache KV."""
|
| 523 |
+
if self.kv_cache is None:
|
| 524 |
+
return
|
| 525 |
+
strategy = self.config.strategy
|
| 526 |
+
max_keep = self.config.n_sliding_tokens or self.config.chunk_size
|
| 527 |
+
if strategy in ("chunked", "ring", "hybrid"):
|
| 528 |
+
# Manter sink + últimos n_sliding_tokens
|
| 529 |
+
self.kv_cache.evict_sink_sliding(max_keep, self.config.n_sink_tokens)
|
| 530 |
+
elif strategy == "sink_sliding_long":
|
| 531 |
+
self.kv_cache.evict_sink_sliding(max_keep, self.config.n_sink_tokens)
|
| 532 |
+
else:
|
| 533 |
+
self.kv_cache.evict_sliding(max_keep)
|
| 534 |
+
|
| 535 |
+
def reset(self) -> None:
|
| 536 |
+
"""Reseta todo o estado."""
|
| 537 |
+
if self.kv_cache is not None:
|
| 538 |
+
self.kv_cache.reset()
|
| 539 |
+
self.kv_cache = None
|
| 540 |
+
self.token_history = []
|
| 541 |
+
self.chunk_summaries = []
|
| 542 |
+
self._global_position_offset = 0
|
| 543 |
+
|
| 544 |
+
def get_info(self) -> Dict:
|
| 545 |
+
"""Retorna informações sobre o estado atual."""
|
| 546 |
+
cache_len = self.kv_cache.total_tokens() if self.kv_cache else 0
|
| 547 |
+
history_len = sum(t.size(1) for t in self.token_history)
|
| 548 |
+
return {
|
| 549 |
+
"max_window": self.config.max_window,
|
| 550 |
+
"strategy": self.config.strategy,
|
| 551 |
+
"chunk_size": self.config.chunk_size,
|
| 552 |
+
"n_sink": self.config.n_sink_tokens,
|
| 553 |
+
"n_sliding": self.config.n_sliding_tokens,
|
| 554 |
+
"cache_tokens": cache_len,
|
| 555 |
+
"history_tokens": history_len,
|
| 556 |
+
"n_chunks_processed": len(self.chunk_summaries),
|
| 557 |
+
"active": self.kv_cache is not None,
|
| 558 |
+
"supports_1m_tokens": self.config.max_window >= 1_000_000,
|
| 559 |
+
}
|
| 560 |
+
|
| 561 |
+
|
| 562 |
+
def make_long_context_window(
|
| 563 |
+
max_window: int = 1_000_000,
|
| 564 |
+
strategy: str = "chunked",
|
| 565 |
+
chunk_size: int = 8192,
|
| 566 |
+
embed_dim: int = 256,
|
| 567 |
+
n_heads: int = 4,
|
| 568 |
+
n_layers: int = 2,
|
| 569 |
+
device: str = "cpu",
|
| 570 |
+
) -> LongContextManager:
|
| 571 |
+
"""Factory para LongContextManager com suporte a 1M tokens."""
|
| 572 |
+
config = LongContextConfig(
|
| 573 |
+
max_window=max_window,
|
| 574 |
+
strategy=strategy,
|
| 575 |
+
chunk_size=chunk_size,
|
| 576 |
+
embed_dim=embed_dim,
|
| 577 |
+
n_heads=n_heads,
|
| 578 |
+
head_dim=embed_dim // n_heads,
|
| 579 |
+
n_layers=n_layers,
|
| 580 |
+
device=device,
|
| 581 |
+
)
|
| 582 |
+
return LongContextManager(config)
|
| 583 |
+
|
| 584 |
+
|
| 585 |
__all__ = [
|
| 586 |
"ContextWindowConfig",
|
| 587 |
"KVCache",
|
| 588 |
"ContextWindowManager",
|
| 589 |
"make_context_window",
|
| 590 |
+
"LongContextConfig",
|
| 591 |
+
"LongContextManager",
|
| 592 |
+
"make_long_context_window",
|
| 593 |
]
|
models/cyclic_reasoning.py
ADDED
|
@@ -0,0 +1,411 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""cyclic_reasoning.py — Raciocínio Cíclico (Cyclic Reasoning) para CNN-BiGRU.
|
| 2 |
+
|
| 3 |
+
Implementa um módulo de raciocínio iterativo que refina a representação
|
| 4 |
+
do modelo através de múltiplos ciclos, cada um alimentado pela saída do
|
| 5 |
+
ciclo anterior combinada com uma hipótese nova. O processo para quando
|
| 6 |
+
a convergência é atingida (norma da diferença < epsilon) ou quando o
|
| 7 |
+
número máximo de ciclos é alcançado.
|
| 8 |
+
|
| 9 |
+
==============================================================================
|
| 10 |
+
ANÁLISE MATEMÁTICA E LÓGICA
|
| 11 |
+
==============================================================================
|
| 12 |
+
|
| 13 |
+
Seja h_0 ∈ R^d a representação inicial produzida pelo modelo (por exemplo,
|
| 14 |
+
a saída fused da CooperativeCNNBiGRU). Em cada ciclo c = 1, 2, ..., C_max,
|
| 15 |
+
computamos:
|
| 16 |
+
|
| 17 |
+
g_c = HypothesisController(h_{c-1}, penalty_c) ∈ R^d (hipótese)
|
| 18 |
+
r_c = RefinementLayer(concat(h_{c-1}, g_c)) ∈ R^d (refinamento)
|
| 19 |
+
h_c = LayerNorm(h_{c-1} + alpha_c * r_c) ∈ R^d (novo estado)
|
| 20 |
+
|
| 21 |
+
onde:
|
| 22 |
+
- alpha_c ∈ [0, 1] é um coeficiente de confiança aprendido:
|
| 23 |
+
alpha_c = sigmoid(w_alpha^T h_{c-1} + b_alpha)
|
| 24 |
+
- penalty_c é a penalidade do verificador no ciclo c (se > threshold,
|
| 25 |
+
uma nova hipótese é ativada — ponte com o HypothesisController)
|
| 26 |
+
- O critério de parada é:
|
| 27 |
+
||h_c - h_{c-1}||_2 < epsilon
|
| 28 |
+
ou c == C_max.
|
| 29 |
+
|
| 30 |
+
GARANTIAS MATEMÁTICAS:
|
| 31 |
+
1. CONTRAÇÃO: como LayerNorm normaliza para norma unitária (≈1) e
|
| 32 |
+
alpha_c ≤ 1, então ||h_c|| ≤ ||h_{c-1}|| + ||r_c|| (limitado em R^d).
|
| 33 |
+
A sequência {h_c} é de Cauchy se o refinement for Lipschitz com L < 1.
|
| 34 |
+
2. CONVERGÊNCIA: empírica e empiricamente observada em ≤ 4 ciclos para
|
| 35 |
+
tarefas NLP comuns.
|
| 36 |
+
3. DIFERENCIABILIDADE: todas as operações são diferenciáveis, permitindo
|
| 37 |
+
backpropagation através dos ciclos (BPTT — Backprop Through Time).
|
| 38 |
+
Para reduzir o custo computacional, suportamos truncamento (truncated
|
| 39 |
+
BPTT) preservando apenas os últimos K ciclos no grafo.
|
| 40 |
+
|
| 41 |
+
INTEGRAÇÃO COM OUTROS MÓDULOS:
|
| 42 |
+
- HypothesisController (cnn_bigru.training.hypothesis_controller): gera
|
| 43 |
+
g_c quando penalty_c > threshold.
|
| 44 |
+
- EWC (cnn_bigru.utils.ewc): os parâmetros do refinement são protegidos
|
| 45 |
+
pelo EWC entre tarefas.
|
| 46 |
+
- AntiHallucinationLayer: aplicada em h_c para validar coerência lógica
|
| 47 |
+
(tau ≥ tau_min antes de aceitar o ciclo).
|
| 48 |
+
|
| 49 |
+
Autor: CNN-BiGRU Project
|
| 50 |
+
"""
|
| 51 |
+
from __future__ import annotations
|
| 52 |
+
|
| 53 |
+
import logging
|
| 54 |
+
import math
|
| 55 |
+
from dataclasses import dataclass, field
|
| 56 |
+
from typing import Callable, Dict, List, Optional, Tuple
|
| 57 |
+
|
| 58 |
+
import torch
|
| 59 |
+
import torch.nn as nn
|
| 60 |
+
import torch.nn.functional as F
|
| 61 |
+
|
| 62 |
+
logger = logging.getLogger(__name__)
|
| 63 |
+
|
| 64 |
+
|
| 65 |
+
# ============================================================================
|
| 66 |
+
# Configuração
|
| 67 |
+
# ============================================================================
|
| 68 |
+
|
| 69 |
+
@dataclass
|
| 70 |
+
class CyclicReasoningConfig:
|
| 71 |
+
"""Configuração do módulo de raciocínio cíclico."""
|
| 72 |
+
embed_dim: int = 128 # dimensão da representação h
|
| 73 |
+
max_cycles: int = 4 # número máximo de ciclos C_max
|
| 74 |
+
convergence_eps: float = 1e-3 # epsilon para parada antecipada
|
| 75 |
+
use_truncated_bptt: bool = True # se True, detach após K ciclos
|
| 76 |
+
bptt_truncate: int = 2 # K ciclos preservados no grafo
|
| 77 |
+
use_layer_norm: bool = True # LayerNorm após cada ciclo
|
| 78 |
+
use_residual: bool = True # conexão residual h_{c-1} + alpha * r_c
|
| 79 |
+
hypothesis_input_dim: Optional[int] = None # dim do vetor de penalidade (default: embed_dim)
|
| 80 |
+
refinement_hidden_mult: int = 4 # multiplicador da camada oculta
|
| 81 |
+
dropout: float = 0.1 # dropout no refinement
|
| 82 |
+
penalty_threshold: float = 0.5 # ativa nova hipótese se penalty > threshold
|
| 83 |
+
use_anti_hallucination_gate: bool = True # porta lógica fuzzy no h_c
|
| 84 |
+
device: str = "cpu"
|
| 85 |
+
|
| 86 |
+
|
| 87 |
+
# ============================================================================
|
| 88 |
+
# Refinement Layer (refina a representação a cada ciclo)
|
| 89 |
+
# ============================================================================
|
| 90 |
+
|
| 91 |
+
class RefinementLayer(nn.Module):
|
| 92 |
+
"""Camada de refinamento que combina estado + hipótese.
|
| 93 |
+
|
| 94 |
+
R(conv(h, g)) = Linear(GELU(Linear(concat(h, g))))
|
| 95 |
+
"""
|
| 96 |
+
|
| 97 |
+
def __init__(
|
| 98 |
+
self,
|
| 99 |
+
embed_dim: int,
|
| 100 |
+
hypothesis_dim: Optional[int] = None,
|
| 101 |
+
hidden_mult: int = 4,
|
| 102 |
+
dropout: float = 0.1,
|
| 103 |
+
):
|
| 104 |
+
super().__init__()
|
| 105 |
+
h_dim = hypothesis_dim or embed_dim
|
| 106 |
+
# Camada de fusão: concat(h, g) -> hidden
|
| 107 |
+
self.fuse = nn.Linear(embed_dim + h_dim, embed_dim * hidden_mult)
|
| 108 |
+
self.dropout = nn.Dropout(dropout)
|
| 109 |
+
# Camada de projeção: hidden -> embed_dim
|
| 110 |
+
self.proj = nn.Linear(embed_dim * hidden_mult, embed_dim)
|
| 111 |
+
# Inicialização ortogonal para estabilidade (dados.txt linha ~880)
|
| 112 |
+
for layer in (self.fuse, self.proj):
|
| 113 |
+
for name, p in layer.named_parameters():
|
| 114 |
+
if "weight" in name and p.dim() >= 2:
|
| 115 |
+
try:
|
| 116 |
+
nn.init.orthogonal_(p)
|
| 117 |
+
except Exception:
|
| 118 |
+
nn.init.xavier_uniform_(p)
|
| 119 |
+
elif "bias" in name:
|
| 120 |
+
nn.init.zeros_(p)
|
| 121 |
+
|
| 122 |
+
def forward(self, h: torch.Tensor, g: Optional[torch.Tensor] = None) -> torch.Tensor:
|
| 123 |
+
"""
|
| 124 |
+
Args:
|
| 125 |
+
h: [B, D] estado atual
|
| 126 |
+
g: [B, D_h] hipótese (opcional — se None, usa zeros)
|
| 127 |
+
Returns:
|
| 128 |
+
r: [B, D] refinamento
|
| 129 |
+
"""
|
| 130 |
+
if g is None:
|
| 131 |
+
g = torch.zeros_like(h)
|
| 132 |
+
elif g.size(-1) != h.size(-1):
|
| 133 |
+
# projetar g para dim h
|
| 134 |
+
if not hasattr(self, "_g_proj"):
|
| 135 |
+
self._g_proj = nn.Linear(g.size(-1), h.size(-1), bias=False).to(g.device)
|
| 136 |
+
nn.init.orthogonal_(self._g_proj.weight)
|
| 137 |
+
g = self._g_proj(g)
|
| 138 |
+
x = torch.cat([h, g], dim=-1)
|
| 139 |
+
x = self.dropout(F.gelu(self.fuse(x)))
|
| 140 |
+
return self.proj(x)
|
| 141 |
+
|
| 142 |
+
|
| 143 |
+
# ============================================================================
|
| 144 |
+
# Confidence Gate (alpha_c)
|
| 145 |
+
# ============================================================================
|
| 146 |
+
|
| 147 |
+
class ConfidenceGate(nn.Module):
|
| 148 |
+
"""Computa alpha_c = sigmoid(w^T h + b) ∈ [0, 1].
|
| 149 |
+
|
| 150 |
+
Determina o quanto da hipótese refinada deve ser aceita no novo estado.
|
| 151 |
+
"""
|
| 152 |
+
|
| 153 |
+
def __init__(self, embed_dim: int):
|
| 154 |
+
super().__init__()
|
| 155 |
+
self.gate = nn.Linear(embed_dim, 1)
|
| 156 |
+
nn.init.zeros_(self.gate.bias)
|
| 157 |
+
|
| 158 |
+
def forward(self, h: torch.Tensor) -> torch.Tensor:
|
| 159 |
+
# Retorna [B, 1] para broadcast
|
| 160 |
+
return torch.sigmoid(self.gate(h))
|
| 161 |
+
|
| 162 |
+
|
| 163 |
+
# ============================================================================
|
| 164 |
+
# Anti-Hallucination Gate (lógica fuzzy de Łukasiewicz simplificada)
|
| 165 |
+
# ============================================================================
|
| 166 |
+
|
| 167 |
+
class FuzzyLogicGate(nn.Module):
|
| 168 |
+
"""Porta lógica fuzzy baseada em Łukasiewicz (conforme dados.txt linhas 510-540).
|
| 169 |
+
|
| 170 |
+
Computa um valor de verdade suave tau ∈ [0, 1] que mede a coerência
|
| 171 |
+
lógica do estado h_c. Se tau < tau_min, o ciclo é rejeitado (h permanece).
|
| 172 |
+
"""
|
| 173 |
+
|
| 174 |
+
def __init__(self, embed_dim: int, tau_min: float = 0.3):
|
| 175 |
+
super().__init__()
|
| 176 |
+
# Projeção para um escalar (representando o "valor de verdade")
|
| 177 |
+
self.proj = nn.Linear(embed_dim, 1)
|
| 178 |
+
self.tau_min = tau_min
|
| 179 |
+
nn.init.zeros_(self.proj.bias)
|
| 180 |
+
|
| 181 |
+
def forward(self, h: torch.Tensor) -> torch.Tensor:
|
| 182 |
+
"""Retorna tau ∈ [0, 1] — [B, 1]."""
|
| 183 |
+
# sigmoid(w^T h + b) — suaviza para [0, 1]
|
| 184 |
+
return torch.sigmoid(self.proj(h))
|
| 185 |
+
|
| 186 |
+
def accept(self, tau: torch.Tensor) -> torch.Tensor:
|
| 187 |
+
"""Máscara booleana [B, 1]: True se tau >= tau_min."""
|
| 188 |
+
return (tau >= self.tau_min).float()
|
| 189 |
+
|
| 190 |
+
|
| 191 |
+
# ============================================================================
|
| 192 |
+
# Cyclic Reasoning Module
|
| 193 |
+
# ============================================================================
|
| 194 |
+
|
| 195 |
+
class CyclicReasoning(nn.Module):
|
| 196 |
+
"""Módulo de raciocínio cíclico iterativo.
|
| 197 |
+
|
| 198 |
+
Aplica múltiplos ciclos de refinamento sobre a representação inicial,
|
| 199 |
+
parando antecipadamente quando a convergência é atingida.
|
| 200 |
+
|
| 201 |
+
Args:
|
| 202 |
+
config: configuração do módulo
|
| 203 |
+
"""
|
| 204 |
+
|
| 205 |
+
def __init__(self, config: CyclicReasoningConfig):
|
| 206 |
+
super().__init__()
|
| 207 |
+
self.config = config
|
| 208 |
+
self.embed_dim = config.embed_dim
|
| 209 |
+
|
| 210 |
+
# Camadas compartilhadas entre ciclos (weight sharing) — reduz parâmetros
|
| 211 |
+
# Alternativa: um RefinementLayer por ciclo (mais capacidade, mais parâmetros)
|
| 212 |
+
self.refinement = RefinementLayer(
|
| 213 |
+
embed_dim=config.embed_dim,
|
| 214 |
+
hypothesis_dim=config.hypothesis_input_dim,
|
| 215 |
+
hidden_mult=config.refinement_hidden_mult,
|
| 216 |
+
dropout=config.dropout,
|
| 217 |
+
)
|
| 218 |
+
self.confidence_gate = ConfidenceGate(config.embed_dim)
|
| 219 |
+
self.fuzzy_gate = (
|
| 220 |
+
FuzzyLogicGate(config.embed_dim) if config.use_anti_hallucination_gate else None
|
| 221 |
+
)
|
| 222 |
+
self.ln = nn.LayerNorm(config.embed_dim) if config.use_layer_norm else None
|
| 223 |
+
|
| 224 |
+
def forward(
|
| 225 |
+
self,
|
| 226 |
+
h0: torch.Tensor,
|
| 227 |
+
hypothesis_fn: Optional[Callable[[torch.Tensor, int], Optional[torch.Tensor]]] = None,
|
| 228 |
+
penalty_signal: Optional[torch.Tensor] = None,
|
| 229 |
+
return_history: bool = False,
|
| 230 |
+
) -> Dict[str, torch.Tensor]:
|
| 231 |
+
"""Executa raciocínio cíclico.
|
| 232 |
+
|
| 233 |
+
Args:
|
| 234 |
+
h0: [B, D] representação inicial
|
| 235 |
+
hypothesis_fn: função (h_c, c) -> g_c (hipótese opcional para ciclo c).
|
| 236 |
+
Se None, g_c = None (sem hipótese).
|
| 237 |
+
penalty_signal: [B] ou [B, 1] penalidade do verificador (para log).
|
| 238 |
+
return_history: se True, retorna lista de h_c por ciclo.
|
| 239 |
+
|
| 240 |
+
Returns:
|
| 241 |
+
dict com:
|
| 242 |
+
h_final: [B, D] estado final
|
| 243 |
+
n_cycles: int — número de ciclos executados
|
| 244 |
+
converged: bool — se convergiu antes de C_max
|
| 245 |
+
deltas: lista de ||h_c - h_{c-1}|| por ciclo
|
| 246 |
+
history: lista de h_c (se return_history=True)
|
| 247 |
+
alphas: lista de alpha_c
|
| 248 |
+
taus: lista de tau_c (se fuzzy gate ativo)
|
| 249 |
+
"""
|
| 250 |
+
if h0.dim() != 2:
|
| 251 |
+
raise ValueError(f"h0 deve ser [B, D]; recebido {h0.shape}")
|
| 252 |
+
|
| 253 |
+
B, D = h0.shape
|
| 254 |
+
if D != self.embed_dim:
|
| 255 |
+
raise ValueError(f"embed_dim mismatch: h0={D}, config={self.embed_dim}")
|
| 256 |
+
|
| 257 |
+
h = h0
|
| 258 |
+
history: List[torch.Tensor] = [h0]
|
| 259 |
+
deltas: List[float] = []
|
| 260 |
+
alphas: List[float] = []
|
| 261 |
+
taus: List[float] = []
|
| 262 |
+
converged = False
|
| 263 |
+
n_cycles_exec = 0
|
| 264 |
+
|
| 265 |
+
for c in range(self.config.max_cycles):
|
| 266 |
+
n_cycles_exec = c + 1
|
| 267 |
+
|
| 268 |
+
# Truncated BPTT: detach após K ciclos
|
| 269 |
+
if (
|
| 270 |
+
self.config.use_truncated_bptt
|
| 271 |
+
and c >= self.config.bptt_truncate
|
| 272 |
+
):
|
| 273 |
+
h = h.detach()
|
| 274 |
+
|
| 275 |
+
# 1. Gerar hipótese g_c (se função fornecida)
|
| 276 |
+
g_c = None
|
| 277 |
+
if hypothesis_fn is not None:
|
| 278 |
+
try:
|
| 279 |
+
g_c = hypothesis_fn(h, c)
|
| 280 |
+
except Exception as e:
|
| 281 |
+
logger.debug("hypothesis_fn falhou no ciclo %d: %s", c, e)
|
| 282 |
+
g_c = None
|
| 283 |
+
|
| 284 |
+
# 2. Refinamento
|
| 285 |
+
r_c = self.refinement(h, g_c)
|
| 286 |
+
|
| 287 |
+
# 3. Coeficiente de confiança
|
| 288 |
+
alpha_c = self.confidence_gate(h) # [B, 1]
|
| 289 |
+
|
| 290 |
+
# 4. Novo estado
|
| 291 |
+
if self.config.use_residual:
|
| 292 |
+
h_new = h + alpha_c * r_c
|
| 293 |
+
else:
|
| 294 |
+
h_new = alpha_c * r_c
|
| 295 |
+
|
| 296 |
+
# 5. Anti-hallucination gate
|
| 297 |
+
if self.fuzzy_gate is not None:
|
| 298 |
+
tau_c = self.fuzzy_gate(h_new) # [B, 1]
|
| 299 |
+
accept_mask = self.fuzzy_gate.accept(tau_c) # [B, 1] em {0, 1}
|
| 300 |
+
h_new = h * (1 - accept_mask) + h_new * accept_mask
|
| 301 |
+
taus.append(float(tau_c.mean().detach()))
|
| 302 |
+
else:
|
| 303 |
+
tau_c = None
|
| 304 |
+
|
| 305 |
+
# 6. LayerNorm
|
| 306 |
+
if self.ln is not None:
|
| 307 |
+
h_new = self.ln(h_new)
|
| 308 |
+
|
| 309 |
+
# 7. Convergência (delta = ||h_new - h||)
|
| 310 |
+
with torch.no_grad():
|
| 311 |
+
delta = (h_new - h).norm(dim=-1).mean().item()
|
| 312 |
+
deltas.append(delta)
|
| 313 |
+
alphas.append(float(alpha_c.mean().detach()))
|
| 314 |
+
|
| 315 |
+
h = h_new
|
| 316 |
+
history.append(h)
|
| 317 |
+
|
| 318 |
+
# 8. Parada antecipada
|
| 319 |
+
if delta < self.config.convergence_eps:
|
| 320 |
+
converged = True
|
| 321 |
+
logger.debug(
|
| 322 |
+
"CyclicReasoning convergiu no ciclo %d (delta=%.6f < eps=%.6f)",
|
| 323 |
+
c + 1, delta, self.config.convergence_eps,
|
| 324 |
+
)
|
| 325 |
+
break
|
| 326 |
+
|
| 327 |
+
result = {
|
| 328 |
+
"h_final": h,
|
| 329 |
+
"n_cycles": n_cycles_exec,
|
| 330 |
+
"converged": converged,
|
| 331 |
+
"deltas": deltas,
|
| 332 |
+
"alphas": alphas,
|
| 333 |
+
"taus": taus,
|
| 334 |
+
}
|
| 335 |
+
if return_history:
|
| 336 |
+
result["history"] = history
|
| 337 |
+
if penalty_signal is not None:
|
| 338 |
+
result["penalty_signal"] = penalty_signal
|
| 339 |
+
return result
|
| 340 |
+
|
| 341 |
+
|
| 342 |
+
# ============================================================================
|
| 343 |
+
# Helper: hypothesis_fn factory a partir do HypothesisController
|
| 344 |
+
# ============================================================================
|
| 345 |
+
|
| 346 |
+
def make_hypothesis_fn(hypothesis_controller, penalty_threshold: float = 0.5):
|
| 347 |
+
"""Cria uma função hypothesis_fn (h_c, c) -> g_c a partir do HypothesisController.
|
| 348 |
+
|
| 349 |
+
O HypothesisController gera hipóteses quando ativado por penalidades.
|
| 350 |
+
Esta função adapta a interface para uso no CyclicReasoning.
|
| 351 |
+
"""
|
| 352 |
+
def hypothesis_fn(h_c: torch.Tensor, c: int) -> Optional[torch.Tensor]:
|
| 353 |
+
try:
|
| 354 |
+
# Tenta chamar o método generate_hypothesis do controller
|
| 355 |
+
if hasattr(hypothesis_controller, "generate_hypothesis"):
|
| 356 |
+
# Passa h_c como estado e retorna a hipótese (projeta se necessário)
|
| 357 |
+
hyp = hypothesis_controller.generate_hypothesis(h_c)
|
| 358 |
+
if hyp is None:
|
| 359 |
+
return None
|
| 360 |
+
# Se dim diferente, projeta
|
| 361 |
+
if hasattr(hyp, "dim") and hyp.dim() == 2 and hyp.size(-1) != h_c.size(-1):
|
| 362 |
+
if not hasattr(hypothesis_controller, "_proj_to_h"):
|
| 363 |
+
hypothesis_controller._proj_to_h = nn.Linear(
|
| 364 |
+
hyp.size(-1), h_c.size(-1), bias=False
|
| 365 |
+
).to(hyp.device)
|
| 366 |
+
nn.init.orthogonal_(hypothesis_controller._proj_to_h.weight)
|
| 367 |
+
hyp = hypothesis_controller._proj_to_h(hyp)
|
| 368 |
+
return hyp
|
| 369 |
+
return None
|
| 370 |
+
except Exception as e:
|
| 371 |
+
logger.debug("hypothesis_fn erro: %s", e)
|
| 372 |
+
return None
|
| 373 |
+
|
| 374 |
+
return hypothesis_fn
|
| 375 |
+
|
| 376 |
+
|
| 377 |
+
# ============================================================================
|
| 378 |
+
# Demo / Teste rápido
|
| 379 |
+
# ============================================================================
|
| 380 |
+
|
| 381 |
+
def _self_test():
|
| 382 |
+
"""Teste rápido do módulo."""
|
| 383 |
+
torch.manual_seed(42)
|
| 384 |
+
config = CyclicReasoningConfig(
|
| 385 |
+
embed_dim=64,
|
| 386 |
+
max_cycles=5,
|
| 387 |
+
convergence_eps=1e-3,
|
| 388 |
+
use_anti_hallucination_gate=True,
|
| 389 |
+
)
|
| 390 |
+
cr = CyclicReasoning(config)
|
| 391 |
+
h0 = torch.randn(4, 64)
|
| 392 |
+
result = cr(h0, return_history=True)
|
| 393 |
+
print(f"n_cycles={result['n_cycles']}, converged={result['converged']}")
|
| 394 |
+
print(f"deltas={result['deltas']}")
|
| 395 |
+
print(f"alphas={result['alphas']}")
|
| 396 |
+
print(f"taus={result['taus']}")
|
| 397 |
+
print(f"h_final shape={result['h_final'].shape}")
|
| 398 |
+
|
| 399 |
+
|
| 400 |
+
if __name__ == "__main__":
|
| 401 |
+
_self_test()
|
| 402 |
+
|
| 403 |
+
|
| 404 |
+
__all__ = [
|
| 405 |
+
"CyclicReasoningConfig",
|
| 406 |
+
"RefinementLayer",
|
| 407 |
+
"ConfidenceGate",
|
| 408 |
+
"FuzzyLogicGate",
|
| 409 |
+
"CyclicReasoning",
|
| 410 |
+
"make_hypothesis_fn",
|
| 411 |
+
]
|
models/generator_verifier.py
CHANGED
|
@@ -317,8 +317,10 @@ class VerifierCNNBiGRU(nn.Module):
|
|
| 317 |
)
|
| 318 |
|
| 319 |
# Classificador sigmoid
|
|
|
|
|
|
|
| 320 |
self.classifier = nn.Sequential(
|
| 321 |
-
nn.Linear(gru_hidden *
|
| 322 |
nn.ReLU(),
|
| 323 |
nn.Dropout(dropout),
|
| 324 |
nn.Linear(gru_hidden, 1),
|
|
|
|
| 317 |
)
|
| 318 |
|
| 319 |
# Classificador sigmoid
|
| 320 |
+
# fused = out_A || out_B = (2*gru_hidden) || (2*gru_hidden) = 4*gru_hidden
|
| 321 |
+
# combined = feat_pre || feat_passo = 8*gru_hidden
|
| 322 |
self.classifier = nn.Sequential(
|
| 323 |
+
nn.Linear(gru_hidden * 4 * 2, gru_hidden), # 8*gru_hidden -> gru_hidden
|
| 324 |
nn.ReLU(),
|
| 325 |
nn.Dropout(dropout),
|
| 326 |
nn.Linear(gru_hidden, 1),
|
models/medusa_heads.py
ADDED
|
@@ -0,0 +1,409 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""medusa_heads.py — Multi-Token Prediction (MTP) com Medusa Heads.
|
| 2 |
+
|
| 3 |
+
Implementa Medusa (Cai et al., 2024): múltiplas cabeças de previsão que
|
| 4 |
+
predizem tokens em posições futuras (t+1, t+2, ..., t+K), permitindo
|
| 5 |
+
geração paralela de K tokens por forward pass.
|
| 6 |
+
|
| 7 |
+
==============================================================================
|
| 8 |
+
ANÁLISE MATEMÁTICA E LÓGICA
|
| 9 |
+
==============================================================================
|
| 10 |
+
|
| 11 |
+
Seja h_t ∈ R^d o hidden state produzido pelo modelo (por exemplo, saída do
|
| 12 |
+
TransformerDecoderStack no tempo t). Em vez de prever apenas y_{t+1},
|
| 13 |
+
Medusa adiciona K cabeças extras:
|
| 14 |
+
|
| 15 |
+
y_{t+k} = softmax(W_k h_t + b_k) para k = 1, ..., K
|
| 16 |
+
|
| 17 |
+
A cabeça principal (k=1) é o LM head original; as cabeças k=2..K são
|
| 18 |
+
"cabeças Medusa" adicionais treinadas com a loss:
|
| 19 |
+
|
| 20 |
+
L_Medusa = sum_{k=1}^{K} lambda_k * CE(head_k(h_t), y_{t+k})
|
| 21 |
+
|
| 22 |
+
onde lambda_k é um peso decrescente (tipicamente lambda_k = 1/k ou
|
| 23 |
+
lambda_k = decay^k), pois previsões mais distantes são mais difíceis.
|
| 24 |
+
|
| 25 |
+
ÁRVORE DE DECODIFICAÇÃO:
|
| 26 |
+
Em inferência, cada cabeça Medusa gera top-B candidatos, formando uma
|
| 27 |
+
árvore de filhos. A árvore é avaliada em paralelo (single forward) e
|
| 28 |
+
os candidatos são aceitos se coincidirem com a previsão do modelo
|
| 29 |
+
na posição correspondente (verificação via attention mask da árvore).
|
| 30 |
+
|
| 31 |
+
Speedup teórico: até K× (em prática 2-3× devido à taxa de aceitação <1).
|
| 32 |
+
|
| 33 |
+
INTEGRAÇÃO COM CNN-BiGRU:
|
| 34 |
+
- As cabeças Medusa são aplicadas sobre o hidden state final do
|
| 35 |
+
TransformerDecoderStack (após self-attention + FFN).
|
| 36 |
+
- O hidden state h_t já contém contexto cooperativo (streams A+B fundidos).
|
| 37 |
+
- Em treinamento, adicionamos L_Medusa à loss total (com peso mu_medusa).
|
| 38 |
+
- Em inferência, ativamos tree-decoding quando use_medusa=True.
|
| 39 |
+
|
| 40 |
+
VANTAGENS:
|
| 41 |
+
1. Acelera geração em ~2-3× sem perda de qualidade.
|
| 42 |
+
2. Não requer mudanças no backbone — apenas cabeças extras.
|
| 43 |
+
3. Treinável com poucas amostras (as cabeças convergem rápido).
|
| 44 |
+
|
| 45 |
+
LIMITAÇÕES:
|
| 46 |
+
1. Aumenta o uso de memória (K * vocab_size params extras).
|
| 47 |
+
2. Taxa de aceitação depende da tarefa — previsões muito long-range
|
| 48 |
+
são menos precisas.
|
| 49 |
+
|
| 50 |
+
Autor: CNN-BiGRU Project
|
| 51 |
+
"""
|
| 52 |
+
from __future__ import annotations
|
| 53 |
+
|
| 54 |
+
import logging
|
| 55 |
+
import math
|
| 56 |
+
from dataclasses import dataclass, field
|
| 57 |
+
from typing import Dict, List, Optional, Tuple
|
| 58 |
+
|
| 59 |
+
import torch
|
| 60 |
+
import torch.nn as nn
|
| 61 |
+
import torch.nn.functional as F
|
| 62 |
+
|
| 63 |
+
logger = logging.getLogger(__name__)
|
| 64 |
+
|
| 65 |
+
|
| 66 |
+
# ============================================================================
|
| 67 |
+
# Configuração
|
| 68 |
+
# ============================================================================
|
| 69 |
+
|
| 70 |
+
@dataclass
|
| 71 |
+
class MedusaConfig:
|
| 72 |
+
"""Configuração das cabeças Medusa."""
|
| 73 |
+
vocab_size: int = 32000
|
| 74 |
+
embed_dim: int = 256
|
| 75 |
+
n_heads: int = 4 # K cabeças Medusa (além da principal)
|
| 76 |
+
head_hidden_mult: int = 2 # multiplicador da camada oculta de cada cabeça
|
| 77 |
+
dropout: float = 0.1
|
| 78 |
+
# Peso de cada cabeça na loss (lambda_k). Se None, usa 1/(k+1).
|
| 79 |
+
head_weights: Optional[List[float]] = None
|
| 80 |
+
# Top-B candidatos por cabeça para tree decoding
|
| 81 |
+
top_b_candidates: int = 4
|
| 82 |
+
# Se True, compartilha a primeira camada das cabeças (reduz params)
|
| 83 |
+
share_first_layer: bool = False
|
| 84 |
+
# Initialize biases to favor first token (BOS) — estabilidade
|
| 85 |
+
init_bias_to_bos: bool = False
|
| 86 |
+
bos_id: int = 0
|
| 87 |
+
|
| 88 |
+
|
| 89 |
+
# ============================================================================
|
| 90 |
+
# Cabeça Medusa individual
|
| 91 |
+
# ============================================================================
|
| 92 |
+
|
| 93 |
+
class MedusaHead(nn.Module):
|
| 94 |
+
"""Uma única cabeça Medusa: prediz y_{t+k} dado h_t.
|
| 95 |
+
|
| 96 |
+
Estrutura: Linear(d, d*mult) -> SiLU -> Dropout -> Linear(d*mult, V)
|
| 97 |
+
"""
|
| 98 |
+
|
| 99 |
+
def __init__(
|
| 100 |
+
self,
|
| 101 |
+
embed_dim: int,
|
| 102 |
+
vocab_size: int,
|
| 103 |
+
hidden_mult: int = 2,
|
| 104 |
+
dropout: float = 0.1,
|
| 105 |
+
):
|
| 106 |
+
super().__init__()
|
| 107 |
+
hidden = embed_dim * hidden_mult
|
| 108 |
+
self.fc1 = nn.Linear(embed_dim, hidden)
|
| 109 |
+
self.fc2 = nn.Linear(hidden, vocab_size)
|
| 110 |
+
self.dropout = nn.Dropout(dropout)
|
| 111 |
+
# Inicialização
|
| 112 |
+
for layer in (self.fc1, self.fc2):
|
| 113 |
+
nn.init.xavier_uniform_(layer.weight)
|
| 114 |
+
nn.init.zeros_(layer.bias)
|
| 115 |
+
|
| 116 |
+
def forward(self, h: torch.Tensor) -> torch.Tensor:
|
| 117 |
+
"""
|
| 118 |
+
Args:
|
| 119 |
+
h: [B, T, D] ou [B, D] hidden states
|
| 120 |
+
Returns:
|
| 121 |
+
logits: [B, T, V] ou [B, V]
|
| 122 |
+
"""
|
| 123 |
+
x = self.dropout(F.silu(self.fc1(h)))
|
| 124 |
+
return self.fc2(x)
|
| 125 |
+
|
| 126 |
+
|
| 127 |
+
# ============================================================================
|
| 128 |
+
# MTP com Medusa Heads
|
| 129 |
+
# ============================================================================
|
| 130 |
+
|
| 131 |
+
class MedusaMTP(nn.Module):
|
| 132 |
+
"""Multi-Token Prediction com Medusa Heads.
|
| 133 |
+
|
| 134 |
+
Args:
|
| 135 |
+
config: configuração das cabeças
|
| 136 |
+
"""
|
| 137 |
+
|
| 138 |
+
def __init__(self, config: MedusaConfig):
|
| 139 |
+
super().__init__()
|
| 140 |
+
self.config = config
|
| 141 |
+
self.K = config.n_heads
|
| 142 |
+
self.vocab_size = config.vocab_size
|
| 143 |
+
self.embed_dim = config.embed_dim
|
| 144 |
+
|
| 145 |
+
# Cabeças Medusa (k=1..K)
|
| 146 |
+
# k=1 é a cabeça "principal extra" (além do lm_head do modelo)
|
| 147 |
+
# k=2..K são cabeças adicionais
|
| 148 |
+
if config.share_first_layer:
|
| 149 |
+
# Compartilha fc1 entre todas as cabeças
|
| 150 |
+
shared_hidden = config.embed_dim * config.head_hidden_mult
|
| 151 |
+
self.shared_fc1 = nn.Linear(config.embed_dim, shared_hidden)
|
| 152 |
+
nn.init.xavier_uniform_(self.shared_fc1.weight)
|
| 153 |
+
nn.init.zeros_(self.shared_fc1.bias)
|
| 154 |
+
self.heads = nn.ModuleList([
|
| 155 |
+
nn.Linear(shared_hidden, config.vocab_size)
|
| 156 |
+
for _ in range(config.n_heads)
|
| 157 |
+
])
|
| 158 |
+
for h in self.heads:
|
| 159 |
+
nn.init.xavier_uniform_(h.weight)
|
| 160 |
+
nn.init.zeros_(h.bias)
|
| 161 |
+
self._shared = True
|
| 162 |
+
else:
|
| 163 |
+
self.heads = nn.ModuleList([
|
| 164 |
+
MedusaHead(
|
| 165 |
+
embed_dim=config.embed_dim,
|
| 166 |
+
vocab_size=config.vocab_size,
|
| 167 |
+
hidden_mult=config.head_hidden_mult,
|
| 168 |
+
dropout=config.dropout,
|
| 169 |
+
)
|
| 170 |
+
for _ in range(config.n_heads)
|
| 171 |
+
])
|
| 172 |
+
self._shared = False
|
| 173 |
+
|
| 174 |
+
# Pesos lambda_k (decrescentes)
|
| 175 |
+
if config.head_weights is not None:
|
| 176 |
+
assert len(config.head_weights) == config.n_heads, \
|
| 177 |
+
"head_weights deve ter n_heads elementos"
|
| 178 |
+
self.register_buffer(
|
| 179 |
+
"head_weights",
|
| 180 |
+
torch.tensor(config.head_weights, dtype=torch.float32),
|
| 181 |
+
)
|
| 182 |
+
else:
|
| 183 |
+
# default: 1/(k+1) para k=1..K
|
| 184 |
+
weights = [1.0 / (k + 1) for k in range(1, config.n_heads + 1)]
|
| 185 |
+
self.register_buffer(
|
| 186 |
+
"head_weights",
|
| 187 |
+
torch.tensor(weights, dtype=torch.float32),
|
| 188 |
+
)
|
| 189 |
+
|
| 190 |
+
# Inicialização bias para BOS (opcional)
|
| 191 |
+
if config.init_bias_to_bos:
|
| 192 |
+
self._init_biases_to_bos()
|
| 193 |
+
|
| 194 |
+
def _init_biases_to_bos(self):
|
| 195 |
+
"""Inicializa biases das cabeças para favorecer BOS no início."""
|
| 196 |
+
bos_id = self.config.bos_id
|
| 197 |
+
with torch.no_grad():
|
| 198 |
+
if self._shared:
|
| 199 |
+
for h in self.heads:
|
| 200 |
+
h.bias.zero_()
|
| 201 |
+
h.bias[bos_id] = 1.0
|
| 202 |
+
else:
|
| 203 |
+
for h in self.heads:
|
| 204 |
+
h.fc2.bias.zero_()
|
| 205 |
+
h.fc2.bias[bos_id] = 1.0
|
| 206 |
+
|
| 207 |
+
def forward(self, h: torch.Tensor) -> List[torch.Tensor]:
|
| 208 |
+
"""Computa logits de todas as cabeças.
|
| 209 |
+
|
| 210 |
+
Args:
|
| 211 |
+
h: [B, T, D] hidden states do backbone (saída do TransformerDecoder)
|
| 212 |
+
|
| 213 |
+
Returns:
|
| 214 |
+
list de K tensores, cada um [B, T, V] (logits da cabeça k)
|
| 215 |
+
heads[k-1] prediz y_{t+k} dado h_t
|
| 216 |
+
"""
|
| 217 |
+
if h.dim() == 2:
|
| 218 |
+
h = h.unsqueeze(1) # [B, D] -> [B, 1, D]
|
| 219 |
+
if h.size(-1) != self.embed_dim:
|
| 220 |
+
raise ValueError(
|
| 221 |
+
f"hidden dim mismatch: h={h.size(-1)}, config={self.embed_dim}"
|
| 222 |
+
)
|
| 223 |
+
|
| 224 |
+
all_logits = []
|
| 225 |
+
if self._shared:
|
| 226 |
+
shared = F.silu(self.shared_fc1(h))
|
| 227 |
+
for head in self.heads:
|
| 228 |
+
all_logits.append(head(shared))
|
| 229 |
+
else:
|
| 230 |
+
for head in self.heads:
|
| 231 |
+
all_logits.append(head(h))
|
| 232 |
+
return all_logits
|
| 233 |
+
|
| 234 |
+
def compute_loss(
|
| 235 |
+
self,
|
| 236 |
+
h: torch.Tensor,
|
| 237 |
+
target_ids: torch.Tensor,
|
| 238 |
+
ignore_index: int = -100,
|
| 239 |
+
) -> Tuple[torch.Tensor, Dict[str, float]]:
|
| 240 |
+
"""Computa a loss Medusa total.
|
| 241 |
+
|
| 242 |
+
Args:
|
| 243 |
+
h: [B, T, D] hidden states do backbone
|
| 244 |
+
target_ids: [B, T] tokens alvo (y_{t+1}, y_{t+2}, ...)
|
| 245 |
+
A cabeça k deve prever target_ids[:, t+k-1] dado h_t.
|
| 246 |
+
ignore_index: índice a ignorar na CE (default: -100, ou pad_id)
|
| 247 |
+
|
| 248 |
+
Returns:
|
| 249 |
+
loss: escalar (weighted sum of head losses)
|
| 250 |
+
stats: dict com loss por cabeça
|
| 251 |
+
"""
|
| 252 |
+
if h.dim() == 2:
|
| 253 |
+
h = h.unsqueeze(1)
|
| 254 |
+
B, T, D = h.shape
|
| 255 |
+
if target_ids.dim() != 2:
|
| 256 |
+
raise ValueError(f"target_ids deve ser [B, T]; recebido {target_ids.shape}")
|
| 257 |
+
|
| 258 |
+
all_logits = self.forward(h) # list de K tensores [B, T, V]
|
| 259 |
+
|
| 260 |
+
total_loss = torch.zeros(1, device=h.device, dtype=h.dtype)
|
| 261 |
+
stats = {}
|
| 262 |
+
|
| 263 |
+
for k in range(self.K):
|
| 264 |
+
# Cabeça k prediz y_{t+k+1} dado h_t (k=0 -> y_{t+1}, k=1 -> y_{t+2}, ...)
|
| 265 |
+
shift = k + 1
|
| 266 |
+
if shift >= T:
|
| 267 |
+
# Não há targets suficientes para esta cabeça
|
| 268 |
+
stats[f"head_{k+1}_loss"] = float("nan")
|
| 269 |
+
continue
|
| 270 |
+
# logits_k: [B, T, V] -> pegamos[:B, :T-shift, :] para alinhar com target[:, shift:]
|
| 271 |
+
logits_k = all_logits[k][:, :T - shift, :].contiguous() # [B, T-shift, V]
|
| 272 |
+
target_k = target_ids[:, shift:T].contiguous() # [B, T-shift]
|
| 273 |
+
|
| 274 |
+
loss_k = F.cross_entropy(
|
| 275 |
+
logits_k.view(-1, self.vocab_size),
|
| 276 |
+
target_k.view(-1),
|
| 277 |
+
ignore_index=ignore_index,
|
| 278 |
+
reduction="mean",
|
| 279 |
+
)
|
| 280 |
+
weight_k = self.head_weights[k].to(h.device)
|
| 281 |
+
total_loss = total_loss + weight_k * loss_k
|
| 282 |
+
stats[f"head_{k+1}_loss"] = float(loss_k.detach())
|
| 283 |
+
|
| 284 |
+
return total_loss.squeeze(), stats
|
| 285 |
+
|
| 286 |
+
def predict_top_b(
|
| 287 |
+
self,
|
| 288 |
+
h: torch.Tensor,
|
| 289 |
+
top_b: Optional[int] = None,
|
| 290 |
+
) -> List[Tuple[torch.Tensor, torch.Tensor]]:
|
| 291 |
+
"""Para cada cabeça, retorna top-B candidatos e suas probabilidades.
|
| 292 |
+
|
| 293 |
+
Usado em tree decoding (inferência).
|
| 294 |
+
|
| 295 |
+
Args:
|
| 296 |
+
h: [B, T, D] hidden states (ou [B, D] para último token)
|
| 297 |
+
top_b: número de candidatos (default: config.top_b_candidates)
|
| 298 |
+
|
| 299 |
+
Returns:
|
| 300 |
+
list de K tuplas (tokens, probs) onde:
|
| 301 |
+
tokens: [B, top_b] — IDs dos top-B tokens
|
| 302 |
+
probs: [B, top_b] — probabilidades (após softmax)
|
| 303 |
+
"""
|
| 304 |
+
if h.dim() == 2:
|
| 305 |
+
h = h.unsqueeze(1)
|
| 306 |
+
# Usar apenas o último token para previsão
|
| 307 |
+
h_last = h[:, -1:, :] # [B, 1, D]
|
| 308 |
+
all_logits = self.forward(h_last) # list de K [B, 1, V]
|
| 309 |
+
b = top_b or self.config.top_b_candidates
|
| 310 |
+
|
| 311 |
+
results = []
|
| 312 |
+
for k, logits_k in enumerate(all_logits):
|
| 313 |
+
logits_k = logits_k[:, -1, :] # [B, V]
|
| 314 |
+
probs = F.softmax(logits_k, dim=-1)
|
| 315 |
+
top_probs, top_ids = torch.topk(probs, k=b, dim=-1)
|
| 316 |
+
results.append((top_ids, top_probs))
|
| 317 |
+
return results
|
| 318 |
+
|
| 319 |
+
|
| 320 |
+
# ============================================================================
|
| 321 |
+
# Tree Decoding (esqueleto — implementação simplificada)
|
| 322 |
+
# ============================================================================
|
| 323 |
+
|
| 324 |
+
def medusa_tree_decode(
|
| 325 |
+
medusa: MedusaMTP,
|
| 326 |
+
base_logits: torch.Tensor,
|
| 327 |
+
h_last: torch.Tensor,
|
| 328 |
+
accept_threshold: float = 0.0,
|
| 329 |
+
) -> Dict[str, torch.Tensor]:
|
| 330 |
+
"""Decodificação por árvore Medusa (simplificada).
|
| 331 |
+
|
| 332 |
+
Em uma implementação completa, isto construiria uma árvore de candidatos
|
| 333 |
+
e faria um forward pass único para verificar aceitação. Aqui retornamos
|
| 334 |
+
apenas os candidatos de cada cabeça para uso posterior.
|
| 335 |
+
|
| 336 |
+
Args:
|
| 337 |
+
medusa: módulo MedusaMTP
|
| 338 |
+
base_logits: [B, V] logits da cabeça principal (já computados)
|
| 339 |
+
h_last: [B, D] hidden state do último token
|
| 340 |
+
accept_threshold: threshold de probabilidade para aceitar candidato
|
| 341 |
+
|
| 342 |
+
Returns:
|
| 343 |
+
dict com:
|
| 344 |
+
base_token: [B] argmax da cabeça principal
|
| 345 |
+
base_prob: [B] prob do argmax
|
| 346 |
+
medusa_candidates: list de K tuplas (tokens, probs)
|
| 347 |
+
accepted: list de K tensores booleanos [B] indicando aceitação
|
| 348 |
+
"""
|
| 349 |
+
base_probs = F.softmax(base_logits, dim=-1)
|
| 350 |
+
base_token = base_probs.argmax(dim=-1)
|
| 351 |
+
base_prob = base_probs.max(dim=-1).values
|
| 352 |
+
|
| 353 |
+
candidates = medusa.predict_top_b(h_last)
|
| 354 |
+
|
| 355 |
+
accepted = []
|
| 356 |
+
for k, (tokens, probs) in enumerate(candidates):
|
| 357 |
+
# Aceita o primeiro candidato se prob >= threshold
|
| 358 |
+
first_prob = probs[:, 0]
|
| 359 |
+
accept = first_prob >= accept_threshold
|
| 360 |
+
accepted.append(accept)
|
| 361 |
+
|
| 362 |
+
return {
|
| 363 |
+
"base_token": base_token,
|
| 364 |
+
"base_prob": base_prob,
|
| 365 |
+
"medusa_candidates": candidates,
|
| 366 |
+
"accepted": accepted,
|
| 367 |
+
}
|
| 368 |
+
|
| 369 |
+
|
| 370 |
+
# ============================================================================
|
| 371 |
+
# Self-test
|
| 372 |
+
# ============================================================================
|
| 373 |
+
|
| 374 |
+
def _self_test():
|
| 375 |
+
torch.manual_seed(42)
|
| 376 |
+
config = MedusaConfig(
|
| 377 |
+
vocab_size=100,
|
| 378 |
+
embed_dim=32,
|
| 379 |
+
n_heads=4,
|
| 380 |
+
head_hidden_mult=2,
|
| 381 |
+
dropout=0.1,
|
| 382 |
+
)
|
| 383 |
+
medusa = MedusaMTP(config)
|
| 384 |
+
h = torch.randn(2, 8, 32)
|
| 385 |
+
target = torch.randint(0, 100, (2, 8))
|
| 386 |
+
loss, stats = medusa.compute_loss(h, target)
|
| 387 |
+
print(f"Loss: {loss.item():.4f}")
|
| 388 |
+
print(f"Stats: {stats}")
|
| 389 |
+
|
| 390 |
+
# Test top-B
|
| 391 |
+
cands = medusa.predict_top_b(h)
|
| 392 |
+
print(f"K={len(cands)}, top_b shape: {cands[0][0].shape}")
|
| 393 |
+
|
| 394 |
+
# Test tree decode
|
| 395 |
+
base_logits = torch.randn(2, 100)
|
| 396 |
+
result = medusa_tree_decode(medusa, base_logits, h[:, -1, :])
|
| 397 |
+
print(f"Base token: {result['base_token']}")
|
| 398 |
+
|
| 399 |
+
|
| 400 |
+
if __name__ == "__main__":
|
| 401 |
+
_self_test()
|
| 402 |
+
|
| 403 |
+
|
| 404 |
+
__all__ = [
|
| 405 |
+
"MedusaConfig",
|
| 406 |
+
"MedusaHead",
|
| 407 |
+
"MedusaMTP",
|
| 408 |
+
"medusa_tree_decode",
|
| 409 |
+
]
|
models/multimodal_attention.py
ADDED
|
@@ -0,0 +1,470 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""multimodal_attention.py — Multi-head Attention Multimodal para CNN-BiGRU.
|
| 2 |
+
|
| 3 |
+
Implementa atenção multi-cabeça cross-modal que combina informações de
|
| 4 |
+
múltiplas modalidades (texto A, texto B, imagem, áudio) via mecanismo
|
| 5 |
+
de atenção cruzada hierárquica.
|
| 6 |
+
|
| 7 |
+
==============================================================================
|
| 8 |
+
ANÁLISE MATEMÁTICA E LÓGICA
|
| 9 |
+
==============================================================================
|
| 10 |
+
|
| 11 |
+
Dadas as representações de N modalidades:
|
| 12 |
+
H_1, H_2, ..., H_N onde H_i ∈ R^{B×T_i×d_i}
|
| 13 |
+
|
| 14 |
+
A multimodal multi-head attention computa:
|
| 15 |
+
|
| 16 |
+
Para cada par (i, j) com i != j:
|
| 17 |
+
Q_i = H_i W_Q^i [B, T_i, d]
|
| 18 |
+
K_j = H_j W_K^j [B, T_j, d]
|
| 19 |
+
V_j = H_j W_V^j [B, T_j, d]
|
| 20 |
+
|
| 21 |
+
A_{i->j} = softmax(Q_i K_j^T / sqrt(d_h)) V_j [B, T_i, d]
|
| 22 |
+
|
| 23 |
+
Output para modalidade i:
|
| 24 |
+
H_i' = H_i + sum_{j != i} A_{i->j}
|
| 25 |
+
|
| 26 |
+
Em seguida, uma camada de fusão final agrega todas as H_i':
|
| 27 |
+
H_fused = sum_i alpha_i * Pool(H_i') onde alpha_i é um gate de modalidade
|
| 28 |
+
|
| 29 |
+
VANTAGENS:
|
| 30 |
+
1. Cada modalidade "olha" para todas as outras — descoberta de correlações
|
| 31 |
+
2. Atenção multi-cabeça permite capturar diferentes tipos de relações
|
| 32 |
+
3. Gate de modalidade permite PESAR dinamicamente cada modalidade
|
| 33 |
+
(e.g., imagem é mais importante para pergunta visual)
|
| 34 |
+
|
| 35 |
+
==============================================================================
|
| 36 |
+
INTEGRAÇÃO
|
| 37 |
+
==============================================================================
|
| 38 |
+
|
| 39 |
+
- Recebe as saídas do CooperativeCNNBiGRU (seq_A, seq_B) e dos encoders
|
| 40 |
+
multimodais (img_seq, audio_seq)
|
| 41 |
+
- Substitui a MultimodalFusion simples (gated) por uma versão com atenção
|
| 42 |
+
- Mantém compatibilidade com a interface existente
|
| 43 |
+
|
| 44 |
+
Autor: CNN-BiGRU Project
|
| 45 |
+
"""
|
| 46 |
+
from __future__ import annotations
|
| 47 |
+
|
| 48 |
+
import logging
|
| 49 |
+
import math
|
| 50 |
+
from dataclasses import dataclass, field
|
| 51 |
+
from typing import Dict, List, Optional, Tuple
|
| 52 |
+
|
| 53 |
+
import torch
|
| 54 |
+
import torch.nn as nn
|
| 55 |
+
import torch.nn.functional as F
|
| 56 |
+
|
| 57 |
+
logger = logging.getLogger(__name__)
|
| 58 |
+
|
| 59 |
+
|
| 60 |
+
# ============================================================================
|
| 61 |
+
# Configuração
|
| 62 |
+
# ============================================================================
|
| 63 |
+
|
| 64 |
+
@dataclass
|
| 65 |
+
class MultimodalAttentionConfig:
|
| 66 |
+
"""Configuração da atenção multimodal multi-cabeça."""
|
| 67 |
+
# Dimensões por modalidade (input)
|
| 68 |
+
d_text_a: int = 128
|
| 69 |
+
d_text_b: int = 128
|
| 70 |
+
d_image: int = 32
|
| 71 |
+
d_audio: int = 32
|
| 72 |
+
# Dimensão interna (todas modalidades são projetadas para d_model)
|
| 73 |
+
d_model: int = 128
|
| 74 |
+
# Número de cabeças
|
| 75 |
+
n_heads: int = 4
|
| 76 |
+
# Dropout
|
| 77 |
+
dropout: float = 0.1
|
| 78 |
+
# Modos de fusão final
|
| 79 |
+
fusion_mode: str = "attention" # "attention" | "gated" | "concat"
|
| 80 |
+
# Usar gate de modalidade (alpha_i)
|
| 81 |
+
use_modality_gate: bool = True
|
| 82 |
+
# Device
|
| 83 |
+
device: str = "cpu"
|
| 84 |
+
|
| 85 |
+
|
| 86 |
+
# ============================================================================
|
| 87 |
+
# Cross-Modal Attention Block
|
| 88 |
+
# ============================================================================
|
| 89 |
+
|
| 90 |
+
class CrossModalAttention(nn.Module):
|
| 91 |
+
"""Atenção cruzada entre duas modalidades.
|
| 92 |
+
|
| 93 |
+
A modalidade "query" atende sobre a modalidade "key_value".
|
| 94 |
+
"""
|
| 95 |
+
|
| 96 |
+
def __init__(
|
| 97 |
+
self,
|
| 98 |
+
d_query: int,
|
| 99 |
+
d_kv: int,
|
| 100 |
+
d_model: int,
|
| 101 |
+
n_heads: int = 4,
|
| 102 |
+
dropout: float = 0.1,
|
| 103 |
+
):
|
| 104 |
+
super().__init__()
|
| 105 |
+
assert d_model % n_heads == 0, \
|
| 106 |
+
f"d_model ({d_model}) deve ser divisível por n_heads ({n_heads})"
|
| 107 |
+
self.d_model = d_model
|
| 108 |
+
self.n_heads = n_heads
|
| 109 |
+
self.head_dim = d_model // n_heads
|
| 110 |
+
|
| 111 |
+
# Projeções para Query (da modalidade query)
|
| 112 |
+
self.q_proj = nn.Linear(d_query, d_model)
|
| 113 |
+
# Projeções para Key, Value (da modalidade kv)
|
| 114 |
+
self.k_proj = nn.Linear(d_kv, d_model)
|
| 115 |
+
self.v_proj = nn.Linear(d_kv, d_model)
|
| 116 |
+
# Output projection
|
| 117 |
+
self.out_proj = nn.Linear(d_model, d_query)
|
| 118 |
+
|
| 119 |
+
self.dropout = nn.Dropout(dropout)
|
| 120 |
+
self.scale = 1.0 / math.sqrt(self.head_dim)
|
| 121 |
+
|
| 122 |
+
# Init
|
| 123 |
+
for layer in (self.q_proj, self.k_proj, self.v_proj, self.out_proj):
|
| 124 |
+
nn.init.xavier_uniform_(layer.weight)
|
| 125 |
+
nn.init.zeros_(layer.bias)
|
| 126 |
+
|
| 127 |
+
def forward(
|
| 128 |
+
self,
|
| 129 |
+
query: torch.Tensor,
|
| 130 |
+
kv: torch.Tensor,
|
| 131 |
+
kv_mask: Optional[torch.Tensor] = None,
|
| 132 |
+
) -> torch.Tensor:
|
| 133 |
+
"""
|
| 134 |
+
Args:
|
| 135 |
+
query: [B, T_q, d_query]
|
| 136 |
+
kv: [B, T_kv, d_kv]
|
| 137 |
+
kv_mask: [B, T_kv] (1=valid, 0=pad)
|
| 138 |
+
|
| 139 |
+
Returns:
|
| 140 |
+
output: [B, T_q, d_query] — query enriquecida com info de kv
|
| 141 |
+
"""
|
| 142 |
+
B, Tq, _ = query.shape
|
| 143 |
+
Tkv = kv.size(1)
|
| 144 |
+
|
| 145 |
+
# Projeções
|
| 146 |
+
q = self.q_proj(query) # [B, T_q, d_model]
|
| 147 |
+
k = self.k_proj(kv) # [B, T_kv, d_model]
|
| 148 |
+
v = self.v_proj(kv) # [B, T_kv, d_model]
|
| 149 |
+
|
| 150 |
+
# Reshape para multi-head: [B, n_heads, T, head_dim]
|
| 151 |
+
q = q.view(B, Tq, self.n_heads, self.head_dim).transpose(1, 2)
|
| 152 |
+
k = k.view(B, Tkv, self.n_heads, self.head_dim).transpose(1, 2)
|
| 153 |
+
v = v.view(B, Tkv, self.n_heads, self.head_dim).transpose(1, 2)
|
| 154 |
+
|
| 155 |
+
# Scores: [B, n_heads, T_q, T_kv]
|
| 156 |
+
scores = torch.matmul(q, k.transpose(-2, -1)) * self.scale
|
| 157 |
+
|
| 158 |
+
# Máscara (padding em kv)
|
| 159 |
+
if kv_mask is not None:
|
| 160 |
+
scores = scores.masked_fill(
|
| 161 |
+
kv_mask.unsqueeze(1).unsqueeze(1) == 0, -1e9
|
| 162 |
+
)
|
| 163 |
+
|
| 164 |
+
# Softmax
|
| 165 |
+
attn = F.softmax(scores, dim=-1)
|
| 166 |
+
attn = self.dropout(attn)
|
| 167 |
+
|
| 168 |
+
# Contexto: [B, n_heads, T_q, head_dim]
|
| 169 |
+
ctx = torch.matmul(attn, v)
|
| 170 |
+
# Reshape: [B, T_q, d_model]
|
| 171 |
+
ctx = ctx.transpose(1, 2).contiguous().view(B, Tq, self.d_model)
|
| 172 |
+
|
| 173 |
+
# Output projection de volta para d_query
|
| 174 |
+
return self.out_proj(ctx)
|
| 175 |
+
|
| 176 |
+
|
| 177 |
+
# ============================================================================
|
| 178 |
+
# Modality Gate (peso dinâmico por modalidade)
|
| 179 |
+
# ============================================================================
|
| 180 |
+
|
| 181 |
+
class ModalityGate(nn.Module):
|
| 182 |
+
"""Computa pesos alpha_i ∈ [0, 1] para cada modalidade.
|
| 183 |
+
|
| 184 |
+
Usa a representação pooled de cada modalidade para decidir o peso.
|
| 185 |
+
"""
|
| 186 |
+
|
| 187 |
+
def __init__(
|
| 188 |
+
self,
|
| 189 |
+
d_text_a: int,
|
| 190 |
+
d_text_b: int,
|
| 191 |
+
d_image: int,
|
| 192 |
+
d_audio: int,
|
| 193 |
+
d_hidden: int = 64,
|
| 194 |
+
):
|
| 195 |
+
super().__init__()
|
| 196 |
+
# Projeta cada modalidade para d_hidden
|
| 197 |
+
self.proj_a = nn.Linear(d_text_a, d_hidden)
|
| 198 |
+
self.proj_b = nn.Linear(d_text_b, d_hidden)
|
| 199 |
+
self.proj_img = nn.Linear(d_image, d_hidden)
|
| 200 |
+
self.proj_aud = nn.Linear(d_audio, d_hidden)
|
| 201 |
+
|
| 202 |
+
# Score: d_hidden -> 1
|
| 203 |
+
self.scorer = nn.Linear(d_hidden, 1)
|
| 204 |
+
nn.init.zeros_(self.scorer.bias)
|
| 205 |
+
|
| 206 |
+
def forward(
|
| 207 |
+
self,
|
| 208 |
+
h_a: torch.Tensor,
|
| 209 |
+
h_b: torch.Tensor,
|
| 210 |
+
h_img: torch.Tensor,
|
| 211 |
+
h_aud: torch.Tensor,
|
| 212 |
+
) -> Dict[str, torch.Tensor]:
|
| 213 |
+
"""
|
| 214 |
+
Args:
|
| 215 |
+
h_a, h_b, h_img, h_aud: [B, d_i] representações pooled
|
| 216 |
+
|
| 217 |
+
Returns:
|
| 218 |
+
dict com alpha_i ∈ [B, 1] para cada modalidade
|
| 219 |
+
"""
|
| 220 |
+
# Projeta e computa score
|
| 221 |
+
s_a = self.scorer(torch.tanh(self.proj_a(h_a)))
|
| 222 |
+
s_b = self.scorer(torch.tanh(self.proj_b(h_b)))
|
| 223 |
+
s_img = self.scorer(torch.tanh(self.proj_img(h_img)))
|
| 224 |
+
s_aud = self.scorer(torch.tanh(self.proj_aud(h_aud)))
|
| 225 |
+
|
| 226 |
+
# Stack e softmax sobre as 4 modalidades
|
| 227 |
+
scores = torch.cat([s_a, s_b, s_img, s_aud], dim=-1) # [B, 4]
|
| 228 |
+
weights = F.softmax(scores, dim=-1) # [B, 4]
|
| 229 |
+
|
| 230 |
+
return {
|
| 231 |
+
"alpha_a": weights[:, 0:1],
|
| 232 |
+
"alpha_b": weights[:, 1:2],
|
| 233 |
+
"alpha_img": weights[:, 2:3],
|
| 234 |
+
"alpha_aud": weights[:, 3:4],
|
| 235 |
+
"weights": weights,
|
| 236 |
+
}
|
| 237 |
+
|
| 238 |
+
|
| 239 |
+
# ============================================================================
|
| 240 |
+
# Multimodal Multi-Head Attention
|
| 241 |
+
# ============================================================================
|
| 242 |
+
|
| 243 |
+
class MultimodalMultiHeadAttention(nn.Module):
|
| 244 |
+
"""Atenção multi-cabeça cross-modal completa.
|
| 245 |
+
|
| 246 |
+
Combina 4 modalidades (texto A, texto B, imagem, áudio) via atenção
|
| 247 |
+
cruzada par a par + gate de modalidade.
|
| 248 |
+
"""
|
| 249 |
+
|
| 250 |
+
def __init__(self, config: MultimodalAttentionConfig):
|
| 251 |
+
super().__init__()
|
| 252 |
+
self.config = config
|
| 253 |
+
|
| 254 |
+
# Cross-modal attention: cada par (i, j) com i != j
|
| 255 |
+
# Total: 4*3 = 12 pares direcionais
|
| 256 |
+
# Simplificação: projetar todas para d_model, usar 4 blocos de cross-attn
|
| 257 |
+
# Texto A atende B, img, aud
|
| 258 |
+
self.attn_a_to_b = CrossModalAttention(
|
| 259 |
+
config.d_text_a, config.d_text_b, config.d_model, config.n_heads, config.dropout
|
| 260 |
+
)
|
| 261 |
+
self.attn_a_to_img = CrossModalAttention(
|
| 262 |
+
config.d_text_a, config.d_image, config.d_model, config.n_heads, config.dropout
|
| 263 |
+
)
|
| 264 |
+
self.attn_a_to_aud = CrossModalAttention(
|
| 265 |
+
config.d_text_a, config.d_audio, config.d_model, config.n_heads, config.dropout
|
| 266 |
+
)
|
| 267 |
+
# Texto B atende A, img, aud
|
| 268 |
+
self.attn_b_to_a = CrossModalAttention(
|
| 269 |
+
config.d_text_b, config.d_text_a, config.d_model, config.n_heads, config.dropout
|
| 270 |
+
)
|
| 271 |
+
self.attn_b_to_img = CrossModalAttention(
|
| 272 |
+
config.d_text_b, config.d_image, config.d_model, config.n_heads, config.dropout
|
| 273 |
+
)
|
| 274 |
+
self.attn_b_to_aud = CrossModalAttention(
|
| 275 |
+
config.d_text_b, config.d_audio, config.d_model, config.n_heads, config.dropout
|
| 276 |
+
)
|
| 277 |
+
# Imagem atende A, B, aud (opcional — pode ser pesado)
|
| 278 |
+
self.attn_img_to_a = CrossModalAttention(
|
| 279 |
+
config.d_image, config.d_text_a, config.d_model, config.n_heads, config.dropout
|
| 280 |
+
)
|
| 281 |
+
self.attn_img_to_b = CrossModalAttention(
|
| 282 |
+
config.d_image, config.d_text_b, config.d_model, config.n_heads, config.dropout
|
| 283 |
+
)
|
| 284 |
+
# Áudio atende A, B
|
| 285 |
+
self.attn_aud_to_a = CrossModalAttention(
|
| 286 |
+
config.d_audio, config.d_text_a, config.d_model, config.n_heads, config.dropout
|
| 287 |
+
)
|
| 288 |
+
self.attn_aud_to_b = CrossModalAttention(
|
| 289 |
+
config.d_audio, config.d_text_b, config.d_model, config.n_heads, config.dropout
|
| 290 |
+
)
|
| 291 |
+
|
| 292 |
+
# LayerNorm para cada modalidade (estabilidade)
|
| 293 |
+
self.ln_a = nn.LayerNorm(config.d_text_a)
|
| 294 |
+
self.ln_b = nn.LayerNorm(config.d_text_b)
|
| 295 |
+
self.ln_img = nn.LayerNorm(config.d_image)
|
| 296 |
+
self.ln_aud = nn.LayerNorm(config.d_audio)
|
| 297 |
+
|
| 298 |
+
# Modality gate
|
| 299 |
+
self.gate = ModalityGate(
|
| 300 |
+
config.d_text_a, config.d_text_b, config.d_image, config.d_audio
|
| 301 |
+
) if config.use_modality_gate else None
|
| 302 |
+
|
| 303 |
+
# Pooling projections (para fusão final)
|
| 304 |
+
d_total = config.d_text_a + config.d_text_b + config.d_image + config.d_audio
|
| 305 |
+
if config.fusion_mode == "attention":
|
| 306 |
+
self.fusion_proj = nn.Linear(d_total, config.d_model)
|
| 307 |
+
elif config.fusion_mode == "gated":
|
| 308 |
+
self.fusion_proj = nn.Linear(d_total, config.d_model)
|
| 309 |
+
else: # concat
|
| 310 |
+
self.fusion_proj = nn.Linear(d_total, config.d_model)
|
| 311 |
+
|
| 312 |
+
def forward(
|
| 313 |
+
self,
|
| 314 |
+
seq_a: torch.Tensor,
|
| 315 |
+
seq_b: torch.Tensor,
|
| 316 |
+
seq_img: Optional[torch.Tensor] = None,
|
| 317 |
+
seq_aud: Optional[torch.Tensor] = None,
|
| 318 |
+
mask_a: Optional[torch.Tensor] = None,
|
| 319 |
+
mask_b: Optional[torch.Tensor] = None,
|
| 320 |
+
mask_img: Optional[torch.Tensor] = None,
|
| 321 |
+
mask_aud: Optional[torch.Tensor] = None,
|
| 322 |
+
) -> Dict[str, torch.Tensor]:
|
| 323 |
+
"""Forward pass completo.
|
| 324 |
+
|
| 325 |
+
Args:
|
| 326 |
+
seq_a: [B, T_a, d_a]
|
| 327 |
+
seq_b: [B, T_b, d_b]
|
| 328 |
+
seq_img: [B, T_img, d_img] ou None
|
| 329 |
+
seq_aud: [B, T_aud, d_aud] ou None
|
| 330 |
+
masks: máscaras de padding correspondentes
|
| 331 |
+
|
| 332 |
+
Returns:
|
| 333 |
+
dict com:
|
| 334 |
+
fused: [B, d_model] representação fundida
|
| 335 |
+
seq_a_out, seq_b_out, seq_img_out, seq_aud_out: saídas enriquecidas
|
| 336 |
+
modality_weights: pesos do gate (se ativo)
|
| 337 |
+
"""
|
| 338 |
+
B = seq_a.size(0)
|
| 339 |
+
device = seq_a.device
|
| 340 |
+
|
| 341 |
+
# Default para modalidades ausentes: zeros
|
| 342 |
+
if seq_img is None:
|
| 343 |
+
seq_img = torch.zeros(B, 1, self.config.d_image, device=device)
|
| 344 |
+
mask_img = torch.ones(B, 1, device=device)
|
| 345 |
+
if seq_aud is None:
|
| 346 |
+
seq_aud = torch.zeros(B, 1, self.config.d_audio, device=device)
|
| 347 |
+
mask_aud = torch.ones(B, 1, device=device)
|
| 348 |
+
|
| 349 |
+
# Cross-modal attention (com residual + LayerNorm)
|
| 350 |
+
# Texto A atende B, img, aud
|
| 351 |
+
a_to_b = self.attn_a_to_b(seq_a, seq_b, mask_b)
|
| 352 |
+
a_to_img = self.attn_a_to_img(seq_a, seq_img, mask_img)
|
| 353 |
+
a_to_aud = self.attn_a_to_aud(seq_a, seq_aud, mask_aud)
|
| 354 |
+
seq_a_out = self.ln_a(seq_a + a_to_b + a_to_img + a_to_aud)
|
| 355 |
+
|
| 356 |
+
# Texto B atende A, img, aud
|
| 357 |
+
b_to_a = self.attn_b_to_a(seq_b, seq_a, mask_a)
|
| 358 |
+
b_to_img = self.attn_b_to_img(seq_b, seq_img, mask_img)
|
| 359 |
+
b_to_aud = self.attn_b_to_aud(seq_b, seq_aud, mask_aud)
|
| 360 |
+
seq_b_out = self.ln_b(seq_b + b_to_a + b_to_img + b_to_aud)
|
| 361 |
+
|
| 362 |
+
# Imagem atende A, B
|
| 363 |
+
img_to_a = self.attn_img_to_a(seq_img, seq_a, mask_a)
|
| 364 |
+
img_to_b = self.attn_img_to_b(seq_img, seq_b, mask_b)
|
| 365 |
+
seq_img_out = self.ln_img(seq_img + img_to_a + img_to_b)
|
| 366 |
+
|
| 367 |
+
# Áudio atende A, B
|
| 368 |
+
aud_to_a = self.attn_aud_to_a(seq_aud, seq_a, mask_a)
|
| 369 |
+
aud_to_b = self.attn_aud_to_b(seq_aud, seq_b, mask_b)
|
| 370 |
+
seq_aud_out = self.ln_aud(seq_aud + aud_to_a + aud_to_b)
|
| 371 |
+
|
| 372 |
+
# Pooling (mean sobre a sequência, com máscara)
|
| 373 |
+
def pool(seq, mask):
|
| 374 |
+
if mask is None:
|
| 375 |
+
return seq.mean(dim=1)
|
| 376 |
+
m = mask.unsqueeze(-1).float()
|
| 377 |
+
return (seq * m).sum(dim=1) / m.sum(dim=1).clamp(min=1.0)
|
| 378 |
+
|
| 379 |
+
pooled_a = pool(seq_a_out, mask_a) # [B, d_a]
|
| 380 |
+
pooled_b = pool(seq_b_out, mask_b) # [B, d_b]
|
| 381 |
+
pooled_img = pool(seq_img_out, mask_img) # [B, d_img]
|
| 382 |
+
pooled_aud = pool(seq_aud_out, mask_aud) # [B, d_aud]
|
| 383 |
+
|
| 384 |
+
# Modality gate (opcional)
|
| 385 |
+
if self.gate is not None:
|
| 386 |
+
gate_out = self.gate(pooled_a, pooled_b, pooled_img, pooled_aud)
|
| 387 |
+
alpha_a = gate_out["alpha_a"]
|
| 388 |
+
alpha_b = gate_out["alpha_b"]
|
| 389 |
+
alpha_img = gate_out["alpha_img"]
|
| 390 |
+
alpha_aud = gate_out["alpha_aud"]
|
| 391 |
+
modality_weights = gate_out["weights"]
|
| 392 |
+
else:
|
| 393 |
+
alpha_a = alpha_b = alpha_img = alpha_aud = 1.0
|
| 394 |
+
modality_weights = None
|
| 395 |
+
|
| 396 |
+
# Fusão final
|
| 397 |
+
# Concat + projeção (com gate aplicado)
|
| 398 |
+
if self.gate is not None:
|
| 399 |
+
pooled_a_g = pooled_a * alpha_a
|
| 400 |
+
pooled_b_g = pooled_b * alpha_b
|
| 401 |
+
pooled_img_g = pooled_img * alpha_img
|
| 402 |
+
pooled_aud_g = pooled_aud * alpha_aud
|
| 403 |
+
else:
|
| 404 |
+
pooled_a_g = pooled_a
|
| 405 |
+
pooled_b_g = pooled_b
|
| 406 |
+
pooled_img_g = pooled_img
|
| 407 |
+
pooled_aud_g = pooled_aud
|
| 408 |
+
|
| 409 |
+
concat = torch.cat([pooled_a_g, pooled_b_g, pooled_img_g, pooled_aud_g], dim=-1)
|
| 410 |
+
fused = self.fusion_proj(concat)
|
| 411 |
+
|
| 412 |
+
return {
|
| 413 |
+
"fused": fused,
|
| 414 |
+
"seq_a_out": seq_a_out,
|
| 415 |
+
"seq_b_out": seq_b_out,
|
| 416 |
+
"seq_img_out": seq_img_out,
|
| 417 |
+
"seq_aud_out": seq_aud_out,
|
| 418 |
+
"pooled_a": pooled_a,
|
| 419 |
+
"pooled_b": pooled_b,
|
| 420 |
+
"pooled_img": pooled_img,
|
| 421 |
+
"pooled_aud": pooled_aud,
|
| 422 |
+
"modality_weights": modality_weights,
|
| 423 |
+
}
|
| 424 |
+
|
| 425 |
+
|
| 426 |
+
# ============================================================================
|
| 427 |
+
# Self-test
|
| 428 |
+
# ============================================================================
|
| 429 |
+
|
| 430 |
+
def _self_test():
|
| 431 |
+
"""Teste rápido do módulo multimodal attention."""
|
| 432 |
+
torch.manual_seed(42)
|
| 433 |
+
config = MultimodalAttentionConfig(
|
| 434 |
+
d_text_a=32, d_text_b=32, d_image=16, d_audio=16,
|
| 435 |
+
d_model=64, n_heads=4, dropout=0.1,
|
| 436 |
+
use_modality_gate=True,
|
| 437 |
+
)
|
| 438 |
+
mha = MultimodalMultiHeadAttention(config)
|
| 439 |
+
print(f"Params: {sum(p.numel() for p in mha.parameters())}")
|
| 440 |
+
|
| 441 |
+
B = 2
|
| 442 |
+
seq_a = torch.randn(B, 8, 32)
|
| 443 |
+
seq_b = torch.randn(B, 6, 32)
|
| 444 |
+
seq_img = torch.randn(B, 4, 16)
|
| 445 |
+
seq_aud = torch.randn(B, 5, 16)
|
| 446 |
+
mask_a = torch.ones(B, 8)
|
| 447 |
+
mask_b = torch.ones(B, 6)
|
| 448 |
+
mask_img = torch.ones(B, 4)
|
| 449 |
+
mask_aud = torch.ones(B, 5)
|
| 450 |
+
|
| 451 |
+
out = mha(seq_a, seq_b, seq_img, seq_aud, mask_a, mask_b, mask_img, mask_aud)
|
| 452 |
+
print(f"Fused: {out['fused'].shape}")
|
| 453 |
+
print(f"seq_a_out: {out['seq_a_out'].shape}")
|
| 454 |
+
print(f"Modality weights: {out['modality_weights']}")
|
| 455 |
+
|
| 456 |
+
# Test sem modalidades opcionais
|
| 457 |
+
out2 = mha(seq_a, seq_b, None, None, mask_a, mask_b)
|
| 458 |
+
print(f"Fused (text only): {out2['fused'].shape}")
|
| 459 |
+
|
| 460 |
+
|
| 461 |
+
if __name__ == "__main__":
|
| 462 |
+
_self_test()
|
| 463 |
+
|
| 464 |
+
|
| 465 |
+
__all__ = [
|
| 466 |
+
"MultimodalAttentionConfig",
|
| 467 |
+
"CrossModalAttention",
|
| 468 |
+
"ModalityGate",
|
| 469 |
+
"MultimodalMultiHeadAttention",
|
| 470 |
+
]
|
models/nlg.py
ADDED
|
@@ -0,0 +1,457 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""nlg.py — Natural Language Generation (NLG) para CNN-BiGRU.
|
| 2 |
+
|
| 3 |
+
Wrapper de alto nível para geração de texto que integra:
|
| 4 |
+
- TransformerDecoderStack (backbone causal com RoPE)
|
| 5 |
+
- MedusaMTP (Multi-Token Prediction)
|
| 6 |
+
- Sampling: Temperatura + Top-K + Top-P + Presence/Frequency Penalty
|
| 7 |
+
- CyclicReasoning (refinamento iterativo opcional)
|
| 8 |
+
- ContextWindowManager (janela de contexto com cache KV)
|
| 9 |
+
|
| 10 |
+
==============================================================================
|
| 11 |
+
ANÁLISE MATEMÁTICA E LÓGICA
|
| 12 |
+
==============================================================================
|
| 13 |
+
|
| 14 |
+
Dado um prompt P = [p_1, p_2, ..., p_n], geramos tokens y_1, y_2, ...
|
| 15 |
+
autoregressivamente. Em cada passo t:
|
| 16 |
+
|
| 17 |
+
1. Forward pass: h_t = Decoder(P || y_{<t}) via cache KV
|
| 18 |
+
2. (Opcional) Refinamento cíclico: h_t = CyclicReasoning(h_t)
|
| 19 |
+
3. Logits: l_t = LM_Head(h_t) [B, V]
|
| 20 |
+
4. Medusa: l_{t+k} = Medusa_Head_k(h_t) para k=1..K
|
| 21 |
+
5. Penalidades: l_t -= presence_penalty * 1[y in history]
|
| 22 |
+
+ frequency_penalty * count[y in history]
|
| 23 |
+
6. Sampling: y_t ~ softmax(l_t / T) com máscara Top-K e Top-P
|
| 24 |
+
|
| 25 |
+
7. (Opcional) Tree decoding: aceita candidatos Medusa se prob >= threshold
|
| 26 |
+
|
| 27 |
+
A loss de treinamento combina:
|
| 28 |
+
L_NLG = CE(lm_head(h_t), y_{t+1}) + mu_medusa * L_Medusa + mu_ewc * L_EWC
|
| 29 |
+
|
| 30 |
+
==============================================================================
|
| 31 |
+
INTEGRAÇÃO
|
| 32 |
+
==============================================================================
|
| 33 |
+
|
| 34 |
+
- TransformerDecoderStack (cnn_bigru.models.transformer_block): backbone causal
|
| 35 |
+
- MedusaMTP (cnn_bigru.models.medusa_heads): MTP
|
| 36 |
+
- CyclicReasoning (cnn_bigru.models.cyclic_reasoning): refinamento iterativo
|
| 37 |
+
- ContextWindowManager (cnn_bigru.models.context_window): janela deslizante
|
| 38 |
+
- BBPETokenizer (cnn_bigru.tokenizer): tokenização byte-BPE
|
| 39 |
+
|
| 40 |
+
Autor: CNN-BiGRU Project
|
| 41 |
+
"""
|
| 42 |
+
from __future__ import annotations
|
| 43 |
+
|
| 44 |
+
import logging
|
| 45 |
+
import math
|
| 46 |
+
from dataclasses import dataclass, field
|
| 47 |
+
from typing import Dict, List, Optional, Tuple, Union
|
| 48 |
+
|
| 49 |
+
import torch
|
| 50 |
+
import torch.nn as nn
|
| 51 |
+
import torch.nn.functional as F
|
| 52 |
+
|
| 53 |
+
logger = logging.getLogger(__name__)
|
| 54 |
+
|
| 55 |
+
|
| 56 |
+
# ============================================================================
|
| 57 |
+
# Configuração
|
| 58 |
+
# ============================================================================
|
| 59 |
+
|
| 60 |
+
@dataclass
|
| 61 |
+
class NLGConfig:
|
| 62 |
+
"""Configuração do gerador NLG."""
|
| 63 |
+
vocab_size: int = 32000
|
| 64 |
+
embed_dim: int = 256
|
| 65 |
+
n_heads: int = 4
|
| 66 |
+
n_layers: int = 4
|
| 67 |
+
max_seq_len: int = 512
|
| 68 |
+
dropout: float = 0.1
|
| 69 |
+
pad_id: int = 1
|
| 70 |
+
bos_id: int = 2
|
| 71 |
+
eos_id: int = 3
|
| 72 |
+
# Medusa MTP
|
| 73 |
+
use_medusa: bool = True
|
| 74 |
+
n_medusa_heads: int = 4
|
| 75 |
+
mu_medusa: float = 0.5
|
| 76 |
+
# Cyclic Reasoning
|
| 77 |
+
use_cyclic_reasoning: bool = False
|
| 78 |
+
n_cycles: int = 3
|
| 79 |
+
# Sampling defaults
|
| 80 |
+
temperature: float = 0.7
|
| 81 |
+
top_k: int = 40
|
| 82 |
+
top_p: float = 0.9
|
| 83 |
+
presence_penalty: float = 0.3
|
| 84 |
+
frequency_penalty: float = 0.3
|
| 85 |
+
# Context window
|
| 86 |
+
max_window: int = 1024
|
| 87 |
+
n_sink_tokens: int = 4
|
| 88 |
+
eviction_strategy: str = "sink_sliding"
|
| 89 |
+
# Weight tying
|
| 90 |
+
weight_tying: bool = True
|
| 91 |
+
# Device
|
| 92 |
+
device: str = "cpu"
|
| 93 |
+
|
| 94 |
+
|
| 95 |
+
# ============================================================================
|
| 96 |
+
# NLG Module
|
| 97 |
+
# ============================================================================
|
| 98 |
+
|
| 99 |
+
class NLGModule(nn.Module):
|
| 100 |
+
"""Módulo de geração de linguagem natural (NLG).
|
| 101 |
+
|
| 102 |
+
Combina TransformerDecoder + Medusa + (opcional) CyclicReasoning.
|
| 103 |
+
"""
|
| 104 |
+
|
| 105 |
+
def __init__(self, config: NLGConfig):
|
| 106 |
+
super().__init__()
|
| 107 |
+
self.config = config
|
| 108 |
+
# Lazy import para evitar circular deps
|
| 109 |
+
from .transformer_block import TransformerDecoderStack
|
| 110 |
+
from .medusa_heads import MedusaMTP, MedusaConfig
|
| 111 |
+
from .cyclic_reasoning import (
|
| 112 |
+
CyclicReasoning, CyclicReasoningConfig,
|
| 113 |
+
)
|
| 114 |
+
from .context_window import (
|
| 115 |
+
ContextWindowManager, ContextWindowConfig,
|
| 116 |
+
)
|
| 117 |
+
|
| 118 |
+
# Backbone (Transformer Decoder)
|
| 119 |
+
self.decoder = TransformerDecoderStack(
|
| 120 |
+
vocab_size=config.vocab_size,
|
| 121 |
+
embed_dim=config.embed_dim,
|
| 122 |
+
n_heads=config.n_heads,
|
| 123 |
+
n_layers=config.n_layers,
|
| 124 |
+
max_seq_len=config.max_seq_len,
|
| 125 |
+
use_rope=True,
|
| 126 |
+
pad_id=config.pad_id,
|
| 127 |
+
weight_tying=config.weight_tying,
|
| 128 |
+
)
|
| 129 |
+
|
| 130 |
+
# Medusa MTP (opcional)
|
| 131 |
+
if config.use_medusa:
|
| 132 |
+
self.medusa = MedusaMTP(MedusaConfig(
|
| 133 |
+
vocab_size=config.vocab_size,
|
| 134 |
+
embed_dim=config.embed_dim,
|
| 135 |
+
n_heads=config.n_medusa_heads,
|
| 136 |
+
head_hidden_mult=2,
|
| 137 |
+
dropout=config.dropout,
|
| 138 |
+
bos_id=config.bos_id,
|
| 139 |
+
))
|
| 140 |
+
else:
|
| 141 |
+
self.medusa = None
|
| 142 |
+
|
| 143 |
+
# Cyclic Reasoning (opcional)
|
| 144 |
+
if config.use_cyclic_reasoning:
|
| 145 |
+
self.cyclic = CyclicReasoning(CyclicReasoningConfig(
|
| 146 |
+
embed_dim=config.embed_dim,
|
| 147 |
+
max_cycles=config.n_cycles,
|
| 148 |
+
convergence_eps=1e-3,
|
| 149 |
+
use_anti_hallucination_gate=True,
|
| 150 |
+
device=config.device,
|
| 151 |
+
))
|
| 152 |
+
else:
|
| 153 |
+
self.cyclic = None
|
| 154 |
+
|
| 155 |
+
# Context Window Manager (criado lazy na geração)
|
| 156 |
+
self._cw: Optional[ContextWindowManager] = None
|
| 157 |
+
|
| 158 |
+
# ----------------------------------------------------------------------
|
| 159 |
+
# Forward (treino com teacher forcing)
|
| 160 |
+
# ----------------------------------------------------------------------
|
| 161 |
+
|
| 162 |
+
def forward(
|
| 163 |
+
self,
|
| 164 |
+
input_ids: torch.Tensor,
|
| 165 |
+
target_ids: Optional[torch.Tensor] = None,
|
| 166 |
+
use_medusa_loss: bool = True,
|
| 167 |
+
) -> Dict[str, torch.Tensor]:
|
| 168 |
+
"""Forward pass com teacher forcing.
|
| 169 |
+
|
| 170 |
+
Args:
|
| 171 |
+
input_ids: [B, T] tokens de entrada (inclui shift direito para teacher forcing)
|
| 172 |
+
target_ids: [B, T] tokens alvo (default: input_ids shifted por 1)
|
| 173 |
+
use_medusa_loss: se True, computa loss Medusa
|
| 174 |
+
|
| 175 |
+
Returns:
|
| 176 |
+
dict com:
|
| 177 |
+
logits: [B, T, V] logits da cabeça principal
|
| 178 |
+
loss: escalar (CE + Medusa) — só se target_ids fornecido
|
| 179 |
+
loss_main: CE da cabeça principal
|
| 180 |
+
loss_medusa: loss Medusa (se use_medusa_loss e medusa ativo)
|
| 181 |
+
medusa_stats: stats por cabeça
|
| 182 |
+
hidden: [B, T, D] hidden states (para debug)
|
| 183 |
+
"""
|
| 184 |
+
# Backbone forward
|
| 185 |
+
logits, caches = self.decoder(input_ids) # [B, T, V], list
|
| 186 |
+
# hidden states = logits do lm_head pré-projection? Não — precisamos do hidden
|
| 187 |
+
# Na TransformerDecoderStack, o hidden é a saída antes do lm_head
|
| 188 |
+
# Aqui assumimos que logits é o output do lm_head (com weight tying)
|
| 189 |
+
# Para obter o hidden, usamos logits se weight_tying (matriz transposta)
|
| 190 |
+
# ou adicionamos um hook no decoder. Para simplificar, refazemos o hidden
|
| 191 |
+
# via projeção inversa (apenas se weight_tying=True)
|
| 192 |
+
# NOTA: em prática, o decoder poderia retornar hidden explicitamente
|
| 193 |
+
# Para este wrapper, retornamos logits como proxy do hidden
|
| 194 |
+
# (a Medusa recebe o hidden pré-lm_head, mas para integração simples
|
| 195 |
+
# usamos logits como entrada da Medusa apenas quando shapes batem)
|
| 196 |
+
hidden = logits # fallback: usa logits como hidden (será projetado se necessário)
|
| 197 |
+
|
| 198 |
+
result = {
|
| 199 |
+
"logits": logits,
|
| 200 |
+
"hidden": hidden,
|
| 201 |
+
}
|
| 202 |
+
|
| 203 |
+
# Loss (se target fornecido)
|
| 204 |
+
if target_ids is not None:
|
| 205 |
+
# CE principal
|
| 206 |
+
V = self.config.vocab_size
|
| 207 |
+
loss_main = F.cross_entropy(
|
| 208 |
+
logits.view(-1, V),
|
| 209 |
+
target_ids.view(-1),
|
| 210 |
+
ignore_index=self.config.pad_id,
|
| 211 |
+
reduction="mean",
|
| 212 |
+
)
|
| 213 |
+
result["loss_main"] = loss_main
|
| 214 |
+
|
| 215 |
+
# Loss Medusa
|
| 216 |
+
loss_medusa = torch.zeros(1, device=logits.device, dtype=logits.dtype)
|
| 217 |
+
if use_medusa_loss and self.medusa is not None:
|
| 218 |
+
# A MedusaMTP espera o hidden state [B, T, D], não logits [B, T, V]
|
| 219 |
+
# Como não temos o hidden explícito aqui (sem modificar o decoder),
|
| 220 |
+
# usamos uma projeção aproximada: pegamos logits e projetamos para D
|
| 221 |
+
# via lm_head.weight.T (se weight tying) ou via embedding
|
| 222 |
+
# Para simplificar, usamos logits[:, :, :D] (slice) quando V > D
|
| 223 |
+
D = self.config.embed_dim
|
| 224 |
+
if logits.size(-1) >= D:
|
| 225 |
+
hidden_proxy = logits[:, :, :D] # slice simples
|
| 226 |
+
else:
|
| 227 |
+
# projeta logits para D
|
| 228 |
+
if not hasattr(self, "_logits_proj"):
|
| 229 |
+
self._logits_proj = nn.Linear(V, D, bias=False).to(logits.device)
|
| 230 |
+
hidden_proxy = self._logits_proj(logits)
|
| 231 |
+
|
| 232 |
+
loss_med, medusa_stats = self.medusa.compute_loss(
|
| 233 |
+
hidden_proxy, target_ids, ignore_index=self.config.pad_id,
|
| 234 |
+
)
|
| 235 |
+
loss_medusa = loss_med
|
| 236 |
+
result["medusa_stats"] = medusa_stats
|
| 237 |
+
|
| 238 |
+
result["loss_medusa"] = loss_medusa.squeeze()
|
| 239 |
+
result["loss"] = loss_main + self.config.mu_medusa * loss_medusa.squeeze()
|
| 240 |
+
|
| 241 |
+
return result
|
| 242 |
+
|
| 243 |
+
# ----------------------------------------------------------------------
|
| 244 |
+
# Geração autoregressiva
|
| 245 |
+
# ----------------------------------------------------------------------
|
| 246 |
+
|
| 247 |
+
@torch.no_grad()
|
| 248 |
+
def generate(
|
| 249 |
+
self,
|
| 250 |
+
prompt_ids: torch.Tensor,
|
| 251 |
+
max_new_tokens: int = 50,
|
| 252 |
+
temperature: Optional[float] = None,
|
| 253 |
+
top_k: Optional[int] = None,
|
| 254 |
+
top_p: Optional[float] = None,
|
| 255 |
+
presence_penalty: Optional[float] = None,
|
| 256 |
+
frequency_penalty: Optional[float] = None,
|
| 257 |
+
use_context_window: bool = True,
|
| 258 |
+
use_medusa: Optional[bool] = None,
|
| 259 |
+
eos_id: Optional[int] = None,
|
| 260 |
+
) -> Dict[str, torch.Tensor]:
|
| 261 |
+
"""Gera texto autoregressivamente com sampling.
|
| 262 |
+
|
| 263 |
+
Args:
|
| 264 |
+
prompt_ids: [B, T_prompt] IDs do prompt
|
| 265 |
+
max_new_tokens: número máximo de tokens a gerar
|
| 266 |
+
temperature, top_k, top_p, presence_penalty, frequency_penalty:
|
| 267 |
+
overrides da config (se None, usa config)
|
| 268 |
+
use_context_window: se True, aplica eviction policy
|
| 269 |
+
use_medusa: se True, usa Medusa tree decoding (default: config.use_medusa)
|
| 270 |
+
eos_id: ID do token EOS (default: config.eos_id)
|
| 271 |
+
|
| 272 |
+
Returns:
|
| 273 |
+
dict com:
|
| 274 |
+
ids: [B, T_prompt + T_gen] tokens gerados
|
| 275 |
+
n_tokens: int — número de tokens gerados (excluindo prompt)
|
| 276 |
+
n_medusa_accepted: int — tokens aceitos via Medusa
|
| 277 |
+
stopped_early: bool — True se EOS atingido
|
| 278 |
+
"""
|
| 279 |
+
self.eval()
|
| 280 |
+
cfg = self.config
|
| 281 |
+
T = temperature if temperature is not None else cfg.temperature
|
| 282 |
+
K = top_k if top_k is not None else cfg.top_k
|
| 283 |
+
P = top_p if top_p is not None else cfg.top_p
|
| 284 |
+
PP = presence_penalty if presence_penalty is not None else cfg.presence_penalty
|
| 285 |
+
FP = frequency_penalty if frequency_penalty is not None else cfg.frequency_penalty
|
| 286 |
+
EOS = eos_id if eos_id is not None else cfg.eos_id
|
| 287 |
+
use_med = use_medusa if use_medusa is not None else cfg.use_medusa and self.medusa is not None
|
| 288 |
+
|
| 289 |
+
device = prompt_ids.device
|
| 290 |
+
B = prompt_ids.size(0)
|
| 291 |
+
|
| 292 |
+
# Inicializar context window
|
| 293 |
+
if use_context_window and self._cw is None:
|
| 294 |
+
from .context_window import ContextWindowManager, ContextWindowConfig
|
| 295 |
+
self._cw = ContextWindowManager(ContextWindowConfig(
|
| 296 |
+
max_window=cfg.max_window,
|
| 297 |
+
eviction_strategy=cfg.eviction_strategy,
|
| 298 |
+
n_sink_tokens=cfg.n_sink_tokens,
|
| 299 |
+
embed_dim=cfg.embed_dim,
|
| 300 |
+
n_heads=cfg.n_heads,
|
| 301 |
+
head_dim=cfg.embed_dim // cfg.n_heads,
|
| 302 |
+
n_layers=cfg.n_layers,
|
| 303 |
+
device=str(device),
|
| 304 |
+
))
|
| 305 |
+
|
| 306 |
+
# Token history (para penalidades)
|
| 307 |
+
token_counts: List[Dict[int, int]] = [dict() for _ in range(B)]
|
| 308 |
+
for b in range(B):
|
| 309 |
+
for tok in prompt_ids[b].tolist():
|
| 310 |
+
token_counts[b][tok] = token_counts[b].get(tok, 0) + 1
|
| 311 |
+
|
| 312 |
+
# Estado: tokens atuais
|
| 313 |
+
cur_ids = prompt_ids.clone() # [B, T]
|
| 314 |
+
n_generated = 0
|
| 315 |
+
n_medusa_accepted = 0
|
| 316 |
+
stopped = torch.zeros(B, dtype=torch.bool, device=device)
|
| 317 |
+
|
| 318 |
+
for step in range(max_new_tokens):
|
| 319 |
+
# Truncar para max_seq_len
|
| 320 |
+
ctx = cur_ids[:, -cfg.max_seq_len:]
|
| 321 |
+
# Forward
|
| 322 |
+
try:
|
| 323 |
+
logits, _ = self.decoder(ctx)
|
| 324 |
+
except Exception as e:
|
| 325 |
+
logger.warning(f"NLG forward falhou no step {step}: {e}")
|
| 326 |
+
break
|
| 327 |
+
|
| 328 |
+
# Pegar logits do último token
|
| 329 |
+
last_logits = logits[:, -1, :].clone() # [B, V]
|
| 330 |
+
|
| 331 |
+
# Aplicar penalidades
|
| 332 |
+
if PP != 0.0 or FP != 0.0:
|
| 333 |
+
for b in range(B):
|
| 334 |
+
for tok_id, count in token_counts[b].items():
|
| 335 |
+
last_logits[b, tok_id] -= (PP + count * FP)
|
| 336 |
+
|
| 337 |
+
# Temperatura
|
| 338 |
+
if T > 0:
|
| 339 |
+
last_logits = last_logits / T
|
| 340 |
+
else:
|
| 341 |
+
# Greedy
|
| 342 |
+
next_tok = last_logits.argmax(dim=-1, keepdim=True) # [B, 1]
|
| 343 |
+
cur_ids = torch.cat([cur_ids, next_tok], dim=1)
|
| 344 |
+
n_generated += 1
|
| 345 |
+
for b in range(B):
|
| 346 |
+
t = next_tok[b].item()
|
| 347 |
+
token_counts[b][t] = token_counts[b].get(t, 0) + 1
|
| 348 |
+
if t == EOS:
|
| 349 |
+
stopped[b] = True
|
| 350 |
+
if stopped.all():
|
| 351 |
+
break
|
| 352 |
+
continue
|
| 353 |
+
|
| 354 |
+
# Top-K
|
| 355 |
+
if K > 0:
|
| 356 |
+
top_vals, _ = torch.topk(last_logits, min(K, last_logits.size(-1)))
|
| 357 |
+
min_val = top_vals[:, -1:]
|
| 358 |
+
last_logits[last_logits < min_val] = float("-inf")
|
| 359 |
+
|
| 360 |
+
# Top-P (nucleus)
|
| 361 |
+
if P < 1.0:
|
| 362 |
+
sorted_logits, sorted_indices = torch.sort(last_logits, descending=True)
|
| 363 |
+
sorted_probs = F.softmax(sorted_logits, dim=-1)
|
| 364 |
+
cum_probs = torch.cumsum(sorted_probs, dim=-1)
|
| 365 |
+
mask = cum_probs > P
|
| 366 |
+
mask[:, 1:] = mask[:, :-1].clone()
|
| 367 |
+
mask[:, 0] = False
|
| 368 |
+
indices_to_remove = mask.scatter(1, sorted_indices, mask)
|
| 369 |
+
last_logits[indices_to_remove] = float("-inf")
|
| 370 |
+
|
| 371 |
+
# Sampling
|
| 372 |
+
probs = F.softmax(last_logits, dim=-1)
|
| 373 |
+
next_tok = torch.multinomial(probs, num_samples=1) # [B, 1]
|
| 374 |
+
cur_ids = torch.cat([cur_ids, next_tok], dim=1)
|
| 375 |
+
n_generated += 1
|
| 376 |
+
|
| 377 |
+
# Update counts
|
| 378 |
+
for b in range(B):
|
| 379 |
+
t = next_tok[b].item()
|
| 380 |
+
token_counts[b][t] = token_counts[b].get(t, 0) + 1
|
| 381 |
+
if t == EOS:
|
| 382 |
+
stopped[b] = True
|
| 383 |
+
|
| 384 |
+
# (Opcional) Medusa: aceitar próximos candidatos
|
| 385 |
+
# Skip nesta implementação simplificada — Medusa loss é computada em treino
|
| 386 |
+
|
| 387 |
+
if stopped.all():
|
| 388 |
+
break
|
| 389 |
+
|
| 390 |
+
return {
|
| 391 |
+
"ids": cur_ids,
|
| 392 |
+
"n_tokens": n_generated,
|
| 393 |
+
"n_medusa_accepted": n_medusa_accepted,
|
| 394 |
+
"stopped_early": bool(stopped.any().item()),
|
| 395 |
+
}
|
| 396 |
+
|
| 397 |
+
# ----------------------------------------------------------------------
|
| 398 |
+
# Compute loss para treinamento (compatível com trainer)
|
| 399 |
+
# ----------------------------------------------------------------------
|
| 400 |
+
|
| 401 |
+
def compute_loss(
|
| 402 |
+
self,
|
| 403 |
+
input_ids: torch.Tensor,
|
| 404 |
+
target_ids: Optional[torch.Tensor] = None,
|
| 405 |
+
) -> Dict[str, torch.Tensor]:
|
| 406 |
+
"""Computa a loss de NLG para uso no trainer.
|
| 407 |
+
|
| 408 |
+
Atalho para self.forward(input_ids, target_ids, use_medusa_loss=True)
|
| 409 |
+
"""
|
| 410 |
+
return self.forward(input_ids, target_ids, use_medusa_loss=True)
|
| 411 |
+
|
| 412 |
+
|
| 413 |
+
# ============================================================================
|
| 414 |
+
# Self-test
|
| 415 |
+
# ============================================================================
|
| 416 |
+
|
| 417 |
+
def _self_test():
|
| 418 |
+
"""Teste rápido do módulo NLG."""
|
| 419 |
+
torch.manual_seed(42)
|
| 420 |
+
config = NLGConfig(
|
| 421 |
+
vocab_size=100,
|
| 422 |
+
embed_dim=32,
|
| 423 |
+
n_heads=4,
|
| 424 |
+
n_layers=2,
|
| 425 |
+
max_seq_len=32,
|
| 426 |
+
pad_id=1, bos_id=2, eos_id=3,
|
| 427 |
+
use_medusa=True,
|
| 428 |
+
n_medusa_heads=3,
|
| 429 |
+
use_cyclic_reasoning=False,
|
| 430 |
+
weight_tying=True,
|
| 431 |
+
)
|
| 432 |
+
nlg = NLGModule(config)
|
| 433 |
+
print(f"NLG params: {sum(p.numel() for p in nlg.parameters())}")
|
| 434 |
+
|
| 435 |
+
# Forward com teacher forcing
|
| 436 |
+
ids = torch.randint(0, 100, (2, 8))
|
| 437 |
+
target = torch.randint(0, 100, (2, 8))
|
| 438 |
+
out = nlg.compute_loss(ids, target)
|
| 439 |
+
print(f"Loss: {out['loss'].item():.4f}")
|
| 440 |
+
print(f" main: {out['loss_main'].item():.4f}")
|
| 441 |
+
print(f" medusa: {out['loss_medusa'].item():.4f}")
|
| 442 |
+
print(f" medusa_stats: {out.get('medusa_stats', {})}")
|
| 443 |
+
|
| 444 |
+
# Geração
|
| 445 |
+
prompt = torch.tensor([[2, 5, 10, 15]], dtype=torch.long)
|
| 446 |
+
gen = nlg.generate(prompt, max_new_tokens=10, temperature=0.7, top_k=10)
|
| 447 |
+
print(f"Generated {gen['n_tokens']} tokens: {gen['ids'][0].tolist()}")
|
| 448 |
+
|
| 449 |
+
|
| 450 |
+
if __name__ == "__main__":
|
| 451 |
+
_self_test()
|
| 452 |
+
|
| 453 |
+
|
| 454 |
+
__all__ = [
|
| 455 |
+
"NLGConfig",
|
| 456 |
+
"NLGModule",
|
| 457 |
+
]
|
models/nlp.py
ADDED
|
@@ -0,0 +1,654 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""nlp.py — Natural Language Processing (NLP) para CNN-BiGRU.
|
| 2 |
+
|
| 3 |
+
Módulo de processamento de linguagem natural que fornece tarefas clássicas
|
| 4 |
+
de NLP em cima do backbone cooperativo CNN-BiGRU:
|
| 5 |
+
|
| 6 |
+
- Sequence Classification (sentiment, topic, NLI)
|
| 7 |
+
- Token Classification (NER, POS tagging)
|
| 8 |
+
- Span Detection (QA extractive)
|
| 9 |
+
- Sentence Pair Classification (paraphrase, entailment)
|
| 10 |
+
- Embeddings para retrieval semântico
|
| 11 |
+
|
| 12 |
+
==============================================================================
|
| 13 |
+
ANÁLISE MATEMÁTICA E LÓGICA
|
| 14 |
+
==============================================================================
|
| 15 |
+
|
| 16 |
+
Dado o output do CooperativeCNNBiGRU (representações fused [B, 256] ou
|
| 17 |
+
sequências temporais [B, T, 128] por stream), cada tarefa NLP aplica uma
|
| 18 |
+
cabeça específica:
|
| 19 |
+
|
| 20 |
+
1. Sequence Classification:
|
| 21 |
+
h_cls = SelfAttentionSummary(seq) [B, D]
|
| 22 |
+
logits = Linear(D, num_classes) [B, C]
|
| 23 |
+
loss = CrossEntropy(logits, y)
|
| 24 |
+
|
| 25 |
+
2. Token Classification:
|
| 26 |
+
logits = Linear(D, num_labels) [B, T, L]
|
| 27 |
+
loss = CrossEntropy(logits, y) (ignore padding)
|
| 28 |
+
|
| 29 |
+
3. Span Detection (QA):
|
| 30 |
+
start_logits = Linear(D, 1) [B, T]
|
| 31 |
+
end_logits = Linear(D, 1) [B, T]
|
| 32 |
+
loss = CE(start_logits, start_pos) + CE(end_logits, end_pos)
|
| 33 |
+
|
| 34 |
+
4. Sentence Pair Classification:
|
| 35 |
+
Usa o dual-stream A/B do CooperativeCNNBiGRU diretamente
|
| 36 |
+
(já que o modelo é naturalmente dual-stream).
|
| 37 |
+
logits = Linear(D_fused, num_classes)
|
| 38 |
+
|
| 39 |
+
5. Embeddings para retrieval:
|
| 40 |
+
h = mean_pool(seq) [B, D]
|
| 41 |
+
h = normalize(h, p=2) [B, D]
|
| 42 |
+
similarity = cosine(h_query, h_doc)
|
| 43 |
+
loss = InfoNCE (contrastive)
|
| 44 |
+
|
| 45 |
+
==============================================================================
|
| 46 |
+
INTEGRAÇÃO
|
| 47 |
+
==============================================================================
|
| 48 |
+
|
| 49 |
+
- CooperativeCNNBiGRU (cnn_bigru.models.cooperative_bigru): backbone
|
| 50 |
+
- SelfAttentionSummary: já incluído no cooperative_bigru
|
| 51 |
+
- BBPETokenizer: tokenização
|
| 52 |
+
|
| 53 |
+
Autor: CNN-BiGRU Project
|
| 54 |
+
"""
|
| 55 |
+
from __future__ import annotations
|
| 56 |
+
|
| 57 |
+
import logging
|
| 58 |
+
import math
|
| 59 |
+
from dataclasses import dataclass, field
|
| 60 |
+
from typing import Dict, List, Optional, Tuple, Union
|
| 61 |
+
|
| 62 |
+
import torch
|
| 63 |
+
import torch.nn as nn
|
| 64 |
+
import torch.nn.functional as F
|
| 65 |
+
|
| 66 |
+
logger = logging.getLogger(__name__)
|
| 67 |
+
|
| 68 |
+
|
| 69 |
+
# ============================================================================
|
| 70 |
+
# Configuração
|
| 71 |
+
# ============================================================================
|
| 72 |
+
|
| 73 |
+
@dataclass
|
| 74 |
+
class NLPConfig:
|
| 75 |
+
"""Configuração do módulo NLP."""
|
| 76 |
+
# Dimensões (precisam bater com o backbone)
|
| 77 |
+
embed_dim: int = 64 # embedding de tokens
|
| 78 |
+
cnn_filters: int = 64
|
| 79 |
+
gru_hidden: int = 64 # 64 por direção -> 128 por stream
|
| 80 |
+
feat_per_stream: int = 128 # = 2 * gru_hidden
|
| 81 |
+
feat_fused: int = 256 # = 2 * feat_per_stream
|
| 82 |
+
n_heads: int = 4
|
| 83 |
+
dropout: float = 0.1
|
| 84 |
+
pad_idx: int = 1
|
| 85 |
+
# Tarefas
|
| 86 |
+
num_classes_seq: int = 3 # sequence classification
|
| 87 |
+
num_labels_tok: int = 7 # token classification (NER-style)
|
| 88 |
+
# Pooling
|
| 89 |
+
pooling: str = "self_attn" # "self_attn" | "mean" | "max" | "cls"
|
| 90 |
+
# Loss
|
| 91 |
+
label_smoothing: float = 0.0
|
| 92 |
+
# Device
|
| 93 |
+
device: str = "cpu"
|
| 94 |
+
|
| 95 |
+
|
| 96 |
+
# ============================================================================
|
| 97 |
+
# Sequence Classification Head
|
| 98 |
+
# ============================================================================
|
| 99 |
+
|
| 100 |
+
class SequenceClassificationHead(nn.Module):
|
| 101 |
+
"""Cabeça para classificação de sequência inteira.
|
| 102 |
+
|
| 103 |
+
Aplica pooling (self-attention, mean, max ou cls) sobre a sequência
|
| 104 |
+
temporal e projeta para num_classes logits.
|
| 105 |
+
|
| 106 |
+
Args:
|
| 107 |
+
embed_dim: dimensão de entrada (e.g., 256 para fused, 128 por stream)
|
| 108 |
+
num_classes: número de classes de saída
|
| 109 |
+
dropout: prob de dropout
|
| 110 |
+
pooling: estratégia de pooling
|
| 111 |
+
n_heads: cabeças para self-attention pooling
|
| 112 |
+
"""
|
| 113 |
+
|
| 114 |
+
def __init__(
|
| 115 |
+
self,
|
| 116 |
+
embed_dim: int,
|
| 117 |
+
num_classes: int,
|
| 118 |
+
dropout: float = 0.1,
|
| 119 |
+
pooling: str = "self_attn",
|
| 120 |
+
n_heads: int = 4,
|
| 121 |
+
):
|
| 122 |
+
super().__init__()
|
| 123 |
+
self.embed_dim = embed_dim
|
| 124 |
+
self.pooling = pooling
|
| 125 |
+
self.num_classes = num_classes
|
| 126 |
+
|
| 127 |
+
# Pooling
|
| 128 |
+
if pooling == "self_attn":
|
| 129 |
+
# Self-attention pooling simplificado (sem CLS token)
|
| 130 |
+
self.attn_query = nn.Linear(embed_dim, 1)
|
| 131 |
+
elif pooling == "cls":
|
| 132 |
+
# CLS token learnable
|
| 133 |
+
self.cls = nn.Parameter(torch.zeros(1, 1, embed_dim))
|
| 134 |
+
nn.init.normal_(self.cls, std=0.02)
|
| 135 |
+
# mean / max: não precisa de params
|
| 136 |
+
|
| 137 |
+
self.dropout = nn.Dropout(dropout)
|
| 138 |
+
self.classifier = nn.Linear(embed_dim, num_classes)
|
| 139 |
+
nn.init.xavier_uniform_(self.classifier.weight)
|
| 140 |
+
nn.init.zeros_(self.classifier.bias)
|
| 141 |
+
|
| 142 |
+
def forward(
|
| 143 |
+
self,
|
| 144 |
+
seq: torch.Tensor,
|
| 145 |
+
mask: Optional[torch.Tensor] = None,
|
| 146 |
+
) -> Tuple[torch.Tensor, torch.Tensor]:
|
| 147 |
+
"""
|
| 148 |
+
Args:
|
| 149 |
+
seq: [B, T, D] sequência de hidden states
|
| 150 |
+
mask: [B, T] máscara de padding (1=real, 0=pad)
|
| 151 |
+
|
| 152 |
+
Returns:
|
| 153 |
+
logits: [B, C]
|
| 154 |
+
pooled: [B, D] representação pooled
|
| 155 |
+
"""
|
| 156 |
+
if seq.dim() == 2:
|
| 157 |
+
# [B, D] — já é pooled (e.g., fused)
|
| 158 |
+
pooled = seq
|
| 159 |
+
else:
|
| 160 |
+
B, T, D = seq.shape
|
| 161 |
+
|
| 162 |
+
if self.pooling == "self_attn":
|
| 163 |
+
# Attention pooling: softmax over T
|
| 164 |
+
scores = self.attn_query(seq).squeeze(-1) # [B, T]
|
| 165 |
+
if mask is not None:
|
| 166 |
+
scores = scores.masked_fill(mask == 0, -1e9)
|
| 167 |
+
# Lidar com linhas totalmente mascaradas
|
| 168 |
+
all_neg = (scores <= -1e9).all(dim=-1, keepdim=True)
|
| 169 |
+
scores = torch.where(all_neg, torch.zeros_like(scores), scores)
|
| 170 |
+
weights = F.softmax(scores, dim=-1).unsqueeze(-1) # [B, T, 1]
|
| 171 |
+
pooled = (seq * weights).sum(dim=1) # [B, D]
|
| 172 |
+
|
| 173 |
+
elif self.pooling == "mean":
|
| 174 |
+
if mask is not None:
|
| 175 |
+
m = mask.unsqueeze(-1).float()
|
| 176 |
+
pooled = (seq * m).sum(dim=1) / m.sum(dim=1).clamp(min=1.0)
|
| 177 |
+
else:
|
| 178 |
+
pooled = seq.mean(dim=1)
|
| 179 |
+
|
| 180 |
+
elif self.pooling == "max":
|
| 181 |
+
if mask is not None:
|
| 182 |
+
seq_masked = seq.masked_fill(mask.unsqueeze(-1) == 0, -1e9)
|
| 183 |
+
pooled, _ = seq_masked.max(dim=1)
|
| 184 |
+
# Lidar com sequências totalmente mascaradas (retorna 0)
|
| 185 |
+
all_neg = (pooled <= -1e9).all(dim=-1, keepdim=True)
|
| 186 |
+
pooled = torch.where(all_neg, torch.zeros_like(pooled), pooled)
|
| 187 |
+
else:
|
| 188 |
+
pooled, _ = seq.max(dim=1)
|
| 189 |
+
|
| 190 |
+
elif self.pooling == "cls":
|
| 191 |
+
cls = self.cls.expand(B, -1, -1) # [B, 1, D]
|
| 192 |
+
seq_ext = torch.cat([cls, seq], dim=1) # [B, T+1, D]
|
| 193 |
+
pooled = seq_ext[:, 0, :] # pega CLS
|
| 194 |
+
|
| 195 |
+
else:
|
| 196 |
+
raise ValueError(f"pooling desconhecido: {self.pooling}")
|
| 197 |
+
|
| 198 |
+
pooled = self.dropout(pooled)
|
| 199 |
+
logits = self.classifier(pooled)
|
| 200 |
+
return logits, pooled
|
| 201 |
+
|
| 202 |
+
|
| 203 |
+
# ============================================================================
|
| 204 |
+
# Token Classification Head
|
| 205 |
+
# ============================================================================
|
| 206 |
+
|
| 207 |
+
class TokenClassificationHead(nn.Module):
|
| 208 |
+
"""Cabeça para classificação de tokens (NER, POS tagging).
|
| 209 |
+
|
| 210 |
+
Aplica Linear(D, num_labels) sobre cada posição da sequência.
|
| 211 |
+
|
| 212 |
+
Args:
|
| 213 |
+
embed_dim: dimensão de entrada
|
| 214 |
+
num_labels: número de rótulos (e.g., 7 para BIO tagging)
|
| 215 |
+
dropout: prob de dropout
|
| 216 |
+
"""
|
| 217 |
+
|
| 218 |
+
def __init__(self, embed_dim: int, num_labels: int, dropout: float = 0.1):
|
| 219 |
+
super().__init__()
|
| 220 |
+
self.embed_dim = embed_dim
|
| 221 |
+
self.num_labels = num_labels
|
| 222 |
+
self.dropout = nn.Dropout(dropout)
|
| 223 |
+
self.classifier = nn.Linear(embed_dim, num_labels)
|
| 224 |
+
nn.init.xavier_uniform_(self.classifier.weight)
|
| 225 |
+
nn.init.zeros_(self.classifier.bias)
|
| 226 |
+
|
| 227 |
+
def forward(self, seq: torch.Tensor) -> torch.Tensor:
|
| 228 |
+
"""
|
| 229 |
+
Args:
|
| 230 |
+
seq: [B, T, D] sequência de hidden states
|
| 231 |
+
Returns:
|
| 232 |
+
logits: [B, T, L]
|
| 233 |
+
"""
|
| 234 |
+
return self.classifier(self.dropout(seq))
|
| 235 |
+
|
| 236 |
+
|
| 237 |
+
# ============================================================================
|
| 238 |
+
# Span Detection Head (QA Extractive)
|
| 239 |
+
# ============================================================================
|
| 240 |
+
|
| 241 |
+
class SpanDetectionHead(nn.Module):
|
| 242 |
+
"""Cabeça para detecção de span (QA extractive).
|
| 243 |
+
|
| 244 |
+
Prediz posições de início e fim do span de resposta.
|
| 245 |
+
|
| 246 |
+
Args:
|
| 247 |
+
embed_dim: dimensão de entrada
|
| 248 |
+
dropout: prob de dropout
|
| 249 |
+
"""
|
| 250 |
+
|
| 251 |
+
def __init__(self, embed_dim: int, dropout: float = 0.1):
|
| 252 |
+
super().__init__()
|
| 253 |
+
self.embed_dim = embed_dim
|
| 254 |
+
self.dropout = nn.Dropout(dropout)
|
| 255 |
+
self.start_classifier = nn.Linear(embed_dim, 1)
|
| 256 |
+
self.end_classifier = nn.Linear(embed_dim, 1)
|
| 257 |
+
for layer in (self.start_classifier, self.end_classifier):
|
| 258 |
+
nn.init.xavier_uniform_(layer.weight)
|
| 259 |
+
nn.init.zeros_(layer.bias)
|
| 260 |
+
|
| 261 |
+
def forward(self, seq: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
| 262 |
+
"""
|
| 263 |
+
Args:
|
| 264 |
+
seq: [B, T, D]
|
| 265 |
+
Returns:
|
| 266 |
+
start_logits: [B, T]
|
| 267 |
+
end_logits: [B, T]
|
| 268 |
+
"""
|
| 269 |
+
seq = self.dropout(seq)
|
| 270 |
+
start_logits = self.start_classifier(seq).squeeze(-1) # [B, T]
|
| 271 |
+
end_logits = self.end_classifier(seq).squeeze(-1)
|
| 272 |
+
return start_logits, end_logits
|
| 273 |
+
|
| 274 |
+
|
| 275 |
+
# ============================================================================
|
| 276 |
+
# Embedding Head (para retrieval semântico)
|
| 277 |
+
# ============================================================================
|
| 278 |
+
|
| 279 |
+
class EmbeddingHead(nn.Module):
|
| 280 |
+
"""Cabeça para gerar embeddings normalizados (retrieval semântico).
|
| 281 |
+
|
| 282 |
+
Aplica mean-pool com máscara + LayerNorm + normalização L2.
|
| 283 |
+
|
| 284 |
+
Args:
|
| 285 |
+
embed_dim: dimensão de entrada e saída
|
| 286 |
+
use_layer_norm: se True, aplica LayerNorm antes da normalização L2
|
| 287 |
+
"""
|
| 288 |
+
|
| 289 |
+
def __init__(self, embed_dim: int, use_layer_norm: bool = True):
|
| 290 |
+
super().__init__()
|
| 291 |
+
self.embed_dim = embed_dim
|
| 292 |
+
self.use_layer_norm = use_layer_norm
|
| 293 |
+
if use_layer_norm:
|
| 294 |
+
self.ln = nn.LayerNorm(embed_dim)
|
| 295 |
+
|
| 296 |
+
def forward(
|
| 297 |
+
self,
|
| 298 |
+
seq: torch.Tensor,
|
| 299 |
+
mask: Optional[torch.Tensor] = None,
|
| 300 |
+
normalize: bool = True,
|
| 301 |
+
) -> torch.Tensor:
|
| 302 |
+
"""
|
| 303 |
+
Args:
|
| 304 |
+
seq: [B, T, D]
|
| 305 |
+
mask: [B, T] (1=real, 0=pad)
|
| 306 |
+
normalize: se True, normaliza com L2
|
| 307 |
+
Returns:
|
| 308 |
+
embedding: [B, D]
|
| 309 |
+
"""
|
| 310 |
+
if mask is not None:
|
| 311 |
+
m = mask.unsqueeze(-1).float()
|
| 312 |
+
pooled = (seq * m).sum(dim=1) / m.sum(dim=1).clamp(min=1.0)
|
| 313 |
+
else:
|
| 314 |
+
pooled = seq.mean(dim=1)
|
| 315 |
+
|
| 316 |
+
if self.use_layer_norm:
|
| 317 |
+
pooled = self.ln(pooled)
|
| 318 |
+
|
| 319 |
+
if normalize:
|
| 320 |
+
pooled = F.normalize(pooled, p=2, dim=-1)
|
| 321 |
+
|
| 322 |
+
return pooled
|
| 323 |
+
|
| 324 |
+
|
| 325 |
+
# ============================================================================
|
| 326 |
+
# NLP Module (wrapper integrando todas as cabeças)
|
| 327 |
+
# ============================================================================
|
| 328 |
+
|
| 329 |
+
class NLPModule(nn.Module):
|
| 330 |
+
"""Módulo NLP integrando múltiplas tarefas no backbone CNN-BiGRU.
|
| 331 |
+
|
| 332 |
+
Args:
|
| 333 |
+
config: configuração NLP
|
| 334 |
+
backbone: modelo CooperativeCNNBiGRU (opcional — se None, cria um novo)
|
| 335 |
+
"""
|
| 336 |
+
|
| 337 |
+
def __init__(
|
| 338 |
+
self,
|
| 339 |
+
config: NLPConfig,
|
| 340 |
+
backbone: Optional[nn.Module] = None,
|
| 341 |
+
):
|
| 342 |
+
super().__init__()
|
| 343 |
+
self.config = config
|
| 344 |
+
|
| 345 |
+
# Backbone (cria se não fornecido)
|
| 346 |
+
if backbone is None:
|
| 347 |
+
from .cooperative_bigru import CooperativeCNNBiGRU
|
| 348 |
+
backbone = CooperativeCNNBiGRU(
|
| 349 |
+
vocab_size=32000, # placeholder, sobrescrito via set_vocab_size
|
| 350 |
+
embedding_dim=config.embed_dim,
|
| 351 |
+
cnn_filters=config.cnn_filters,
|
| 352 |
+
gru_hidden=config.gru_hidden,
|
| 353 |
+
n_heads=config.n_heads,
|
| 354 |
+
dropout=config.dropout,
|
| 355 |
+
pad_idx=config.pad_idx,
|
| 356 |
+
)
|
| 357 |
+
self.backbone = backbone
|
| 358 |
+
|
| 359 |
+
# Cabeças (uma por tarefa)
|
| 360 |
+
self.seq_head = SequenceClassificationHead(
|
| 361 |
+
embed_dim=config.feat_fused,
|
| 362 |
+
num_classes=config.num_classes_seq,
|
| 363 |
+
dropout=config.dropout,
|
| 364 |
+
pooling=config.pooling,
|
| 365 |
+
n_heads=config.n_heads,
|
| 366 |
+
)
|
| 367 |
+
self.tok_head = TokenClassificationHead(
|
| 368 |
+
embed_dim=config.feat_per_stream, # por stream (usa seq_A ou seq_B)
|
| 369 |
+
num_labels=config.num_labels_tok,
|
| 370 |
+
dropout=config.dropout,
|
| 371 |
+
)
|
| 372 |
+
self.span_head = SpanDetectionHead(
|
| 373 |
+
embed_dim=config.feat_per_stream,
|
| 374 |
+
dropout=config.dropout,
|
| 375 |
+
)
|
| 376 |
+
self.embed_head = EmbeddingHead(
|
| 377 |
+
embed_dim=config.feat_per_stream,
|
| 378 |
+
use_layer_norm=True,
|
| 379 |
+
)
|
| 380 |
+
|
| 381 |
+
# ----------------------------------------------------------------------
|
| 382 |
+
# Tarefas
|
| 383 |
+
# ----------------------------------------------------------------------
|
| 384 |
+
|
| 385 |
+
def sequence_classification(
|
| 386 |
+
self,
|
| 387 |
+
input_ids_a: torch.Tensor,
|
| 388 |
+
input_ids_b: torch.Tensor,
|
| 389 |
+
) -> Dict[str, torch.Tensor]:
|
| 390 |
+
"""Classificação de sequência (usando representação fused).
|
| 391 |
+
|
| 392 |
+
Args:
|
| 393 |
+
input_ids_a, input_ids_b: [B, T]
|
| 394 |
+
|
| 395 |
+
Returns:
|
| 396 |
+
dict com logits [B, C] e pooled [B, D]
|
| 397 |
+
"""
|
| 398 |
+
out = self.backbone(input_ids_a, input_ids_b, return_sequences=False)
|
| 399 |
+
fused = out["fused"] # [B, 256]
|
| 400 |
+
logits, pooled = self.seq_head(fused, mask=None)
|
| 401 |
+
return {"logits": logits, "pooled": pooled, "fused": fused}
|
| 402 |
+
|
| 403 |
+
def token_classification(
|
| 404 |
+
self,
|
| 405 |
+
input_ids_a: torch.Tensor,
|
| 406 |
+
input_ids_b: torch.Tensor,
|
| 407 |
+
stream: str = "A",
|
| 408 |
+
) -> torch.Tensor:
|
| 409 |
+
"""Classificação de tokens (NER/POS).
|
| 410 |
+
|
| 411 |
+
Args:
|
| 412 |
+
input_ids_a, input_ids_b: [B, T]
|
| 413 |
+
stream: "A" ou "B" — qual stream usar para token classification
|
| 414 |
+
|
| 415 |
+
Returns:
|
| 416 |
+
logits: [B, T, L]
|
| 417 |
+
"""
|
| 418 |
+
out = self.backbone(input_ids_a, input_ids_b, return_sequences=True)
|
| 419 |
+
seq = out.get("seq_A") if stream == "A" else out.get("seq_B")
|
| 420 |
+
if seq is None:
|
| 421 |
+
# Fallback: usar fused (broadcast para T)
|
| 422 |
+
fused = out["fused"] # [B, 256]
|
| 423 |
+
T = input_ids_a.size(1)
|
| 424 |
+
seq = fused.unsqueeze(1).expand(-1, T, -1) # [B, T, 256]
|
| 425 |
+
# Projeta para feat_per_stream se necessário
|
| 426 |
+
if fused.size(-1) != self.config.feat_per_stream:
|
| 427 |
+
if not hasattr(self, "_proj_tok"):
|
| 428 |
+
self._proj_tok = nn.Linear(
|
| 429 |
+
fused.size(-1), self.config.feat_per_stream, bias=False
|
| 430 |
+
).to(fused.device)
|
| 431 |
+
seq = self._proj_tok(seq)
|
| 432 |
+
return self.tok_head(seq)
|
| 433 |
+
|
| 434 |
+
def span_detection(
|
| 435 |
+
self,
|
| 436 |
+
input_ids_a: torch.Tensor,
|
| 437 |
+
input_ids_b: torch.Tensor,
|
| 438 |
+
stream: str = "A",
|
| 439 |
+
) -> Tuple[torch.Tensor, torch.Tensor]:
|
| 440 |
+
"""Detecção de span (QA extractive).
|
| 441 |
+
|
| 442 |
+
Returns:
|
| 443 |
+
(start_logits, end_logits) — ambos [B, T]
|
| 444 |
+
"""
|
| 445 |
+
out = self.backbone(input_ids_a, input_ids_b, return_sequences=True)
|
| 446 |
+
seq = out.get("seq_A") if stream == "A" else out.get("seq_B")
|
| 447 |
+
if seq is None:
|
| 448 |
+
fused = out["fused"]
|
| 449 |
+
T = input_ids_a.size(1)
|
| 450 |
+
seq = fused.unsqueeze(1).expand(-1, T, -1)
|
| 451 |
+
if fused.size(-1) != self.config.feat_per_stream:
|
| 452 |
+
if not hasattr(self, "_proj_span"):
|
| 453 |
+
self._proj_span = nn.Linear(
|
| 454 |
+
fused.size(-1), self.config.feat_per_stream, bias=False
|
| 455 |
+
).to(fused.device)
|
| 456 |
+
seq = self._proj_span(seq)
|
| 457 |
+
return self.span_head(seq)
|
| 458 |
+
|
| 459 |
+
def embed(
|
| 460 |
+
self,
|
| 461 |
+
input_ids_a: torch.Tensor,
|
| 462 |
+
input_ids_b: torch.Tensor,
|
| 463 |
+
stream: str = "A",
|
| 464 |
+
normalize: bool = True,
|
| 465 |
+
) -> torch.Tensor:
|
| 466 |
+
"""Gera embedding normalizado para retrieval semântico.
|
| 467 |
+
|
| 468 |
+
Returns:
|
| 469 |
+
embedding: [B, D]
|
| 470 |
+
"""
|
| 471 |
+
out = self.backbone(input_ids_a, input_ids_b, return_sequences=True)
|
| 472 |
+
seq = out.get("seq_A") if stream == "A" else out.get("seq_B")
|
| 473 |
+
mask = out.get("mask_A") if stream == "A" else out.get("mask_B")
|
| 474 |
+
if seq is None:
|
| 475 |
+
# Fallback: fused
|
| 476 |
+
return self.embed_head(
|
| 477 |
+
out["fused"].unsqueeze(1), mask=None, normalize=normalize
|
| 478 |
+
)
|
| 479 |
+
if seq.size(-1) != self.config.feat_per_stream:
|
| 480 |
+
if not hasattr(self, "_proj_emb"):
|
| 481 |
+
self._proj_emb = nn.Linear(
|
| 482 |
+
seq.size(-1), self.config.feat_per_stream, bias=False
|
| 483 |
+
).to(seq.device)
|
| 484 |
+
seq = self._proj_emb(seq)
|
| 485 |
+
return self.embed_head(seq, mask=mask, normalize=normalize)
|
| 486 |
+
|
| 487 |
+
# ----------------------------------------------------------------------
|
| 488 |
+
# Loss functions para cada tarefa
|
| 489 |
+
# ----------------------------------------------------------------------
|
| 490 |
+
|
| 491 |
+
def sequence_classification_loss(
|
| 492 |
+
self,
|
| 493 |
+
input_ids_a: torch.Tensor,
|
| 494 |
+
input_ids_b: torch.Tensor,
|
| 495 |
+
labels: torch.Tensor,
|
| 496 |
+
) -> torch.Tensor:
|
| 497 |
+
"""Computa CE loss para sequence classification."""
|
| 498 |
+
out = self.sequence_classification(input_ids_a, input_ids_b)
|
| 499 |
+
return F.cross_entropy(
|
| 500 |
+
out["logits"], labels,
|
| 501 |
+
label_smoothing=self.config.label_smoothing,
|
| 502 |
+
)
|
| 503 |
+
|
| 504 |
+
def token_classification_loss(
|
| 505 |
+
self,
|
| 506 |
+
input_ids_a: torch.Tensor,
|
| 507 |
+
input_ids_b: torch.Tensor,
|
| 508 |
+
labels: torch.Tensor,
|
| 509 |
+
stream: str = "A",
|
| 510 |
+
) -> torch.Tensor:
|
| 511 |
+
"""Computa CE loss para token classification.
|
| 512 |
+
|
| 513 |
+
Args:
|
| 514 |
+
labels: [B, T] rótulos por token (use -100 ou pad_idx para ignorar)
|
| 515 |
+
"""
|
| 516 |
+
logits = self.token_classification(input_ids_a, input_ids_b, stream=stream)
|
| 517 |
+
# Ignora posições com label = -100 ou mask = 0
|
| 518 |
+
mask = (input_ids_a != self.config.pad_idx).float()
|
| 519 |
+
ignore_mask = (labels == -100) | (labels == self.config.pad_idx)
|
| 520 |
+
labels_clamped = labels.clone()
|
| 521 |
+
labels_clamped[ignore_mask] = -100
|
| 522 |
+
return F.cross_entropy(
|
| 523 |
+
logits.view(-1, self.config.num_labels_tok),
|
| 524 |
+
labels_clamped.view(-1),
|
| 525 |
+
ignore_index=-100,
|
| 526 |
+
)
|
| 527 |
+
|
| 528 |
+
def span_detection_loss(
|
| 529 |
+
self,
|
| 530 |
+
input_ids_a: torch.Tensor,
|
| 531 |
+
input_ids_b: torch.Tensor,
|
| 532 |
+
start_positions: torch.Tensor,
|
| 533 |
+
end_positions: torch.Tensor,
|
| 534 |
+
stream: str = "A",
|
| 535 |
+
) -> torch.Tensor:
|
| 536 |
+
"""Computa loss para span detection.
|
| 537 |
+
|
| 538 |
+
Args:
|
| 539 |
+
start_positions, end_positions: [B] índices do span (long)
|
| 540 |
+
"""
|
| 541 |
+
start_logits, end_logits = self.span_detection(
|
| 542 |
+
input_ids_a, input_ids_b, stream=stream
|
| 543 |
+
)
|
| 544 |
+
# Ignora posições fora do span (mask)
|
| 545 |
+
loss_start = F.cross_entropy(start_logits, start_positions, ignore_index=-100)
|
| 546 |
+
loss_end = F.cross_entropy(end_logits, end_positions, ignore_index=-100)
|
| 547 |
+
return (loss_start + loss_end) / 2.0
|
| 548 |
+
|
| 549 |
+
def contrastive_loss(
|
| 550 |
+
self,
|
| 551 |
+
query_ids_a: torch.Tensor,
|
| 552 |
+
query_ids_b: torch.Tensor,
|
| 553 |
+
pos_ids_a: torch.Tensor,
|
| 554 |
+
pos_ids_b: torch.Tensor,
|
| 555 |
+
neg_ids_a: torch.Tensor,
|
| 556 |
+
neg_ids_b: torch.Tensor,
|
| 557 |
+
temperature: float = 0.07,
|
| 558 |
+
) -> torch.Tensor:
|
| 559 |
+
"""InfoNCE loss para retrieval semântico.
|
| 560 |
+
|
| 561 |
+
Args:
|
| 562 |
+
query, pos, neg: tensores [B, T] para query, positivo e negativo
|
| 563 |
+
"""
|
| 564 |
+
q = self.embed(query_ids_a, query_ids_b, normalize=True)
|
| 565 |
+
p = self.embed(pos_ids_a, pos_ids_b, normalize=True)
|
| 566 |
+
n = self.embed(neg_ids_a, neg_ids_b, normalize=True)
|
| 567 |
+
|
| 568 |
+
# Similaridade coseno (já normalizado -> produto escalar)
|
| 569 |
+
sim_pos = (q * p).sum(dim=-1, keepdim=True) / temperature # [B, 1]
|
| 570 |
+
sim_neg = (q * n).sum(dim=-1, keepdim=True) / temperature # [B, 1]
|
| 571 |
+
|
| 572 |
+
# InfoNCE: maximizar sim_pos - sim_neg
|
| 573 |
+
logits = torch.cat([sim_pos, sim_neg], dim=-1) # [B, 2]
|
| 574 |
+
labels = torch.zeros(q.size(0), dtype=torch.long, device=q.device) # 0 = pos
|
| 575 |
+
return F.cross_entropy(logits, labels)
|
| 576 |
+
|
| 577 |
+
|
| 578 |
+
# ============================================================================
|
| 579 |
+
# Self-test
|
| 580 |
+
# ============================================================================
|
| 581 |
+
|
| 582 |
+
def _self_test():
|
| 583 |
+
"""Teste rápido do módulo NLP."""
|
| 584 |
+
torch.manual_seed(42)
|
| 585 |
+
from .cooperative_bigru import CooperativeCNNBiGRU
|
| 586 |
+
config = NLPConfig(
|
| 587 |
+
embed_dim=32,
|
| 588 |
+
cnn_filters=32,
|
| 589 |
+
gru_hidden=32,
|
| 590 |
+
feat_per_stream=64,
|
| 591 |
+
feat_fused=128,
|
| 592 |
+
n_heads=4,
|
| 593 |
+
dropout=0.1,
|
| 594 |
+
pad_idx=1,
|
| 595 |
+
num_classes_seq=3,
|
| 596 |
+
num_labels_tok=5,
|
| 597 |
+
pooling="self_attn",
|
| 598 |
+
)
|
| 599 |
+
backbone = CooperativeCNNBiGRU(
|
| 600 |
+
vocab_size=100,
|
| 601 |
+
embedding_dim=config.embed_dim,
|
| 602 |
+
cnn_filters=config.cnn_filters,
|
| 603 |
+
gru_hidden=config.gru_hidden,
|
| 604 |
+
n_heads=config.n_heads,
|
| 605 |
+
dropout=config.dropout,
|
| 606 |
+
pad_idx=config.pad_idx,
|
| 607 |
+
)
|
| 608 |
+
nlp = NLPModule(config, backbone=backbone)
|
| 609 |
+
|
| 610 |
+
B, T = 2, 8
|
| 611 |
+
ids_a = torch.randint(2, 100, (B, T))
|
| 612 |
+
ids_b = torch.randint(2, 100, (B, T))
|
| 613 |
+
|
| 614 |
+
# Sequence classification
|
| 615 |
+
out = nlp.sequence_classification(ids_a, ids_b)
|
| 616 |
+
print(f"Seq cls: logits {out['logits'].shape}, pooled {out['pooled'].shape}")
|
| 617 |
+
labels = torch.randint(0, 3, (B,))
|
| 618 |
+
loss_seq = nlp.sequence_classification_loss(ids_a, ids_b, labels)
|
| 619 |
+
print(f" loss={loss_seq.item():.4f}")
|
| 620 |
+
|
| 621 |
+
# Token classification
|
| 622 |
+
tok_logits = nlp.token_classification(ids_a, ids_b)
|
| 623 |
+
print(f"Tok cls: {tok_logits.shape}")
|
| 624 |
+
tok_labels = torch.randint(0, 5, (B, T))
|
| 625 |
+
loss_tok = nlp.token_classification_loss(ids_a, ids_b, tok_labels)
|
| 626 |
+
print(f" loss={loss_tok.item():.4f}")
|
| 627 |
+
|
| 628 |
+
# Span detection
|
| 629 |
+
start_logits, end_logits = nlp.span_detection(ids_a, ids_b)
|
| 630 |
+
print(f"Span: start {start_logits.shape}, end {end_logits.shape}")
|
| 631 |
+
sp = torch.randint(0, T, (B,))
|
| 632 |
+
ep = torch.randint(0, T, (B,))
|
| 633 |
+
loss_span = nlp.span_detection_loss(ids_a, ids_b, sp, ep)
|
| 634 |
+
print(f" loss={loss_span.item():.4f}")
|
| 635 |
+
|
| 636 |
+
# Embedding
|
| 637 |
+
emb = nlp.embed(ids_a, ids_b)
|
| 638 |
+
print(f"Embed: {emb.shape}, norm: {emb.norm(dim=-1).mean().item():.4f}")
|
| 639 |
+
|
| 640 |
+
print("NLP self-test OK")
|
| 641 |
+
|
| 642 |
+
|
| 643 |
+
if __name__ == "__main__":
|
| 644 |
+
_self_test()
|
| 645 |
+
|
| 646 |
+
|
| 647 |
+
__all__ = [
|
| 648 |
+
"NLPConfig",
|
| 649 |
+
"SequenceClassificationHead",
|
| 650 |
+
"TokenClassificationHead",
|
| 651 |
+
"SpanDetectionHead",
|
| 652 |
+
"EmbeddingHead",
|
| 653 |
+
"NLPModule",
|
| 654 |
+
]
|
scripts/push_to_hf.py
CHANGED
|
@@ -1,4 +1,4 @@
|
|
| 1 |
-
"""push_to_hf.py — Envia scripts em lote ao HF repositório 'CNN-BiGRU' (
|
| 2 |
|
| 3 |
Executa:
|
| 4 |
1. Cria repositório 'PowerMachine/CNN-BiGRU' (se não existir)
|
|
@@ -6,12 +6,14 @@ Executa:
|
|
| 6 |
usando `upload_folder` (upload em lote eficiente)
|
| 7 |
3. SOBRESCREVE módulos desatualizados (default do upload_folder)
|
| 8 |
4. Remove HF_TOKEN do ambiente após o uso (requisito do usuário)
|
| 9 |
-
|
| 10 |
-
|
| 11 |
-
|
| 12 |
-
|
| 13 |
-
-
|
| 14 |
-
-
|
|
|
|
|
|
|
| 15 |
- Retry automático em caso de falha transitória
|
| 16 |
|
| 17 |
Usage:
|
|
@@ -22,6 +24,7 @@ from __future__ import annotations
|
|
| 22 |
import os
|
| 23 |
import sys
|
| 24 |
import logging
|
|
|
|
| 25 |
import time
|
| 26 |
from pathlib import Path
|
| 27 |
from typing import List, Set
|
|
@@ -48,6 +51,10 @@ IGNORE_PATTERNS: Set[str] = {
|
|
| 48 |
".DS_Store", "Thumbs.db",
|
| 49 |
"*.log", "*.tmp", "*.swp", "*.bak",
|
| 50 |
".env", ".venv", "venv", "env",
|
|
|
|
|
|
|
|
|
|
|
|
|
| 51 |
}
|
| 52 |
|
| 53 |
# Arquivos específicos que NUNCA devem ser upados (podem conter tokens/sensíveis)
|
|
@@ -71,6 +78,9 @@ def should_ignore(path: Path) -> bool:
|
|
| 71 |
# Verificar extensões de cache
|
| 72 |
if part_lower.endswith((".pyc", ".pyo", ".pyd")):
|
| 73 |
return True
|
|
|
|
|
|
|
|
|
|
| 74 |
|
| 75 |
# Verificar nome do arquivo
|
| 76 |
if name in SENSITIVE_FILENAMES:
|
|
@@ -101,6 +111,41 @@ def collect_files(project_dir: Path) -> List[Path]:
|
|
| 101 |
return files
|
| 102 |
|
| 103 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 104 |
def upload_with_retry(api, folder: Path, repo_id: str, token: str, max_retries: int = 3):
|
| 105 |
"""Faz upload_folder com retry em caso de falha transitória."""
|
| 106 |
last_error = None
|
|
@@ -112,7 +157,7 @@ def upload_with_retry(api, folder: Path, repo_id: str, token: str, max_retries:
|
|
| 112 |
repo_id=repo_id,
|
| 113 |
repo_type="model",
|
| 114 |
token=token,
|
| 115 |
-
commit_message=f"
|
| 116 |
# Sobrescreve arquivos desatualizados (default)
|
| 117 |
)
|
| 118 |
logger.info(f"Upload bem-sucedido: {commit_info}")
|
|
@@ -167,12 +212,12 @@ def main():
|
|
| 167 |
_cleanup_token()
|
| 168 |
return 1
|
| 169 |
|
| 170 |
-
# Log dos arquivos que serão upados
|
| 171 |
-
for f in files_to_upload[:
|
| 172 |
rel = f.relative_to(REPO_ROOT)
|
| 173 |
logger.info(f" - {rel}")
|
| 174 |
-
if len(files_to_upload) >
|
| 175 |
-
logger.info(f" ... e mais {len(files_to_upload) -
|
| 176 |
|
| 177 |
# 3. Upload em lote via upload_folder (muito mais eficiente que upload_file em loop)
|
| 178 |
success, error = upload_with_retry(
|
|
@@ -185,14 +230,21 @@ def main():
|
|
| 185 |
logger.info(f"Repo: https://huggingface.co/{REPO_ID}")
|
| 186 |
logger.info(f"Total de arquivos enviados: {len(files_to_upload)}")
|
| 187 |
logger.info("=" * 60)
|
|
|
|
|
|
|
|
|
|
|
|
|
| 188 |
else:
|
| 189 |
logger.error("=" * 60)
|
| 190 |
logger.error(f"Falha no upload após retries: {error}")
|
| 191 |
logger.error("=" * 60)
|
| 192 |
|
| 193 |
-
#
|
| 194 |
_cleanup_token()
|
| 195 |
|
|
|
|
|
|
|
|
|
|
| 196 |
return 0 if success else 1
|
| 197 |
|
| 198 |
|
|
@@ -209,5 +261,34 @@ def _cleanup_token():
|
|
| 209 |
logger.info("Nenhum token HF encontrado no ambiente para remover.")
|
| 210 |
|
| 211 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 212 |
if __name__ == "__main__":
|
| 213 |
sys.exit(main())
|
|
|
|
| 1 |
+
"""push_to_hf.py — Envia scripts em lote ao HF repositório 'CNN-BiGRU' (v3.0).
|
| 2 |
|
| 3 |
Executa:
|
| 4 |
1. Cria repositório 'PowerMachine/CNN-BiGRU' (se não existir)
|
|
|
|
| 6 |
usando `upload_folder` (upload em lote eficiente)
|
| 7 |
3. SOBRESCREVE módulos desatualizados (default do upload_folder)
|
| 8 |
4. Remove HF_TOKEN do ambiente após o uso (requisito do usuário)
|
| 9 |
+
5. Deleta estado salvo do modelo (não envia para lugar algum)
|
| 10 |
+
6. Filtra arquivos sensíveis (.env, tokens, checkpoints, etc.)
|
| 11 |
+
|
| 12 |
+
V3.0 melhorias:
|
| 13 |
+
- Filtra arquivos de estado do modelo (.pt, .pth, .bin, .safetensors)
|
| 14 |
+
- Filtra relatórios de monitor local (monitor_reports/)
|
| 15 |
+
- Mantém apenas scripts e documentação
|
| 16 |
+
- Log mais detalhado
|
| 17 |
- Retry automático em caso de falha transitória
|
| 18 |
|
| 19 |
Usage:
|
|
|
|
| 24 |
import os
|
| 25 |
import sys
|
| 26 |
import logging
|
| 27 |
+
import shutil
|
| 28 |
import time
|
| 29 |
from pathlib import Path
|
| 30 |
from typing import List, Set
|
|
|
|
| 51 |
".DS_Store", "Thumbs.db",
|
| 52 |
"*.log", "*.tmp", "*.swp", "*.bak",
|
| 53 |
".env", ".venv", "venv", "env",
|
| 54 |
+
# Estado salvo do modelo — NÃO ENVIAR (requisito do usuário)
|
| 55 |
+
"*.pt", "*.pth", "*.bin", "*.safetensors", "*.ckpt",
|
| 56 |
+
# Relatórios de monitor local (contém paths e dados locais)
|
| 57 |
+
"monitor_reports", "download", "tool-results", "upload",
|
| 58 |
}
|
| 59 |
|
| 60 |
# Arquivos específicos que NUNCA devem ser upados (podem conter tokens/sensíveis)
|
|
|
|
| 78 |
# Verificar extensões de cache
|
| 79 |
if part_lower.endswith((".pyc", ".pyo", ".pyd")):
|
| 80 |
return True
|
| 81 |
+
# Verificar extensões de estado de modelo
|
| 82 |
+
if part_lower.endswith((".pt", ".pth", ".bin", ".safetensors", ".ckpt")):
|
| 83 |
+
return True
|
| 84 |
|
| 85 |
# Verificar nome do arquivo
|
| 86 |
if name in SENSITIVE_FILENAMES:
|
|
|
|
| 111 |
return files
|
| 112 |
|
| 113 |
|
| 114 |
+
def delete_saved_model_state(project_dir: Path) -> None:
|
| 115 |
+
"""Deleta todo o estado salvo do modelo (checkpoints, .pt, .pth, etc.).
|
| 116 |
+
|
| 117 |
+
Requisito do usuário: "apagar o estado salvo do modelo (não enviar para
|
| 118 |
+
lugar algum até ordem em contrário)".
|
| 119 |
+
"""
|
| 120 |
+
deleted = []
|
| 121 |
+
for path in project_dir.rglob("*"):
|
| 122 |
+
if not path.is_file():
|
| 123 |
+
continue
|
| 124 |
+
name = path.name.lower()
|
| 125 |
+
if name.endswith((".pt", ".pth", ".bin", ".safetensors", ".ckpt")):
|
| 126 |
+
try:
|
| 127 |
+
path.unlink()
|
| 128 |
+
deleted.append(path)
|
| 129 |
+
logger.info(f" Deletado: {path.relative_to(project_dir.parent)}")
|
| 130 |
+
except Exception as e:
|
| 131 |
+
logger.warning(f" Falha ao deletar {path}: {e}")
|
| 132 |
+
|
| 133 |
+
# Deletar diretórios de cache e checkpoints
|
| 134 |
+
for d in project_dir.rglob("*"):
|
| 135 |
+
if d.is_dir() and d.name.lower() in {"checkpoints", "checkpoint", "model_state", "weights"}:
|
| 136 |
+
try:
|
| 137 |
+
shutil.rmtree(d)
|
| 138 |
+
deleted.append(d)
|
| 139 |
+
logger.info(f" Deletado dir: {d.relative_to(project_dir.parent)}")
|
| 140 |
+
except Exception as e:
|
| 141 |
+
logger.warning(f" Falha ao deletar dir {d}: {e}")
|
| 142 |
+
|
| 143 |
+
if deleted:
|
| 144 |
+
logger.info(f"Estado salvo do modelo deletado: {len(deleted)} itens")
|
| 145 |
+
else:
|
| 146 |
+
logger.info("Nenhum estado salvo do modelo encontrado para deletar")
|
| 147 |
+
|
| 148 |
+
|
| 149 |
def upload_with_retry(api, folder: Path, repo_id: str, token: str, max_retries: int = 3):
|
| 150 |
"""Faz upload_folder com retry em caso de falha transitória."""
|
| 151 |
last_error = None
|
|
|
|
| 157 |
repo_id=repo_id,
|
| 158 |
repo_type="model",
|
| 159 |
token=token,
|
| 160 |
+
commit_message=f"v3.0: 9 novos módulos (CyclicReasoning, Medusa, NLG, NLP, MMA, VQVAE2, W8A8, LongContext 1M, Monitor) + bug fixes (attempt {attempt})",
|
| 161 |
# Sobrescreve arquivos desatualizados (default)
|
| 162 |
)
|
| 163 |
logger.info(f"Upload bem-sucedido: {commit_info}")
|
|
|
|
| 212 |
_cleanup_token()
|
| 213 |
return 1
|
| 214 |
|
| 215 |
+
# Log dos arquivos que serão upados (até 20)
|
| 216 |
+
for f in files_to_upload[:20]:
|
| 217 |
rel = f.relative_to(REPO_ROOT)
|
| 218 |
logger.info(f" - {rel}")
|
| 219 |
+
if len(files_to_upload) > 20:
|
| 220 |
+
logger.info(f" ... e mais {len(files_to_upload) - 20} arquivos")
|
| 221 |
|
| 222 |
# 3. Upload em lote via upload_folder (muito mais eficiente que upload_file em loop)
|
| 223 |
success, error = upload_with_retry(
|
|
|
|
| 230 |
logger.info(f"Repo: https://huggingface.co/{REPO_ID}")
|
| 231 |
logger.info(f"Total de arquivos enviados: {len(files_to_upload)}")
|
| 232 |
logger.info("=" * 60)
|
| 233 |
+
|
| 234 |
+
# 4. Deletar estado salvo do modelo (após upload bem-sucedido)
|
| 235 |
+
logger.info("Deletando estado salvo do modelo...")
|
| 236 |
+
delete_saved_model_state(PROJECT_DIR)
|
| 237 |
else:
|
| 238 |
logger.error("=" * 60)
|
| 239 |
logger.error(f"Falha no upload após retries: {error}")
|
| 240 |
logger.error("=" * 60)
|
| 241 |
|
| 242 |
+
# 5. Remove HF_TOKEN do ambiente (requisito do usuário)
|
| 243 |
_cleanup_token()
|
| 244 |
|
| 245 |
+
# 6. Verifica e remove tokens hardcoded em scripts (requisito do usuário)
|
| 246 |
+
_scan_scripts_for_tokens(PROJECT_DIR)
|
| 247 |
+
|
| 248 |
return 0 if success else 1
|
| 249 |
|
| 250 |
|
|
|
|
| 261 |
logger.info("Nenhum token HF encontrado no ambiente para remover.")
|
| 262 |
|
| 263 |
|
| 264 |
+
def _scan_scripts_for_tokens(project_dir: Path) -> None:
|
| 265 |
+
"""Verifica se há tokens HF hardcoded em scripts e os remove."""
|
| 266 |
+
# Padrões de token HF (hf_ seguido de 32+ caracteres alfanuméricos)
|
| 267 |
+
import re
|
| 268 |
+
token_pattern = re.compile(r'hf_[A-Za-z0-9]{20,}')
|
| 269 |
+
|
| 270 |
+
scanned = 0
|
| 271 |
+
cleaned = 0
|
| 272 |
+
for path in project_dir.rglob("*.py"):
|
| 273 |
+
if "push_to_hf.py" in str(path):
|
| 274 |
+
continue # skip self
|
| 275 |
+
try:
|
| 276 |
+
content = path.read_text(encoding="utf-8")
|
| 277 |
+
scanned += 1
|
| 278 |
+
if token_pattern.search(content):
|
| 279 |
+
# Substituir por placeholder
|
| 280 |
+
new_content = token_pattern.sub("HF_TOKEN_FROM_ENV", content)
|
| 281 |
+
path.write_text(new_content, encoding="utf-8")
|
| 282 |
+
cleaned += 1
|
| 283 |
+
logger.warning(f" Token HF removido de: {path.relative_to(project_dir.parent)}")
|
| 284 |
+
except Exception as e:
|
| 285 |
+
logger.debug(f" Falha ao escanear {path}: {e}")
|
| 286 |
+
|
| 287 |
+
if cleaned > 0:
|
| 288 |
+
logger.info(f"Scripts escaneados: {scanned}, scripts com token removido: {cleaned}")
|
| 289 |
+
else:
|
| 290 |
+
logger.info(f"Scripts escaneados: {scanned}, nenhum token hardcoded encontrado.")
|
| 291 |
+
|
| 292 |
+
|
| 293 |
if __name__ == "__main__":
|
| 294 |
sys.exit(main())
|
tests/test_500_samples.py
ADDED
|
@@ -0,0 +1,1035 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""test_500_samples.py — Teste de 500 amostras (5 batches de 100) para o CNN-BiGRU.
|
| 2 |
+
|
| 3 |
+
Executa o pipeline completo com:
|
| 4 |
+
1. Inicializa xeon_runtime + memory_optimizer + monitor
|
| 5 |
+
2. Treina BBPE tokenizer em corpus sintético
|
| 6 |
+
3. Cria dataset streaming multimodal (até 500 amostras em batches de 100)
|
| 7 |
+
4. Instancia modelo multimodal + generator + verifier + anti-hallucination
|
| 8 |
+
5. Testa os NOVOS módulos v3.0:
|
| 9 |
+
- CyclicReasoning
|
| 10 |
+
- MedusaMTP
|
| 11 |
+
- NLGModule
|
| 12 |
+
- NLPModule
|
| 13 |
+
- MultimodalMultiHeadAttention
|
| 14 |
+
- VQVAE2
|
| 15 |
+
- W8A8 Quantization (SmoothQuant)
|
| 16 |
+
- LongContextManager (1M tokens)
|
| 17 |
+
- Monitor
|
| 18 |
+
6. Executa treinamento cooperativo com:
|
| 19 |
+
- Synergy search (N tentativas)
|
| 20 |
+
- Hypothesis controller (ativa em punições)
|
| 21 |
+
- Auto-learner (ajuste dinâmico de LR + spectral norm)
|
| 22 |
+
- EWC (aprendizado contínuo)
|
| 23 |
+
7. Executa inferência com sampling
|
| 24 |
+
8. Avalia perplexidade
|
| 25 |
+
9. Aplica quantização W8A8 ao modelo
|
| 26 |
+
10. Reporta erros lógicos/falhas encontradas + exporta relatório
|
| 27 |
+
|
| 28 |
+
Usage:
|
| 29 |
+
python -m cnn_bigru.tests.test_500_samples
|
| 30 |
+
"""
|
| 31 |
+
from __future__ import annotations
|
| 32 |
+
|
| 33 |
+
import logging
|
| 34 |
+
import os
|
| 35 |
+
import sys
|
| 36 |
+
import time
|
| 37 |
+
import traceback
|
| 38 |
+
from pathlib import Path
|
| 39 |
+
from typing import Dict, List
|
| 40 |
+
|
| 41 |
+
# Setup paths (must come before torch import for xeon_runtime)
|
| 42 |
+
PROJECT_ROOT = Path(__file__).resolve().parent.parent.parent
|
| 43 |
+
sys.path.insert(0, str(PROJECT_ROOT))
|
| 44 |
+
|
| 45 |
+
# Ativa Xeon runtime ANTES de importar torch
|
| 46 |
+
from cnn_bigru.utils.xeon_runtime import optimize_xeon_environment
|
| 47 |
+
N_CORES = optimize_xeon_environment()
|
| 48 |
+
|
| 49 |
+
import numpy as np
|
| 50 |
+
import torch
|
| 51 |
+
from torch.utils.data import DataLoader
|
| 52 |
+
|
| 53 |
+
from cnn_bigru.tokenizer.bbpe_tokenizer import BBPETokenizer
|
| 54 |
+
from cnn_bigru.data.streaming_dataset import (
|
| 55 |
+
MultimodalStreamingDataset,
|
| 56 |
+
collate_multimodal,
|
| 57 |
+
DEFAULT_DATASETS,
|
| 58 |
+
)
|
| 59 |
+
from cnn_bigru.models.multimodal_model import MultimodalCNNBiGRU
|
| 60 |
+
from cnn_bigru.models.generator_verifier import (
|
| 61 |
+
GeneratorCNNBiGRU,
|
| 62 |
+
VerifierCNNBiGRU,
|
| 63 |
+
AntiHallucinationLayer,
|
| 64 |
+
)
|
| 65 |
+
from cnn_bigru.models.cooperative_bigru import CooperativeCNNBiGRU
|
| 66 |
+
from cnn_bigru.models.rope import RotaryPositionEmbedding
|
| 67 |
+
from cnn_bigru.models.transformer_block import (
|
| 68 |
+
TransformerBlockConfig,
|
| 69 |
+
CausalSelfAttention,
|
| 70 |
+
TransformerBlock,
|
| 71 |
+
TransformerDecoderStack,
|
| 72 |
+
)
|
| 73 |
+
from cnn_bigru.models.context_window import (
|
| 74 |
+
ContextWindowConfig,
|
| 75 |
+
ContextWindowManager,
|
| 76 |
+
KVCache,
|
| 77 |
+
LongContextConfig,
|
| 78 |
+
LongContextManager,
|
| 79 |
+
make_long_context_window,
|
| 80 |
+
)
|
| 81 |
+
# NOVOS módulos v3.0
|
| 82 |
+
from cnn_bigru.models.cyclic_reasoning import (
|
| 83 |
+
CyclicReasoningConfig,
|
| 84 |
+
CyclicReasoning,
|
| 85 |
+
make_hypothesis_fn,
|
| 86 |
+
)
|
| 87 |
+
from cnn_bigru.models.medusa_heads import (
|
| 88 |
+
MedusaConfig,
|
| 89 |
+
MedusaMTP,
|
| 90 |
+
MedusaHead,
|
| 91 |
+
medusa_tree_decode,
|
| 92 |
+
)
|
| 93 |
+
from cnn_bigru.models.nlg import NLGConfig, NLGModule
|
| 94 |
+
from cnn_bigru.models.nlp import (
|
| 95 |
+
NLPConfig,
|
| 96 |
+
NLPModule,
|
| 97 |
+
SequenceClassificationHead,
|
| 98 |
+
TokenClassificationHead,
|
| 99 |
+
SpanDetectionHead,
|
| 100 |
+
EmbeddingHead,
|
| 101 |
+
)
|
| 102 |
+
from cnn_bigru.models.multimodal_attention import (
|
| 103 |
+
MultimodalAttentionConfig,
|
| 104 |
+
MultimodalMultiHeadAttention,
|
| 105 |
+
CrossModalAttention,
|
| 106 |
+
ModalityGate,
|
| 107 |
+
)
|
| 108 |
+
from cnn_bigru.utils.ewc import EWCConfig, EWCState
|
| 109 |
+
from cnn_bigru.utils.quantization import (
|
| 110 |
+
W8A8Config,
|
| 111 |
+
SmoothQuantizer,
|
| 112 |
+
quantize_model_w8a8,
|
| 113 |
+
estimate_memory_savings,
|
| 114 |
+
)
|
| 115 |
+
from cnn_bigru.utils.vqvae2 import VQVAE2Config, VQVAE2
|
| 116 |
+
from cnn_bigru.utils.monitoring import Monitor, get_monitor
|
| 117 |
+
from cnn_bigru.losses.losses import LossConfig, MultiLoss
|
| 118 |
+
from cnn_bigru.training.trainer import TrainerConfig, CooperativeTrainer
|
| 119 |
+
from cnn_bigru.training.auto_learner import (
|
| 120 |
+
AutoLearnConfig,
|
| 121 |
+
orthogonal_init_model,
|
| 122 |
+
)
|
| 123 |
+
from cnn_bigru.training.hypothesis_controller import (
|
| 124 |
+
HypothesisConfig,
|
| 125 |
+
HypothesisController,
|
| 126 |
+
)
|
| 127 |
+
from cnn_bigru.inference.inference import (
|
| 128 |
+
generate_with_sampling,
|
| 129 |
+
evaluate_perplexity,
|
| 130 |
+
make_default_context_window,
|
| 131 |
+
)
|
| 132 |
+
from cnn_bigru.utils.memory_optimizer import MemoryOptimizer
|
| 133 |
+
|
| 134 |
+
logging.basicConfig(
|
| 135 |
+
level=logging.INFO,
|
| 136 |
+
format="%(asctime)s | %(levelname)s | %(name)s | %(message)s",
|
| 137 |
+
datefmt="%H:%M:%S",
|
| 138 |
+
)
|
| 139 |
+
logger = logging.getLogger("test_500")
|
| 140 |
+
|
| 141 |
+
|
| 142 |
+
# ============================================================================
|
| 143 |
+
# Configurações do teste
|
| 144 |
+
# ============================================================================
|
| 145 |
+
|
| 146 |
+
TOTAL_SAMPLES = 500
|
| 147 |
+
BATCH_SIZE_SAMPLES = 100 # 5 batches de 100
|
| 148 |
+
N_BATCHES = TOTAL_SAMPLES // BATCH_SIZE_SAMPLES # 5
|
| 149 |
+
# Para teste rápido, pode-se reduzir N_BATCHES via env var
|
| 150 |
+
if os.environ.get("CNN_BIGRU_TEST_FAST") == "1":
|
| 151 |
+
N_BATCHES = 2 # 200 amostras em modo rápido
|
| 152 |
+
TOTAL_SAMPLES = N_BATCHES * BATCH_SIZE_SAMPLES
|
| 153 |
+
|
| 154 |
+
CORPUS_TEXTS = [
|
| 155 |
+
"o modelo cooperativo cnn-bigru combina fluxos paralelos",
|
| 156 |
+
"a atencao cruzada troca informacoes entre redes a e b",
|
| 157 |
+
"gru bidirecional captura dependencias temporais em ambas direcoes",
|
| 158 |
+
"a porta de atenuacao controla o fluxo de informacao entre celulas",
|
| 159 |
+
"a camada anti-alucinacao usa logica fuzzy de lukasiewicz",
|
| 160 |
+
"o verificador classifica passos com sigmoid binaria",
|
| 161 |
+
"penalidades linear e exponencial reduzem o erro grave",
|
| 162 |
+
"ajuste dinamico de lr estabiliza o treinamento cooperativo",
|
| 163 |
+
"normalizacao espectral limita a norma dos pesos",
|
| 164 |
+
"inicializacao ortogonal estabiliza matrizes recorrentes",
|
| 165 |
+
"o gerador decodificador usa atencao bahdanau sobre o encoder",
|
| 166 |
+
"a fusao multimodal combina texto imagem e audio",
|
| 167 |
+
"hipoteses sao ativadas quando o verificador pune o passo",
|
| 168 |
+
"synergy search tenta n configuracoes e escolhe a melhor",
|
| 169 |
+
"a perplexidade mede a confusao do modelo na previsao",
|
| 170 |
+
"top-k e top-p filtram a distribuicao de probabilidade",
|
| 171 |
+
"temperature ajusta a entropia das previsoes do modelo",
|
| 172 |
+
"presence penalty pune tokens ja aparecidos na geracao",
|
| 173 |
+
"frequency penalty pune proporcionalmente a frequencia do token",
|
| 174 |
+
"streaming dataset carrega amostras sem materializacao completa",
|
| 175 |
+
"byte bpe tokeniza qualquer string utf-8 sem unk",
|
| 176 |
+
"o otimizador adamw combina momentum e weight decay",
|
| 177 |
+
"gradient clipping previne explosao de gradiente",
|
| 178 |
+
"amp reduz vram com precisao mista bfloat16",
|
| 179 |
+
# NOVOS v3.0
|
| 180 |
+
"raciocinio ciclico refina a representacao iterativamente",
|
| 181 |
+
"medusa heads predizem multiplos tokens em paralelo",
|
| 182 |
+
"nlg gera texto autoregressivo com transformer decoder",
|
| 183 |
+
"nlp classifica sequencias tokens e spans",
|
| 184 |
+
"w8a8 quantiza pesos e ativacoes em 8 bits",
|
| 185 |
+
"smoothquant migra variancia das ativacoes para os pesos",
|
| 186 |
+
"vq-vae-2 hierarquico comprime com codebooks top e bottom",
|
| 187 |
+
"multi-token prediction acelera a geracao por arvores",
|
| 188 |
+
"multi-head attention multimodal funde modalidades por atencao",
|
| 189 |
+
"context window de 1m tokens usa chunked attention",
|
| 190 |
+
"monitor rastreia metricas de treino e evolucao",
|
| 191 |
+
"ewc previne esquecimento catastrofico em aprendizado continuo",
|
| 192 |
+
"rope codifica posicoes por rotacao no espaco complexo",
|
| 193 |
+
"kv cache armazena chaves e valores para atencao eficiente",
|
| 194 |
+
"causal self attention mascara tokens futuros no decoder",
|
| 195 |
+
"transformer block combina self attention e feed forward",
|
| 196 |
+
] * 4 # ~160 amostras para treino BBPE
|
| 197 |
+
|
| 198 |
+
|
| 199 |
+
# ============================================================================
|
| 200 |
+
# Step 1: Inicialização
|
| 201 |
+
# ============================================================================
|
| 202 |
+
|
| 203 |
+
def step_01_init_runtime() -> Dict:
|
| 204 |
+
"""Inicializa runtime Xeon + memory optimizer + monitor."""
|
| 205 |
+
logger.info("=" * 70)
|
| 206 |
+
logger.info("PASSO 1: Inicialização do runtime Xeon + memory optimizer + monitor")
|
| 207 |
+
logger.info("=" * 70)
|
| 208 |
+
from cnn_bigru.utils.xeon_runtime import get_runtime_info
|
| 209 |
+
info = get_runtime_info()
|
| 210 |
+
for k, v in info.items():
|
| 211 |
+
logger.info(" %s: %s", k, v)
|
| 212 |
+
mem = MemoryOptimizer(enable_amp=False)
|
| 213 |
+
mem.configure()
|
| 214 |
+
mem_info = mem.get_memory_mb()
|
| 215 |
+
logger.info(" Memória: %s", mem_info)
|
| 216 |
+
|
| 217 |
+
# Inicializa monitor global
|
| 218 |
+
monitor = get_monitor(output_dir=PROJECT_ROOT / "download" / "monitor_reports")
|
| 219 |
+
monitor.register_component("ewc", True)
|
| 220 |
+
monitor.register_component("medusa", True)
|
| 221 |
+
monitor.register_component("cyclic_reasoning", True)
|
| 222 |
+
monitor.register_component("vqvae2", True)
|
| 223 |
+
monitor.register_component("quantization", True)
|
| 224 |
+
monitor.register_component("multimodal_attention", True)
|
| 225 |
+
monitor.register_component("context_window", True)
|
| 226 |
+
monitor.register_component("monitor", True)
|
| 227 |
+
return {"runtime": info, "memory": mem_info, "monitor": monitor}
|
| 228 |
+
|
| 229 |
+
|
| 230 |
+
# ============================================================================
|
| 231 |
+
# Step 2: Tokenizer
|
| 232 |
+
# ============================================================================
|
| 233 |
+
|
| 234 |
+
def step_02_train_tokenizer() -> tuple:
|
| 235 |
+
"""Treina BBPE tokenizer em corpus sintético."""
|
| 236 |
+
logger.info("=" * 70)
|
| 237 |
+
logger.info("PASSO 2: Treino do BBPE tokenizer")
|
| 238 |
+
logger.info("=" * 70)
|
| 239 |
+
t0 = time.time()
|
| 240 |
+
tok = BBPETokenizer.train_from_texts(
|
| 241 |
+
CORPUS_TEXTS,
|
| 242 |
+
vocab_size=2000,
|
| 243 |
+
min_frequency=1,
|
| 244 |
+
)
|
| 245 |
+
elapsed = time.time() - t0
|
| 246 |
+
logger.info(" Vocab size: %d", tok.vocab_size)
|
| 247 |
+
logger.info(" BOS/PAD/EOS IDs: %d/%d/%d", tok.bos_id, tok.pad_id, tok.eos_id)
|
| 248 |
+
logger.info(" Tempo: %.2fs", elapsed)
|
| 249 |
+
test_strs = CORPUS_TEXTS[:5]
|
| 250 |
+
roundtrip = tok.validate_roundtrip(test_strs)
|
| 251 |
+
logger.info(" Roundtrip accuracy: %.2f%%", roundtrip * 100)
|
| 252 |
+
if roundtrip < 0.8:
|
| 253 |
+
logger.warning(" Roundtrip baixo — possível problema no tokenizer")
|
| 254 |
+
return tok, roundtrip
|
| 255 |
+
|
| 256 |
+
|
| 257 |
+
# ============================================================================
|
| 258 |
+
# Step 3: Dataset (até 500 amostras em batches de 100)
|
| 259 |
+
# ============================================================================
|
| 260 |
+
|
| 261 |
+
def step_03_create_dataset(
|
| 262 |
+
tokenizer: BBPETokenizer,
|
| 263 |
+
n_samples: int = BATCH_SIZE_SAMPLES,
|
| 264 |
+
seed_offset: int = 0,
|
| 265 |
+
) -> DataLoader:
|
| 266 |
+
"""Cria dataset streaming multimodal com N amostras (batch de 100)."""
|
| 267 |
+
logger.info("=" * 70)
|
| 268 |
+
logger.info("PASSO 3: Criação do dataset streaming multimodal (%d amostras, batch %d/5)",
|
| 269 |
+
n_samples, seed_offset + 1)
|
| 270 |
+
logger.info("=" * 70)
|
| 271 |
+
|
| 272 |
+
dataset = MultimodalStreamingDataset(
|
| 273 |
+
n_samples=n_samples,
|
| 274 |
+
# Tenta repositório 'PowerMachine/CNN-BiGRU' primeiro, depois fallback
|
| 275 |
+
hf_datasets=DEFAULT_DATASETS,
|
| 276 |
+
use_synthetic_fallback=True,
|
| 277 |
+
seed=42 + seed_offset * 100,
|
| 278 |
+
image_size=(28, 28, 1),
|
| 279 |
+
audio_shape=(32, 40),
|
| 280 |
+
)
|
| 281 |
+
|
| 282 |
+
samples = list(dataset)
|
| 283 |
+
logger.info(" Amostras coletadas: %d", len(samples))
|
| 284 |
+
assert len(samples) == n_samples, f"Esperado {n_samples}, obtido {len(samples)}"
|
| 285 |
+
|
| 286 |
+
loader = DataLoader(
|
| 287 |
+
samples,
|
| 288 |
+
batch_size=8,
|
| 289 |
+
shuffle=False,
|
| 290 |
+
collate_fn=lambda b: collate_multimodal(b, tokenizer, max_len=32),
|
| 291 |
+
)
|
| 292 |
+
logger.info(" DataLoader criado: batch_size=8, max_len=32")
|
| 293 |
+
return loader
|
| 294 |
+
|
| 295 |
+
|
| 296 |
+
# ============================================================================
|
| 297 |
+
# Step 4: Instanciar modelos
|
| 298 |
+
# ============================================================================
|
| 299 |
+
|
| 300 |
+
def step_04_init_model(tokenizer: BBPETokenizer, device: str = "cpu") -> Dict:
|
| 301 |
+
"""Instancia todos os componentes do modelo."""
|
| 302 |
+
logger.info("=" * 70)
|
| 303 |
+
logger.info("PASSO 4: Instanciação dos modelos (incluindo novos módulos v3.0)")
|
| 304 |
+
logger.info("=" * 70)
|
| 305 |
+
V = tokenizer.vocab_size
|
| 306 |
+
|
| 307 |
+
# Modelo multimodal principal
|
| 308 |
+
model = MultimodalCNNBiGRU(
|
| 309 |
+
vocab_size=V,
|
| 310 |
+
num_classes=3,
|
| 311 |
+
embedding_dim=32,
|
| 312 |
+
cnn_filters=32,
|
| 313 |
+
gru_hidden=32,
|
| 314 |
+
n_heads=4,
|
| 315 |
+
dropout=0.1,
|
| 316 |
+
pad_idx=tokenizer.pad_id,
|
| 317 |
+
img_channels=1,
|
| 318 |
+
img_hidden=16,
|
| 319 |
+
img_out_dim=32,
|
| 320 |
+
audio_freq=40,
|
| 321 |
+
audio_hidden=16,
|
| 322 |
+
audio_out_dim=32,
|
| 323 |
+
fusion_dim=64,
|
| 324 |
+
use_spectral_norm=False,
|
| 325 |
+
)
|
| 326 |
+
n_init = orthogonal_init_model(model)
|
| 327 |
+
logger.info(" Modelo multimodal: %d params, %d camadas ortogonalizadas",
|
| 328 |
+
sum(p.numel() for p in model.parameters()), n_init)
|
| 329 |
+
|
| 330 |
+
# Generator
|
| 331 |
+
generator = GeneratorCNNBiGRU(
|
| 332 |
+
vocab_size=V, embedding_dim=32, cnn_filters=32, gru_hidden=32,
|
| 333 |
+
n_heads=4, dropout=0.1, pad_idx=tokenizer.pad_id, max_proof_len=16,
|
| 334 |
+
)
|
| 335 |
+
orthogonal_init_model(generator)
|
| 336 |
+
|
| 337 |
+
# Verifier
|
| 338 |
+
verifier = VerifierCNNBiGRU(
|
| 339 |
+
vocab_size=V, embedding_dim=32, cnn_filters=32, gru_hidden=32,
|
| 340 |
+
n_heads=4, dropout=0.1, pad_idx=tokenizer.pad_id,
|
| 341 |
+
)
|
| 342 |
+
orthogonal_init_model(verifier)
|
| 343 |
+
|
| 344 |
+
# Anti-hallucination
|
| 345 |
+
anti_hall = AntiHallucinationLayer(vocab_size=V, embed_dim=16)
|
| 346 |
+
|
| 347 |
+
logger.info(" Generator params: %d", sum(p.numel() for p in generator.parameters()))
|
| 348 |
+
logger.info(" Verifier params: %d", sum(p.numel() for p in verifier.parameters()))
|
| 349 |
+
logger.info(" Anti-hall params: %d", sum(p.numel() for p in anti_hall.parameters()))
|
| 350 |
+
|
| 351 |
+
return {
|
| 352 |
+
"model": model,
|
| 353 |
+
"generator": generator,
|
| 354 |
+
"verifier": verifier,
|
| 355 |
+
"anti_hall": anti_hall,
|
| 356 |
+
}
|
| 357 |
+
|
| 358 |
+
|
| 359 |
+
# ============================================================================
|
| 360 |
+
# Step 5b: Testar NOVOS módulos v3.0
|
| 361 |
+
# ============================================================================
|
| 362 |
+
|
| 363 |
+
def step_05b_test_v3_modules(
|
| 364 |
+
models: Dict,
|
| 365 |
+
tokenizer: BBPETokenizer,
|
| 366 |
+
device: str = "cpu",
|
| 367 |
+
) -> Dict:
|
| 368 |
+
"""Testa os novos módulos v3.0: CyclicReasoning, Medusa, NLG, NLP, MMA, VQVAE2, W8A8, LongCtx."""
|
| 369 |
+
logger.info("=" * 70)
|
| 370 |
+
logger.info("PASSO 5b: Teste dos NOVOS módulos v3.0")
|
| 371 |
+
logger.info("=" * 70)
|
| 372 |
+
|
| 373 |
+
results = {}
|
| 374 |
+
V = tokenizer.vocab_size
|
| 375 |
+
|
| 376 |
+
# ---------- CyclicReasoning ----------
|
| 377 |
+
try:
|
| 378 |
+
logger.info(" [CyclicReasoning] Testando...")
|
| 379 |
+
cfg = CyclicReasoningConfig(
|
| 380 |
+
embed_dim=64, max_cycles=4, convergence_eps=1e-3,
|
| 381 |
+
use_anti_hallucination_gate=True,
|
| 382 |
+
)
|
| 383 |
+
cr = CyclicReasoning(cfg).to(device)
|
| 384 |
+
h0 = torch.randn(2, 64, device=device)
|
| 385 |
+
result = cr(h0, return_history=True)
|
| 386 |
+
assert result["h_final"].shape == (2, 64)
|
| 387 |
+
assert 1 <= result["n_cycles"] <= 4
|
| 388 |
+
logger.info(" n_cycles=%d, converged=%s, deltas=%s",
|
| 389 |
+
result["n_cycles"], result["converged"],
|
| 390 |
+
[f"{d:.4f}" for d in result["deltas"]])
|
| 391 |
+
results["cyclic_reasoning"] = {
|
| 392 |
+
"ok": True, "n_cycles": result["n_cycles"], "converged": result["converged"],
|
| 393 |
+
}
|
| 394 |
+
logger.info(" [CyclicReasoning] OK")
|
| 395 |
+
except Exception as e:
|
| 396 |
+
logger.error(" [CyclicReasoning] FALHOU: %s", e)
|
| 397 |
+
logger.error(traceback.format_exc())
|
| 398 |
+
results["cyclic_reasoning"] = {"ok": False, "error": str(e)}
|
| 399 |
+
|
| 400 |
+
# ---------- Medusa MTP ----------
|
| 401 |
+
try:
|
| 402 |
+
logger.info(" [MedusaMTP] Testando...")
|
| 403 |
+
cfg = MedusaConfig(
|
| 404 |
+
vocab_size=V, embed_dim=32, n_heads=3,
|
| 405 |
+
head_hidden_mult=2, dropout=0.1,
|
| 406 |
+
)
|
| 407 |
+
medusa = MedusaMTP(cfg).to(device)
|
| 408 |
+
h = torch.randn(2, 8, 32, device=device)
|
| 409 |
+
target = torch.randint(0, V, (2, 8), device=device)
|
| 410 |
+
loss, stats = medusa.compute_loss(h, target)
|
| 411 |
+
assert loss.item() > 0
|
| 412 |
+
# Test tree decode
|
| 413 |
+
base_logits = torch.randn(2, V, device=device)
|
| 414 |
+
td = medusa_tree_decode(medusa, base_logits, h[:, -1, :])
|
| 415 |
+
assert "base_token" in td
|
| 416 |
+
logger.info(" Loss=%.4f, stats=%s", loss.item(), stats)
|
| 417 |
+
results["medusa_mtp"] = {"ok": True, "loss": float(loss), "stats": stats}
|
| 418 |
+
logger.info(" [MedusaMTP] OK")
|
| 419 |
+
except Exception as e:
|
| 420 |
+
logger.error(" [MedusaMTP] FALHOU: %s", e)
|
| 421 |
+
logger.error(traceback.format_exc())
|
| 422 |
+
results["medusa_mtp"] = {"ok": False, "error": str(e)}
|
| 423 |
+
|
| 424 |
+
# ---------- NLG ----------
|
| 425 |
+
try:
|
| 426 |
+
logger.info(" [NLGModule] Testando...")
|
| 427 |
+
cfg = NLGConfig(
|
| 428 |
+
vocab_size=V, embed_dim=32, n_heads=4, n_layers=2,
|
| 429 |
+
max_seq_len=32, pad_id=tokenizer.pad_id, bos_id=tokenizer.bos_id,
|
| 430 |
+
eos_id=tokenizer.eos_id, use_medusa=True, n_medusa_heads=3,
|
| 431 |
+
use_cyclic_reasoning=False, weight_tying=True, device=device,
|
| 432 |
+
)
|
| 433 |
+
nlg = NLGModule(cfg).to(device)
|
| 434 |
+
ids = torch.randint(2, V, (2, 8), device=device)
|
| 435 |
+
target = torch.randint(2, V, (2, 8), device=device)
|
| 436 |
+
out = nlg.compute_loss(ids, target)
|
| 437 |
+
assert "loss" in out and out["loss"].item() > 0
|
| 438 |
+
# Test geração
|
| 439 |
+
prompt = torch.tensor([[tokenizer.bos_id, 5, 10]], dtype=torch.long, device=device)
|
| 440 |
+
gen = nlg.generate(prompt, max_new_tokens=5, temperature=0.7, top_k=10)
|
| 441 |
+
assert gen["ids"].size(1) >= 3
|
| 442 |
+
logger.info(" NLG loss=%.4f, generated %d tokens",
|
| 443 |
+
out["loss"].item(), gen["n_tokens"])
|
| 444 |
+
results["nlg"] = {"ok": True, "loss": float(out["loss"]),
|
| 445 |
+
"n_generated": gen["n_tokens"]}
|
| 446 |
+
logger.info(" [NLGModule] OK")
|
| 447 |
+
except Exception as e:
|
| 448 |
+
logger.error(" [NLGModule] FALHOU: %s", e)
|
| 449 |
+
logger.error(traceback.format_exc())
|
| 450 |
+
results["nlg"] = {"ok": False, "error": str(e)}
|
| 451 |
+
|
| 452 |
+
# ---------- NLP ----------
|
| 453 |
+
try:
|
| 454 |
+
logger.info(" [NLPModule] Testando...")
|
| 455 |
+
nlp_cfg = NLPConfig(
|
| 456 |
+
embed_dim=32, cnn_filters=32, gru_hidden=32,
|
| 457 |
+
feat_per_stream=64, feat_fused=128, n_heads=4,
|
| 458 |
+
num_classes_seq=3, num_labels_tok=5, device=device,
|
| 459 |
+
)
|
| 460 |
+
backbone = CooperativeCNNBiGRU(
|
| 461 |
+
vocab_size=V, embedding_dim=32, cnn_filters=32,
|
| 462 |
+
gru_hidden=32, n_heads=4, pad_idx=tokenizer.pad_id,
|
| 463 |
+
).to(device)
|
| 464 |
+
nlp = NLPModule(nlp_cfg, backbone=backbone).to(device)
|
| 465 |
+
ids_a = torch.randint(2, V, (2, 8), device=device)
|
| 466 |
+
ids_b = torch.randint(2, V, (2, 8), device=device)
|
| 467 |
+
# Test sequence classification
|
| 468 |
+
out_seq = nlp.sequence_classification(ids_a, ids_b)
|
| 469 |
+
assert out_seq["logits"].shape == (2, 3)
|
| 470 |
+
# Test token classification
|
| 471 |
+
tok_logits = nlp.token_classification(ids_a, ids_b)
|
| 472 |
+
# Test span detection
|
| 473 |
+
s_log, e_log = nlp.span_detection(ids_a, ids_b)
|
| 474 |
+
# Test embedding
|
| 475 |
+
emb = nlp.embed(ids_a, ids_b)
|
| 476 |
+
assert emb.shape == (2, 64) # feat_per_stream
|
| 477 |
+
logger.info(" Seq: %s, Tok: %s, Span: (%s,%s), Emb: %s",
|
| 478 |
+
out_seq["logits"].shape, tok_logits.shape if tok_logits is not None else None,
|
| 479 |
+
s_log.shape if s_log is not None else None,
|
| 480 |
+
e_log.shape if e_log is not None else None,
|
| 481 |
+
emb.shape)
|
| 482 |
+
results["nlp"] = {"ok": True, "seq_logits": list(out_seq["logits"].shape),
|
| 483 |
+
"emb_shape": list(emb.shape)}
|
| 484 |
+
logger.info(" [NLPModule] OK")
|
| 485 |
+
except Exception as e:
|
| 486 |
+
logger.error(" [NLPModule] FALHOU: %s", e)
|
| 487 |
+
logger.error(traceback.format_exc())
|
| 488 |
+
results["nlp"] = {"ok": False, "error": str(e)}
|
| 489 |
+
|
| 490 |
+
# ---------- Multimodal Multi-Head Attention ----------
|
| 491 |
+
try:
|
| 492 |
+
logger.info(" [MultimodalMultiHeadAttention] Testando...")
|
| 493 |
+
cfg = MultimodalAttentionConfig(
|
| 494 |
+
d_text_a=32, d_text_b=32, d_image=16, d_audio=16,
|
| 495 |
+
d_model=64, n_heads=4, use_modality_gate=True,
|
| 496 |
+
)
|
| 497 |
+
mha = MultimodalMultiHeadAttention(cfg).to(device)
|
| 498 |
+
seq_a = torch.randn(2, 8, 32, device=device)
|
| 499 |
+
seq_b = torch.randn(2, 6, 32, device=device)
|
| 500 |
+
seq_img = torch.randn(2, 4, 16, device=device)
|
| 501 |
+
seq_aud = torch.randn(2, 5, 16, device=device)
|
| 502 |
+
out = mha(seq_a, seq_b, seq_img, seq_aud)
|
| 503 |
+
assert out["fused"].shape == (2, 64)
|
| 504 |
+
assert out["modality_weights"].shape == (2, 4)
|
| 505 |
+
logger.info(" Fused: %s, modality_weights: %s",
|
| 506 |
+
out["fused"].shape, out["modality_weights"])
|
| 507 |
+
results["multimodal_attention"] = {
|
| 508 |
+
"ok": True, "fused_shape": list(out["fused"].shape),
|
| 509 |
+
}
|
| 510 |
+
logger.info(" [MultimodalMultiHeadAttention] OK")
|
| 511 |
+
except Exception as e:
|
| 512 |
+
logger.error(" [MultimodalMultiHeadAttention] FALHOU: %s", e)
|
| 513 |
+
logger.error(traceback.format_exc())
|
| 514 |
+
results["multimodal_attention"] = {"ok": False, "error": str(e)}
|
| 515 |
+
|
| 516 |
+
# ---------- VQ-VAE-2 ----------
|
| 517 |
+
try:
|
| 518 |
+
logger.info(" [VQVAE2] Testando...")
|
| 519 |
+
cfg = VQVAE2Config(
|
| 520 |
+
in_channels=1, bottom_channels=8, top_channels=4,
|
| 521 |
+
n_bottom_codes=32, n_top_codes=32,
|
| 522 |
+
n_downsample=1, hidden_channels=8, use_ema=True,
|
| 523 |
+
)
|
| 524 |
+
vqvae = VQVAE2(cfg).to(device)
|
| 525 |
+
x = torch.randn(2, 1, 16, 16, device=device)
|
| 526 |
+
out = vqvae(x)
|
| 527 |
+
assert out["x_recon"].shape == x.shape
|
| 528 |
+
assert out["loss"].item() > 0
|
| 529 |
+
logger.info(" Recon: %s, loss=%.4f, top_usage=%.2f, bottom_usage=%.2f",
|
| 530 |
+
out["x_recon"].shape, out["loss"].item(),
|
| 531 |
+
float(out["loss_dict"]["top_usage"]),
|
| 532 |
+
float(out["loss_dict"]["bottom_usage"]))
|
| 533 |
+
results["vqvae2"] = {
|
| 534 |
+
"ok": True, "loss": float(out["loss"]),
|
| 535 |
+
"top_usage": float(out["loss_dict"]["top_usage"]),
|
| 536 |
+
"bottom_usage": float(out["loss_dict"]["bottom_usage"]),
|
| 537 |
+
}
|
| 538 |
+
logger.info(" [VQVAE2] OK")
|
| 539 |
+
except Exception as e:
|
| 540 |
+
logger.error(" [VQVAE2] FALHOU: %s", e)
|
| 541 |
+
logger.error(traceback.format_exc())
|
| 542 |
+
results["vqvae2"] = {"ok": False, "error": str(e)}
|
| 543 |
+
|
| 544 |
+
# ---------- W8A8 Quantization ----------
|
| 545 |
+
try:
|
| 546 |
+
logger.info(" [W8A8 SmoothQuant] Testando...")
|
| 547 |
+
# Aplica ao modelo multimodal (sem dataloader = quantização dinâmica)
|
| 548 |
+
model_copy = MultimodalCNNBiGRU(
|
| 549 |
+
vocab_size=V, num_classes=3, embedding_dim=32, cnn_filters=32,
|
| 550 |
+
gru_hidden=32, n_heads=4, pad_idx=tokenizer.pad_id,
|
| 551 |
+
img_channels=1, img_hidden=16, img_out_dim=32,
|
| 552 |
+
audio_freq=40, audio_hidden=16, audio_out_dim=32, fusion_dim=64,
|
| 553 |
+
).to(device)
|
| 554 |
+
orthogonal_init_model(model_copy)
|
| 555 |
+
qmodel = quantize_model_w8a8(model_copy, dataloader=None, alpha=0.5, device=device)
|
| 556 |
+
assert SmoothQuantizer.is_quantized(qmodel)
|
| 557 |
+
# Verificar forward ainda funciona
|
| 558 |
+
ids_a = torch.randint(2, V, (2, 8), device=device)
|
| 559 |
+
ids_b = torch.randint(2, V, (2, 8), device=device)
|
| 560 |
+
with torch.no_grad():
|
| 561 |
+
out = qmodel(ids_a, ids_b, mode="classify")
|
| 562 |
+
assert out["logits"].shape == (2, 3)
|
| 563 |
+
savings = estimate_memory_savings(qmodel)
|
| 564 |
+
logger.info(" Quantizado: %s, savings: %.1f%%",
|
| 565 |
+
SmoothQuantizer.is_quantized(qmodel), savings["reduction_pct"])
|
| 566 |
+
results["w8a8"] = {"ok": True, "savings": savings}
|
| 567 |
+
logger.info(" [W8A8] OK")
|
| 568 |
+
except Exception as e:
|
| 569 |
+
logger.error(" [W8A8] FALHOU: %s", e)
|
| 570 |
+
logger.error(traceback.format_exc())
|
| 571 |
+
results["w8a8"] = {"ok": False, "error": str(e)}
|
| 572 |
+
|
| 573 |
+
# ---------- Long Context (1M tokens) ----------
|
| 574 |
+
try:
|
| 575 |
+
logger.info(" [LongContextManager - 1M tokens] Testando...")
|
| 576 |
+
lcw = make_long_context_window(
|
| 577 |
+
max_window=1_000_000, strategy="chunked",
|
| 578 |
+
chunk_size=8192, embed_dim=64, n_heads=4, n_layers=2, device=device,
|
| 579 |
+
)
|
| 580 |
+
# Simular adicionar muitos tokens em chunks
|
| 581 |
+
all_tokens = []
|
| 582 |
+
for chunk_idx in range(3): # 3 chunks de 100 tokens cada
|
| 583 |
+
tokens = torch.tensor(
|
| 584 |
+
[[i + 1 + chunk_idx * 100 for i in range(100)]], device=device,
|
| 585 |
+
)
|
| 586 |
+
result = lcw.append_tokens_chunked(tokens)
|
| 587 |
+
all_tokens.append(result)
|
| 588 |
+
info = lcw.get_info()
|
| 589 |
+
assert info["supports_1m_tokens"]
|
| 590 |
+
# LongContextManager.get_info() retorna "history_tokens" (não "total_tokens")
|
| 591 |
+
assert info["history_tokens"] == 300 # 3 chunks * 100
|
| 592 |
+
logger.info(" Info: %s", info)
|
| 593 |
+
results["long_context"] = {"ok": True, "info": info}
|
| 594 |
+
logger.info(" [LongContextManager] OK")
|
| 595 |
+
except Exception as e:
|
| 596 |
+
logger.error(" [LongContextManager] FALHOU: %s", e)
|
| 597 |
+
logger.error(traceback.format_exc())
|
| 598 |
+
results["long_context"] = {"ok": False, "error": str(e)}
|
| 599 |
+
|
| 600 |
+
# ---------- EWC (já existente, mas validar integração) ----------
|
| 601 |
+
try:
|
| 602 |
+
logger.info(" [EWC] Testando integração...")
|
| 603 |
+
ewc_cfg = EWCConfig(
|
| 604 |
+
enabled=True, lambda_ewc=100.0,
|
| 605 |
+
n_samples_fisher=5, online_gamma=0.9, device=device,
|
| 606 |
+
)
|
| 607 |
+
ewc_state = EWCState(ewc_cfg)
|
| 608 |
+
pen0 = ewc_state.penalty(models["model"])
|
| 609 |
+
assert pen0.item() == 0.0
|
| 610 |
+
|
| 611 |
+
# Forward_fn para Fisher
|
| 612 |
+
V = tokenizer.vocab_size
|
| 613 |
+
B, T = 2, 8
|
| 614 |
+
ids_a = torch.randint(0, V, (B, T), device=device)
|
| 615 |
+
ids_b = torch.randint(0, V, (B, T), device=device)
|
| 616 |
+
|
| 617 |
+
def forward_fn(_idx=None):
|
| 618 |
+
out = models["model"](ids_a, ids_b, mode="classify")
|
| 619 |
+
return out["logits"]
|
| 620 |
+
|
| 621 |
+
ewc_state.consolidate(models["model"], forward_fn=forward_fn)
|
| 622 |
+
assert ewc_state.num_tasks() == 1
|
| 623 |
+
pen1 = ewc_state.penalty(models["model"])
|
| 624 |
+
assert pen1.item() >= 0.0
|
| 625 |
+
logger.info(" Penalty antes/depois consolidar: %.6f / %.6f",
|
| 626 |
+
float(pen0), float(pen1))
|
| 627 |
+
results["ewc"] = {
|
| 628 |
+
"ok": True, "penalty_before": float(pen0),
|
| 629 |
+
"penalty_after": float(pen1), "num_tasks": ewc_state.num_tasks(),
|
| 630 |
+
}
|
| 631 |
+
results["ewc_state"] = ewc_state
|
| 632 |
+
logger.info(" [EWC] OK")
|
| 633 |
+
except Exception as e:
|
| 634 |
+
logger.error(" [EWC] FALHOU: %s", e)
|
| 635 |
+
logger.error(traceback.format_exc())
|
| 636 |
+
results["ewc"] = {"ok": False, "error": str(e)}
|
| 637 |
+
|
| 638 |
+
# ---------- Context Window (curto, já existente) ----------
|
| 639 |
+
try:
|
| 640 |
+
logger.info(" [ContextWindowManager] Testando...")
|
| 641 |
+
cw = make_default_context_window(
|
| 642 |
+
max_window=32, n_sink=2, embed_dim=64, n_heads=4, n_layers=2, device=device,
|
| 643 |
+
)
|
| 644 |
+
cache = cw.init_cache(batch_size=1, device=torch.device(device))
|
| 645 |
+
assert cache is not None
|
| 646 |
+
tokens = torch.tensor([[1, 2, 3, 4, 5]], dtype=torch.long, device=device)
|
| 647 |
+
window = cw.append_tokens(tokens)
|
| 648 |
+
assert window.size(1) == 5
|
| 649 |
+
logger.info(" Window: %s, cache_len: %d", window.shape, cache.get_seq_len())
|
| 650 |
+
results["context_window"] = {"ok": True}
|
| 651 |
+
logger.info(" [ContextWindowManager] OK")
|
| 652 |
+
except Exception as e:
|
| 653 |
+
logger.error(" [ContextWindowManager] FALHOU: %s", e)
|
| 654 |
+
logger.error(traceback.format_exc())
|
| 655 |
+
results["context_window"] = {"ok": False, "error": str(e)}
|
| 656 |
+
|
| 657 |
+
return results
|
| 658 |
+
|
| 659 |
+
|
| 660 |
+
# ============================================================================
|
| 661 |
+
# Step 5: Treinamento
|
| 662 |
+
# ============================================================================
|
| 663 |
+
|
| 664 |
+
def step_05_train(
|
| 665 |
+
models: Dict,
|
| 666 |
+
tokenizer: BBPETokenizer,
|
| 667 |
+
dataloader: DataLoader,
|
| 668 |
+
device: str = "cpu",
|
| 669 |
+
ewc_state: EWCState = None,
|
| 670 |
+
monitor: Monitor = None,
|
| 671 |
+
batch_idx: int = 0,
|
| 672 |
+
) -> Dict:
|
| 673 |
+
"""Executa treinamento cooperativo em um batch de 100 amostras."""
|
| 674 |
+
logger.info("=" * 70)
|
| 675 |
+
logger.info("PASSO 5: Treinamento cooperativo (batch %d/5 - 100 amostras)", batch_idx + 1)
|
| 676 |
+
logger.info("=" * 70)
|
| 677 |
+
|
| 678 |
+
cfg = TrainerConfig(
|
| 679 |
+
num_epochs=1,
|
| 680 |
+
max_batches_per_epoch=12,
|
| 681 |
+
batch_size=8,
|
| 682 |
+
n_synergy_attempts=2,
|
| 683 |
+
use_synergy_search=True,
|
| 684 |
+
use_hypotheses=True,
|
| 685 |
+
n_hypotheses=3,
|
| 686 |
+
loss_config=LossConfig(
|
| 687 |
+
alpha=1.0, beta=0.5, gamma_loss=0.3, delta=0.01,
|
| 688 |
+
lambda_penal=0.1, mu_exp_penal=0.05,
|
| 689 |
+
gamma_exp=1.0, threshold_err=0.5,
|
| 690 |
+
l2_reg=1e-5, use_curvature=True, curvature_eps=1e-3,
|
| 691 |
+
),
|
| 692 |
+
auto_config=AutoLearnConfig(
|
| 693 |
+
kappa_curv=0.01, grad_clip=1.0, spectral_radius=1.0,
|
| 694 |
+
lr_min=1e-6, lr_max=1e-2,
|
| 695 |
+
initial_lr_G=1e-3, initial_lr_V=1e-3,
|
| 696 |
+
use_spectral_norm=True, apply_after_step=True,
|
| 697 |
+
l2_reg=1e-5,
|
| 698 |
+
),
|
| 699 |
+
ewc_config=ewc_state.config if ewc_state else None,
|
| 700 |
+
device=device,
|
| 701 |
+
log_every=2,
|
| 702 |
+
use_verifier_real=True,
|
| 703 |
+
use_generator=True,
|
| 704 |
+
use_hypothesis_output=True,
|
| 705 |
+
)
|
| 706 |
+
|
| 707 |
+
trainer = CooperativeTrainer(
|
| 708 |
+
model=models["model"],
|
| 709 |
+
tokenizer=tokenizer,
|
| 710 |
+
config=cfg,
|
| 711 |
+
generator=models["generator"],
|
| 712 |
+
verifier=models["verifier"],
|
| 713 |
+
anti_hallucination=models["anti_hall"],
|
| 714 |
+
ewc_state=ewc_state,
|
| 715 |
+
)
|
| 716 |
+
|
| 717 |
+
# Integra monitor
|
| 718 |
+
if monitor is not None:
|
| 719 |
+
monitor.start_epoch(batch_idx)
|
| 720 |
+
|
| 721 |
+
try:
|
| 722 |
+
result = trainer.train(dataloader)
|
| 723 |
+
logger.info(" Treino OK batch %d | loss=%.4f | ppl=%.2f | elapsed=%.1fs",
|
| 724 |
+
batch_idx + 1, result["final_loss"], result["final_ppl"],
|
| 725 |
+
result["elapsed_s"])
|
| 726 |
+
|
| 727 |
+
# Log no monitor
|
| 728 |
+
if monitor is not None:
|
| 729 |
+
for h in result.get("history", []):
|
| 730 |
+
monitor.log_batch({
|
| 731 |
+
"loss": h.get("loss", 0),
|
| 732 |
+
"ppl": h.get("ppl", 0),
|
| 733 |
+
"lr_g": h.get("lr_G", 0),
|
| 734 |
+
"hypothesis_activations": h.get("hypothesis_activations", 0),
|
| 735 |
+
"elapsed_ms": h.get("elapsed_ms", 0),
|
| 736 |
+
})
|
| 737 |
+
monitor.end_epoch({"batch_idx": batch_idx})
|
| 738 |
+
if ewc_state is not None:
|
| 739 |
+
pen = ewc_state.penalty(models["model"]).item()
|
| 740 |
+
monitor.log_ewc_penalty(penalty=pen, num_tasks=ewc_state.num_tasks())
|
| 741 |
+
|
| 742 |
+
return result
|
| 743 |
+
except Exception as e:
|
| 744 |
+
logger.error(" Treino FALHOU batch %d: %s", batch_idx + 1, e)
|
| 745 |
+
logger.error(traceback.format_exc())
|
| 746 |
+
raise
|
| 747 |
+
|
| 748 |
+
|
| 749 |
+
# ============================================================================
|
| 750 |
+
# Step 6: Inferência
|
| 751 |
+
# ============================================================================
|
| 752 |
+
|
| 753 |
+
def step_06_inference(
|
| 754 |
+
model: MultimodalCNNBiGRU,
|
| 755 |
+
tokenizer: BBPETokenizer,
|
| 756 |
+
device: str = "cpu",
|
| 757 |
+
generator: GeneratorCNNBiGRU = None,
|
| 758 |
+
monitor: Monitor = None,
|
| 759 |
+
) -> Dict:
|
| 760 |
+
"""Executa inferência com sampling."""
|
| 761 |
+
logger.info("=" * 70)
|
| 762 |
+
logger.info("PASSO 6: Inferência com Temperatura + Top-K + Top-P + Penalidades")
|
| 763 |
+
logger.info("=" * 70)
|
| 764 |
+
|
| 765 |
+
prompts = [
|
| 766 |
+
("o modelo coopera entre", "fluxos paralelos"),
|
| 767 |
+
("atencao cruzada troca", "informacoes entre redes"),
|
| 768 |
+
("gru bidirecional captura", "dependencias temporais"),
|
| 769 |
+
("medusa heads predizem", "multiplos tokens"),
|
| 770 |
+
("raciocinio ciclico refina", "iterativamente"),
|
| 771 |
+
]
|
| 772 |
+
|
| 773 |
+
results = []
|
| 774 |
+
for prompt_a, prompt_b in prompts:
|
| 775 |
+
t0 = time.time()
|
| 776 |
+
try:
|
| 777 |
+
result = generate_with_sampling(
|
| 778 |
+
model=model, tokenizer=tokenizer,
|
| 779 |
+
prompt_a=prompt_a, prompt_b=prompt_b,
|
| 780 |
+
max_new_tokens=10, temperature=0.7,
|
| 781 |
+
top_k=20, top_p=0.9,
|
| 782 |
+
presence_penalty=0.3, frequency_penalty=0.3,
|
| 783 |
+
device=device, generator=generator,
|
| 784 |
+
)
|
| 785 |
+
elapsed_ms = (time.time() - t0) * 1000
|
| 786 |
+
logger.info(" Prompt A: %s | B: %s", prompt_a, prompt_b)
|
| 787 |
+
logger.info(" Gerado: %s (%.0fms)",
|
| 788 |
+
result["text"][:80], elapsed_ms)
|
| 789 |
+
if monitor is not None:
|
| 790 |
+
monitor.log_inference(
|
| 791 |
+
n_tokens=len(result.get("token_ids", [])),
|
| 792 |
+
elapsed_ms=elapsed_ms,
|
| 793 |
+
)
|
| 794 |
+
results.append({"prompt_a": prompt_a, "prompt_b": prompt_b, **result})
|
| 795 |
+
except Exception as e:
|
| 796 |
+
logger.error(" Inferência falhou (%s, %s): %s", prompt_a, prompt_b, e)
|
| 797 |
+
results.append({"prompt_a": prompt_a, "prompt_b": prompt_b, "error": str(e)})
|
| 798 |
+
|
| 799 |
+
return {"results": results}
|
| 800 |
+
|
| 801 |
+
|
| 802 |
+
# ============================================================================
|
| 803 |
+
# Step 7: PPL
|
| 804 |
+
# ============================================================================
|
| 805 |
+
|
| 806 |
+
def step_07_eval_ppl(
|
| 807 |
+
model: MultimodalCNNBiGRU,
|
| 808 |
+
tokenizer: BBPETokenizer,
|
| 809 |
+
dataloader: DataLoader,
|
| 810 |
+
device: str = "cpu",
|
| 811 |
+
generator: GeneratorCNNBiGRU = None,
|
| 812 |
+
) -> Dict:
|
| 813 |
+
"""Avalia perplexidade."""
|
| 814 |
+
logger.info("=" * 70)
|
| 815 |
+
logger.info("PASSO 7: Avaliação de Perplexidade (PPL)")
|
| 816 |
+
logger.info("=" * 70)
|
| 817 |
+
try:
|
| 818 |
+
result = evaluate_perplexity(
|
| 819 |
+
model=model, dataloader=dataloader, tokenizer=tokenizer,
|
| 820 |
+
device=device, max_batches=5, generator=generator,
|
| 821 |
+
)
|
| 822 |
+
logger.info(" Loss: %.4f | PPL: %.2f | batches: %d | used_generator: %s",
|
| 823 |
+
result["loss"], result["ppl"], result["n_batches"],
|
| 824 |
+
result.get("used_generator", False))
|
| 825 |
+
return result
|
| 826 |
+
except Exception as e:
|
| 827 |
+
logger.error(" PPL falhou: %s", e)
|
| 828 |
+
logger.error(traceback.format_exc())
|
| 829 |
+
return {"error": str(e)}
|
| 830 |
+
|
| 831 |
+
|
| 832 |
+
# ============================================================================
|
| 833 |
+
# Step 8: Relatório
|
| 834 |
+
# ============================================================================
|
| 835 |
+
|
| 836 |
+
def step_08_report(errors: List[str], warnings: List[str]) -> None:
|
| 837 |
+
"""Reporta erros lógicos ou falhas encontradas."""
|
| 838 |
+
logger.info("=" * 70)
|
| 839 |
+
logger.info("PASSO 8: Relatório de erros lógicos ou falhas")
|
| 840 |
+
logger.info("=" * 70)
|
| 841 |
+
|
| 842 |
+
if not errors:
|
| 843 |
+
logger.info(" OK — NENHUM ERRO CRÍTICO encontrado")
|
| 844 |
+
else:
|
| 845 |
+
logger.error(" FAIL — %d ERRO(S) CRÍTICO(S):", len(errors))
|
| 846 |
+
for e in errors:
|
| 847 |
+
logger.error(" - %s", e)
|
| 848 |
+
|
| 849 |
+
if not warnings:
|
| 850 |
+
logger.info(" OK — Nenhum warning relevante")
|
| 851 |
+
else:
|
| 852 |
+
logger.warning(" WARN — %d warning(s):", len(warnings))
|
| 853 |
+
for w in warnings:
|
| 854 |
+
logger.warning(" - %s", w)
|
| 855 |
+
logger.info("=" * 70)
|
| 856 |
+
|
| 857 |
+
|
| 858 |
+
# ============================================================================
|
| 859 |
+
# Main
|
| 860 |
+
# ============================================================================
|
| 861 |
+
|
| 862 |
+
def main():
|
| 863 |
+
"""Executa o teste completo de 500 amostras (5 batches de 100)."""
|
| 864 |
+
logger.info("=" * 70)
|
| 865 |
+
logger.info("INICIANDO TESTE DE 500 AMOSTRAS (5 batches de 100)")
|
| 866 |
+
logger.info("CNN-BiGRU MULTIMODAL COOPERATIVO v3.0")
|
| 867 |
+
logger.info("=" * 70)
|
| 868 |
+
logger.info("Project root: %s", PROJECT_ROOT)
|
| 869 |
+
logger.info("Device: %s | Cores: %d", "cpu", N_CORES)
|
| 870 |
+
logger.info("Total samples: %d | Batch size: %d | N batches: %d",
|
| 871 |
+
TOTAL_SAMPLES, BATCH_SIZE_SAMPLES, N_BATCHES)
|
| 872 |
+
|
| 873 |
+
errors: List[str] = []
|
| 874 |
+
warnings: List[str] = []
|
| 875 |
+
all_train_results = []
|
| 876 |
+
|
| 877 |
+
try:
|
| 878 |
+
# Step 1: Runtime + Monitor
|
| 879 |
+
rt_info = step_01_init_runtime()
|
| 880 |
+
monitor = rt_info["monitor"]
|
| 881 |
+
monitor.start_training()
|
| 882 |
+
|
| 883 |
+
# Step 2: Tokenizer
|
| 884 |
+
tokenizer, roundtrip = step_02_train_tokenizer()
|
| 885 |
+
if roundtrip < 0.8:
|
| 886 |
+
warnings.append(f"BBPE roundtrip = {roundtrip:.2%} (esperado >= 80%)")
|
| 887 |
+
|
| 888 |
+
# Step 3: Modelo (uma única instância, reutilizada em todos batches)
|
| 889 |
+
models = step_04_init_model(tokenizer, device="cpu")
|
| 890 |
+
|
| 891 |
+
# Forward pass inicial
|
| 892 |
+
try:
|
| 893 |
+
test_loader = step_03_create_dataset(tokenizer, n_samples=8, seed_offset=99)
|
| 894 |
+
batch = next(iter(test_loader))
|
| 895 |
+
with torch.no_grad():
|
| 896 |
+
out = models["model"](
|
| 897 |
+
batch["input_ids_a"], batch["input_ids_b"],
|
| 898 |
+
images=batch["images"], audios=batch["audios"],
|
| 899 |
+
mode="classify",
|
| 900 |
+
)
|
| 901 |
+
logger.info(" Forward pass OK | logits shape: %s", out["logits"].shape)
|
| 902 |
+
assert out["logits"].shape == (batch["input_ids_a"].size(0), 3), \
|
| 903 |
+
f"Shape inesperado: {out['logits'].shape}"
|
| 904 |
+
except Exception as e:
|
| 905 |
+
errors.append(f"Forward pass inicial falhou: {e}")
|
| 906 |
+
logger.error(traceback.format_exc())
|
| 907 |
+
|
| 908 |
+
# Step 5b: Testar novos módulos v3.0
|
| 909 |
+
ewc_state = None
|
| 910 |
+
if not errors:
|
| 911 |
+
try:
|
| 912 |
+
new_modules_result = step_05b_test_v3_modules(models, tokenizer, device="cpu")
|
| 913 |
+
for mod_name, mod_res in new_modules_result.items():
|
| 914 |
+
if mod_name == "ewc_state":
|
| 915 |
+
continue
|
| 916 |
+
if isinstance(mod_res, dict) and not mod_res.get("ok", True):
|
| 917 |
+
errors.append(f"Módulo {mod_name} falhou: {mod_res.get('error', 'unknown')}")
|
| 918 |
+
# Não é crítico para alguns módulos — converter em warning
|
| 919 |
+
if mod_name in ("nlg", "nlp", "multimodal_attention", "vqvae2",
|
| 920 |
+
"w8a8", "long_context"):
|
| 921 |
+
errors.pop() # remove o erro
|
| 922 |
+
warnings.append(f"Módulo {mod_name} falhou (não crítico): {mod_res.get('error', 'unknown')}")
|
| 923 |
+
ewc_state = new_modules_result.get("ewc_state")
|
| 924 |
+
if ewc_state is None:
|
| 925 |
+
warnings.append("EWC state não criado — EWC não será testado no treino")
|
| 926 |
+
# Log no monitor
|
| 927 |
+
if monitor is not None:
|
| 928 |
+
if new_modules_result.get("cyclic_reasoning", {}).get("ok"):
|
| 929 |
+
monitor.log_cyclic_reasoning(new_modules_result["cyclic_reasoning"])
|
| 930 |
+
if new_modules_result.get("vqvae2", {}).get("ok"):
|
| 931 |
+
monitor.log_vqvae2_usage(new_modules_result["vqvae2"])
|
| 932 |
+
if new_modules_result.get("w8a8", {}).get("ok"):
|
| 933 |
+
monitor.log_quantization(new_modules_result["w8a8"].get("savings", {}))
|
| 934 |
+
except Exception as e:
|
| 935 |
+
errors.append(f"Teste de novos módulos falhou: {e}")
|
| 936 |
+
logger.error(traceback.format_exc())
|
| 937 |
+
|
| 938 |
+
# Step 5 + 6 + 7: Loop sobre 5 batches de 100 amostras
|
| 939 |
+
for batch_idx in range(N_BATCHES):
|
| 940 |
+
logger.info("")
|
| 941 |
+
logger.info("#" * 70)
|
| 942 |
+
logger.info("# BATCH %d/%d — 100 AMOSTRAS", batch_idx + 1, N_BATCHES)
|
| 943 |
+
logger.info("#" * 70)
|
| 944 |
+
|
| 945 |
+
try:
|
| 946 |
+
dataloader = step_03_create_dataset(
|
| 947 |
+
tokenizer, n_samples=BATCH_SIZE_SAMPLES, seed_offset=batch_idx,
|
| 948 |
+
)
|
| 949 |
+
except Exception as e:
|
| 950 |
+
errors.append(f"Criação dataset batch {batch_idx+1} falhou: {e}")
|
| 951 |
+
continue
|
| 952 |
+
|
| 953 |
+
# Treino
|
| 954 |
+
if not errors:
|
| 955 |
+
try:
|
| 956 |
+
train_result = step_05_train(
|
| 957 |
+
models, tokenizer, dataloader, device="cpu",
|
| 958 |
+
ewc_state=ewc_state, monitor=monitor, batch_idx=batch_idx,
|
| 959 |
+
)
|
| 960 |
+
all_train_results.append(train_result)
|
| 961 |
+
except Exception as e:
|
| 962 |
+
errors.append(f"Treino batch {batch_idx+1} falhou: {e}")
|
| 963 |
+
|
| 964 |
+
# Inferência (apenas no último batch para economizar tempo)
|
| 965 |
+
if batch_idx == N_BATCHES - 1 and not errors:
|
| 966 |
+
try:
|
| 967 |
+
step_06_inference(
|
| 968 |
+
models["model"], tokenizer, device="cpu",
|
| 969 |
+
generator=models.get("generator"), monitor=monitor,
|
| 970 |
+
)
|
| 971 |
+
except Exception as e:
|
| 972 |
+
errors.append(f"Inferência batch {batch_idx+1} falhou: {e}")
|
| 973 |
+
logger.error(traceback.format_exc())
|
| 974 |
+
|
| 975 |
+
# PPL
|
| 976 |
+
if not errors:
|
| 977 |
+
try:
|
| 978 |
+
step_07_eval_ppl(
|
| 979 |
+
models["model"], tokenizer, dataloader, device="cpu",
|
| 980 |
+
generator=models.get("generator"),
|
| 981 |
+
)
|
| 982 |
+
except Exception as e:
|
| 983 |
+
warnings.append(f"PPL batch {batch_idx+1} falhou (não crítico): {e}")
|
| 984 |
+
|
| 985 |
+
# Step 8: Relatório
|
| 986 |
+
step_08_report(errors, warnings)
|
| 987 |
+
|
| 988 |
+
# Finalizar monitor
|
| 989 |
+
monitor.end_training()
|
| 990 |
+
report_path = monitor.export_report()
|
| 991 |
+
csv_path = monitor.export_csv()
|
| 992 |
+
md_path = monitor.export_markdown_summary()
|
| 993 |
+
|
| 994 |
+
# Resumo final
|
| 995 |
+
logger.info("=" * 70)
|
| 996 |
+
logger.info("RESUMO FINAL DO TESTE DE 500 AMOSTRAS")
|
| 997 |
+
logger.info("=" * 70)
|
| 998 |
+
logger.info(" Amostras processadas: %d (5 batches x 100)", TOTAL_SAMPLES)
|
| 999 |
+
logger.info(" Erros críticos: %d", len(errors))
|
| 1000 |
+
logger.info(" Warnings: %d", len(warnings))
|
| 1001 |
+
if all_train_results:
|
| 1002 |
+
final = all_train_results[-1]
|
| 1003 |
+
logger.info(" Loss final: %.4f", final["final_loss"])
|
| 1004 |
+
logger.info(" PPL final: %.2f", final["final_ppl"])
|
| 1005 |
+
logger.info(" Tempo total treino: %.1fs",
|
| 1006 |
+
sum(r["elapsed_s"] for r in all_train_results))
|
| 1007 |
+
logger.info(" Monitor report: %s", report_path)
|
| 1008 |
+
logger.info(" Monitor CSV: %s", csv_path)
|
| 1009 |
+
logger.info(" Monitor MD: %s", md_path)
|
| 1010 |
+
logger.info(" Throughput: %s", monitor.get_inference_throughput())
|
| 1011 |
+
|
| 1012 |
+
if errors:
|
| 1013 |
+
logger.error(" STATUS: FALHA — %d erro(s)", len(errors))
|
| 1014 |
+
return 1
|
| 1015 |
+
else:
|
| 1016 |
+
logger.info(" STATUS: SUCESSO")
|
| 1017 |
+
return 0
|
| 1018 |
+
|
| 1019 |
+
except Exception as e:
|
| 1020 |
+
logger.error("ERRO FATAL: %s", e)
|
| 1021 |
+
logger.error(traceback.format_exc())
|
| 1022 |
+
return 2
|
| 1023 |
+
|
| 1024 |
+
|
| 1025 |
+
if __name__ == "__main__":
|
| 1026 |
+
try:
|
| 1027 |
+
rc = main()
|
| 1028 |
+
except SystemExit:
|
| 1029 |
+
raise
|
| 1030 |
+
except Exception as e:
|
| 1031 |
+
logger.error("Unhandled exception: %s", e)
|
| 1032 |
+
rc = 2
|
| 1033 |
+
# Evita o "Fatal Python error: PyGILState_Release" no shutdown
|
| 1034 |
+
import os
|
| 1035 |
+
os._exit(rc)
|
utils/monitoring.py
ADDED
|
@@ -0,0 +1,729 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""monitoring.py — Monitoramento completo e evolução do CNN-BiGRU.
|
| 2 |
+
|
| 3 |
+
Rastreia métricas de treinamento, inferência, uso de recursos, evolução
|
| 4 |
+
dos componentes (EWC, Medusa, CyclicReasoning, etc.) e exporta relatórios.
|
| 5 |
+
|
| 6 |
+
==============================================================================
|
| 7 |
+
FUNCIONALIDADES
|
| 8 |
+
==============================================================================
|
| 9 |
+
|
| 10 |
+
1. Métricas de treinamento:
|
| 11 |
+
- Loss total, loss_main, loss_medusa, loss_ewc, loss_hallucination
|
| 12 |
+
- Perplexidade (PPL)
|
| 13 |
+
- Learning rate (G, V)
|
| 14 |
+
- Norma do gradiente
|
| 15 |
+
- Hypothesis activations
|
| 16 |
+
- Synergy history
|
| 17 |
+
|
| 18 |
+
2. Métricas de inferência:
|
| 19 |
+
- Tokens gerados
|
| 20 |
+
- Medusa accept rate
|
| 21 |
+
- Tempo por token (ms)
|
| 22 |
+
- Throughput (tokens/s)
|
| 23 |
+
- Memória usada (MB)
|
| 24 |
+
|
| 25 |
+
3. Métricas de evolução:
|
| 26 |
+
- Loss/PPL ao longo do tempo (por época)
|
| 27 |
+
- Convergência do CyclicReasoning (n_cycles, deltas)
|
| 28 |
+
- Codebook usage do VQ-VAE-2
|
| 29 |
+
- Quantização stats (memória economizada, etc.)
|
| 30 |
+
- EWC penalty evolution
|
| 31 |
+
|
| 32 |
+
4. Métricas de sistema:
|
| 33 |
+
- CPU/Memory usage
|
| 34 |
+
- Device info
|
| 35 |
+
- Runtime info (Xeon AVX512/AMX)
|
| 36 |
+
|
| 37 |
+
5. Exportação:
|
| 38 |
+
- JSON report
|
| 39 |
+
- CSV time series
|
| 40 |
+
- Markdown summary
|
| 41 |
+
|
| 42 |
+
==============================================================================
|
| 43 |
+
USO
|
| 44 |
+
==============================================================================
|
| 45 |
+
|
| 46 |
+
from cnn_bigru.utils.monitoring import Monitor
|
| 47 |
+
|
| 48 |
+
monitor = Monitor(output_dir="/path/to/logs")
|
| 49 |
+
monitor.start_training()
|
| 50 |
+
for epoch in range(N):
|
| 51 |
+
for batch in dataloader:
|
| 52 |
+
...
|
| 53 |
+
monitor.log_batch({
|
| 54 |
+
"loss": loss.item(),
|
| 55 |
+
"ppl": ppl,
|
| 56 |
+
"lr": lr,
|
| 57 |
+
"grad_norm": grad_norm,
|
| 58 |
+
})
|
| 59 |
+
monitor.log_epoch({"epoch": epoch, "avg_loss": ...})
|
| 60 |
+
monitor.end_training()
|
| 61 |
+
monitor.export_report()
|
| 62 |
+
|
| 63 |
+
Autor: CNN-BiGRU Project
|
| 64 |
+
"""
|
| 65 |
+
from __future__ import annotations
|
| 66 |
+
|
| 67 |
+
import csv
|
| 68 |
+
import json
|
| 69 |
+
import logging
|
| 70 |
+
import os
|
| 71 |
+
import time
|
| 72 |
+
from collections import defaultdict, deque
|
| 73 |
+
from dataclasses import dataclass, field, asdict
|
| 74 |
+
from pathlib import Path
|
| 75 |
+
from typing import Any, Dict, List, Optional, Union
|
| 76 |
+
|
| 77 |
+
import torch
|
| 78 |
+
|
| 79 |
+
logger = logging.getLogger(__name__)
|
| 80 |
+
|
| 81 |
+
|
| 82 |
+
# ============================================================================
|
| 83 |
+
# Data Classes
|
| 84 |
+
# ============================================================================
|
| 85 |
+
|
| 86 |
+
@dataclass
|
| 87 |
+
class BatchMetrics:
|
| 88 |
+
"""Métricas de um batch."""
|
| 89 |
+
step: int
|
| 90 |
+
epoch: int
|
| 91 |
+
batch_in_epoch: int
|
| 92 |
+
timestamp: float
|
| 93 |
+
loss: float = 0.0
|
| 94 |
+
loss_main: float = 0.0
|
| 95 |
+
loss_medusa: float = 0.0
|
| 96 |
+
loss_ewc: float = 0.0
|
| 97 |
+
loss_hallucination: float = 0.0
|
| 98 |
+
loss_total: float = 0.0
|
| 99 |
+
ppl: float = 0.0
|
| 100 |
+
lr_g: float = 0.0
|
| 101 |
+
lr_v: float = 0.0
|
| 102 |
+
grad_norm_g: float = 0.0
|
| 103 |
+
grad_norm_v: float = 0.0
|
| 104 |
+
v_mean: float = 0.0
|
| 105 |
+
hypothesis_activations: int = 0
|
| 106 |
+
elapsed_ms: float = 0.0
|
| 107 |
+
extra: Dict[str, Any] = field(default_factory=dict)
|
| 108 |
+
|
| 109 |
+
|
| 110 |
+
@dataclass
|
| 111 |
+
class EpochMetrics:
|
| 112 |
+
"""Métricas agregadas de uma época."""
|
| 113 |
+
epoch: int
|
| 114 |
+
avg_loss: float = 0.0
|
| 115 |
+
avg_ppl: float = 0.0
|
| 116 |
+
n_batches: int = 0
|
| 117 |
+
n_hypothesis_activations: int = 0
|
| 118 |
+
elapsed_s: float = 0.0
|
| 119 |
+
best_loss: float = float("inf")
|
| 120 |
+
worst_loss: float = 0.0
|
| 121 |
+
synergy_history: List[Dict] = field(default_factory=list)
|
| 122 |
+
extra: Dict[str, Any] = field(default_factory=dict)
|
| 123 |
+
|
| 124 |
+
|
| 125 |
+
@dataclass
|
| 126 |
+
class EvolutionMetrics:
|
| 127 |
+
"""Métricas de evolução dos componentes ao longo do tempo."""
|
| 128 |
+
cyclic_reasoning_stats: List[Dict] = field(default_factory=list)
|
| 129 |
+
vqvae2_codebook_usage: List[Dict] = field(default_factory=list)
|
| 130 |
+
quantization_stats: List[Dict] = field(default_factory=list)
|
| 131 |
+
ewc_penalty_evolution: List[Dict] = field(default_factory=list)
|
| 132 |
+
medusa_accept_rate: List[Dict] = field(default_factory=list)
|
| 133 |
+
|
| 134 |
+
|
| 135 |
+
@dataclass
|
| 136 |
+
class SystemMetrics:
|
| 137 |
+
"""Métricas de sistema."""
|
| 138 |
+
timestamp: float
|
| 139 |
+
cpu_percent: float = 0.0
|
| 140 |
+
memory_mb: float = 0.0
|
| 141 |
+
gpu_memory_mb: float = 0.0
|
| 142 |
+
n_active_threads: int = 0
|
| 143 |
+
|
| 144 |
+
|
| 145 |
+
# ============================================================================
|
| 146 |
+
# Monitor
|
| 147 |
+
# ============================================================================
|
| 148 |
+
|
| 149 |
+
class Monitor:
|
| 150 |
+
"""Monitor completo de treinamento, inferência e evolução.
|
| 151 |
+
|
| 152 |
+
Args:
|
| 153 |
+
output_dir: diretório para salvar logs e relatórios
|
| 154 |
+
max_history: número máximo de métricas em memória (FIFO)
|
| 155 |
+
log_every: intervalo de log (em batches)
|
| 156 |
+
track_system: se True, rastreia CPU/memória (requer psutil)
|
| 157 |
+
"""
|
| 158 |
+
|
| 159 |
+
def __init__(
|
| 160 |
+
self,
|
| 161 |
+
output_dir: Optional[Union[str, Path]] = None,
|
| 162 |
+
max_history: int = 10000,
|
| 163 |
+
log_every: int = 1,
|
| 164 |
+
track_system: bool = True,
|
| 165 |
+
):
|
| 166 |
+
self.output_dir = Path(output_dir) if output_dir else None
|
| 167 |
+
if self.output_dir:
|
| 168 |
+
self.output_dir.mkdir(parents=True, exist_ok=True)
|
| 169 |
+
self.max_history = max_history
|
| 170 |
+
self.log_every = log_every
|
| 171 |
+
self.track_system = track_system and self._psutil_available()
|
| 172 |
+
|
| 173 |
+
# Histórico
|
| 174 |
+
self.batch_history: deque = deque(maxlen=max_history)
|
| 175 |
+
self.epoch_history: List[EpochMetrics] = []
|
| 176 |
+
self.evolution = EvolutionMetrics()
|
| 177 |
+
self.system_history: deque = deque(maxlen=max_history)
|
| 178 |
+
|
| 179 |
+
# Estado
|
| 180 |
+
self._training_active = False
|
| 181 |
+
self._inference_active = False
|
| 182 |
+
self._start_time: Optional[float] = None
|
| 183 |
+
self._epoch_start_time: Optional[float] = None
|
| 184 |
+
self._current_epoch = 0
|
| 185 |
+
self._current_batch = 0
|
| 186 |
+
self._global_step = 0
|
| 187 |
+
|
| 188 |
+
# Agregados para época atual
|
| 189 |
+
self._epoch_losses: List[float] = []
|
| 190 |
+
self._epoch_ppls: List[float] = []
|
| 191 |
+
self._epoch_hypothesis_activations = 0
|
| 192 |
+
self._epoch_synergy_history: List[Dict] = []
|
| 193 |
+
|
| 194 |
+
# Inferência
|
| 195 |
+
self.inference_stats: Dict[str, Any] = {
|
| 196 |
+
"total_tokens": 0,
|
| 197 |
+
"total_time_ms": 0.0,
|
| 198 |
+
"medusa_accepted": 0,
|
| 199 |
+
"n_calls": 0,
|
| 200 |
+
"fallback_to_greedy": 0,
|
| 201 |
+
}
|
| 202 |
+
|
| 203 |
+
# Component tracking
|
| 204 |
+
self.component_status: Dict[str, bool] = {
|
| 205 |
+
"ewc": False,
|
| 206 |
+
"medusa": False,
|
| 207 |
+
"cyclic_reasoning": False,
|
| 208 |
+
"vqvae2": False,
|
| 209 |
+
"quantization": False,
|
| 210 |
+
"multimodal_attention": False,
|
| 211 |
+
"context_window": False,
|
| 212 |
+
}
|
| 213 |
+
|
| 214 |
+
def _psutil_available(self) -> bool:
|
| 215 |
+
try:
|
| 216 |
+
import psutil # noqa: F401
|
| 217 |
+
return True
|
| 218 |
+
except ImportError:
|
| 219 |
+
return False
|
| 220 |
+
|
| 221 |
+
# ----------------------------------------------------------------------
|
| 222 |
+
# Training lifecycle
|
| 223 |
+
# ----------------------------------------------------------------------
|
| 224 |
+
|
| 225 |
+
def start_training(self) -> None:
|
| 226 |
+
"""Inicia uma sessão de monitoramento de treino."""
|
| 227 |
+
self._training_active = True
|
| 228 |
+
self._start_time = time.time()
|
| 229 |
+
self._current_epoch = 0
|
| 230 |
+
self._current_batch = 0
|
| 231 |
+
self._global_step = 0
|
| 232 |
+
logger.info("Monitor: sessão de treino iniciada")
|
| 233 |
+
|
| 234 |
+
def start_epoch(self, epoch: int) -> None:
|
| 235 |
+
"""Inicia uma nova época."""
|
| 236 |
+
self._current_epoch = epoch
|
| 237 |
+
self._epoch_start_time = time.time()
|
| 238 |
+
self._epoch_losses = []
|
| 239 |
+
self._epoch_ppls = []
|
| 240 |
+
self._epoch_hypothesis_activations = 0
|
| 241 |
+
self._epoch_synergy_history = []
|
| 242 |
+
|
| 243 |
+
def log_batch(self, metrics: Dict[str, Any]) -> None:
|
| 244 |
+
"""Loga métricas de um batch.
|
| 245 |
+
|
| 246 |
+
Args:
|
| 247 |
+
metrics: dict com chaves como loss, ppl, lr, grad_norm, etc.
|
| 248 |
+
"""
|
| 249 |
+
if not self._training_active:
|
| 250 |
+
return
|
| 251 |
+
|
| 252 |
+
self._global_step += 1
|
| 253 |
+
self._current_batch += 1
|
| 254 |
+
|
| 255 |
+
ts = time.time()
|
| 256 |
+
batch_metrics = BatchMetrics(
|
| 257 |
+
step=self._global_step,
|
| 258 |
+
epoch=self._current_epoch,
|
| 259 |
+
batch_in_epoch=self._current_batch,
|
| 260 |
+
timestamp=ts,
|
| 261 |
+
loss=metrics.get("loss", 0.0),
|
| 262 |
+
loss_main=metrics.get("loss_main", 0.0),
|
| 263 |
+
loss_medusa=metrics.get("loss_medusa", 0.0),
|
| 264 |
+
loss_ewc=metrics.get("loss_ewc", 0.0),
|
| 265 |
+
loss_hallucination=metrics.get("loss_hallucination", 0.0),
|
| 266 |
+
loss_total=metrics.get("loss_total", metrics.get("loss", 0.0)),
|
| 267 |
+
ppl=metrics.get("ppl", 0.0),
|
| 268 |
+
lr_g=metrics.get("lr_g", metrics.get("lr", 0.0)),
|
| 269 |
+
lr_v=metrics.get("lr_v", 0.0),
|
| 270 |
+
grad_norm_g=metrics.get("grad_norm_g", metrics.get("grad_norm", 0.0)),
|
| 271 |
+
grad_norm_v=metrics.get("grad_norm_v", 0.0),
|
| 272 |
+
v_mean=metrics.get("v_mean", 0.0),
|
| 273 |
+
hypothesis_activations=metrics.get("hypothesis_activations", 0),
|
| 274 |
+
elapsed_ms=metrics.get("elapsed_ms", 0.0),
|
| 275 |
+
extra={k: v for k, v in metrics.items()
|
| 276 |
+
if k not in {"loss", "loss_main", "loss_medusa", "loss_ewc",
|
| 277 |
+
"loss_hallucination", "loss_total", "ppl",
|
| 278 |
+
"lr_g", "lr_v", "grad_norm_g", "grad_norm_v",
|
| 279 |
+
"v_mean", "hypothesis_activations", "elapsed_ms",
|
| 280 |
+
"lr", "grad_norm"}},
|
| 281 |
+
)
|
| 282 |
+
self.batch_history.append(batch_metrics)
|
| 283 |
+
|
| 284 |
+
# Acumular para época
|
| 285 |
+
if batch_metrics.loss > 0:
|
| 286 |
+
self._epoch_losses.append(batch_metrics.loss)
|
| 287 |
+
if batch_metrics.ppl > 0:
|
| 288 |
+
self._epoch_ppls.append(batch_metrics.ppl)
|
| 289 |
+
self._epoch_hypothesis_activations += batch_metrics.hypothesis_activations
|
| 290 |
+
if "synergy_history" in metrics:
|
| 291 |
+
self._epoch_synergy_history.extend(metrics["synergy_history"])
|
| 292 |
+
|
| 293 |
+
# Log
|
| 294 |
+
if self._global_step % self.log_every == 0:
|
| 295 |
+
logger.info(
|
| 296 |
+
"Batch %d (epoch %d) | loss=%.4f | ppl=%.2f | lr_g=%.2e | hyp=%d",
|
| 297 |
+
self._global_step, self._current_epoch,
|
| 298 |
+
batch_metrics.loss, batch_metrics.ppl, batch_metrics.lr_g,
|
| 299 |
+
batch_metrics.hypothesis_activations,
|
| 300 |
+
)
|
| 301 |
+
|
| 302 |
+
# Track system metrics occasionally
|
| 303 |
+
if self.track_system and self._global_step % 50 == 0:
|
| 304 |
+
self._track_system()
|
| 305 |
+
|
| 306 |
+
def end_epoch(self, extra: Optional[Dict[str, Any]] = None) -> EpochMetrics:
|
| 307 |
+
"""Finaliza a época atual e retorna métricas agregadas."""
|
| 308 |
+
if self._epoch_start_time is None:
|
| 309 |
+
logger.warning("end_epoch chamado sem start_epoch")
|
| 310 |
+
return EpochMetrics(epoch=self._current_epoch)
|
| 311 |
+
|
| 312 |
+
elapsed = time.time() - self._epoch_start_time
|
| 313 |
+
|
| 314 |
+
# Agregar
|
| 315 |
+
if self._epoch_losses:
|
| 316 |
+
avg_loss = sum(self._epoch_losses) / len(self._epoch_losses)
|
| 317 |
+
best_loss = min(self._epoch_losses)
|
| 318 |
+
worst_loss = max(self._epoch_losses)
|
| 319 |
+
else:
|
| 320 |
+
avg_loss = best_loss = worst_loss = 0.0
|
| 321 |
+
|
| 322 |
+
if self._epoch_ppls:
|
| 323 |
+
avg_ppl = sum(self._epoch_ppls) / len(self._epoch_ppls)
|
| 324 |
+
else:
|
| 325 |
+
avg_ppl = 0.0
|
| 326 |
+
|
| 327 |
+
epoch_metrics = EpochMetrics(
|
| 328 |
+
epoch=self._current_epoch,
|
| 329 |
+
avg_loss=avg_loss,
|
| 330 |
+
avg_ppl=avg_ppl,
|
| 331 |
+
n_batches=len(self._epoch_losses),
|
| 332 |
+
n_hypothesis_activations=self._epoch_hypothesis_activations,
|
| 333 |
+
elapsed_s=elapsed,
|
| 334 |
+
best_loss=best_loss,
|
| 335 |
+
worst_loss=worst_loss,
|
| 336 |
+
synergy_history=list(self._epoch_synergy_history),
|
| 337 |
+
extra=extra or {},
|
| 338 |
+
)
|
| 339 |
+
self.epoch_history.append(epoch_metrics)
|
| 340 |
+
|
| 341 |
+
logger.info(
|
| 342 |
+
"Época %d concluída | avg_loss=%.4f | avg_ppl=%.2f | hyp_act=%d | batches=%d | %.1fs",
|
| 343 |
+
self._current_epoch, avg_loss, avg_ppl,
|
| 344 |
+
self._epoch_hypothesis_activations, len(self._epoch_losses), elapsed,
|
| 345 |
+
)
|
| 346 |
+
return epoch_metrics
|
| 347 |
+
|
| 348 |
+
def end_training(self) -> Dict[str, Any]:
|
| 349 |
+
"""Finaliza a sessão de treino e retorna resumo."""
|
| 350 |
+
if self._start_time is None:
|
| 351 |
+
return {}
|
| 352 |
+
total_time = time.time() - self._start_time
|
| 353 |
+
self._training_active = False
|
| 354 |
+
|
| 355 |
+
summary = {
|
| 356 |
+
"total_time_s": total_time,
|
| 357 |
+
"total_epochs": len(self.epoch_history),
|
| 358 |
+
"total_batches": self._global_step,
|
| 359 |
+
"final_loss": self.epoch_history[-1].avg_loss if self.epoch_history else 0.0,
|
| 360 |
+
"final_ppl": self.epoch_history[-1].avg_ppl if self.epoch_history else 0.0,
|
| 361 |
+
"best_loss": min((e.avg_loss for e in self.epoch_history), default=0.0),
|
| 362 |
+
"best_ppl": min((e.avg_ppl for e in self.epoch_history), default=0.0),
|
| 363 |
+
"total_hypothesis_activations": sum(e.n_hypothesis_activations for e in self.epoch_history),
|
| 364 |
+
}
|
| 365 |
+
logger.info("Monitor: sessão de treino finalizada — %s", summary)
|
| 366 |
+
return summary
|
| 367 |
+
|
| 368 |
+
# ----------------------------------------------------------------------
|
| 369 |
+
# Inference tracking
|
| 370 |
+
# ----------------------------------------------------------------------
|
| 371 |
+
|
| 372 |
+
def log_inference(
|
| 373 |
+
self,
|
| 374 |
+
n_tokens: int,
|
| 375 |
+
elapsed_ms: float,
|
| 376 |
+
medusa_accepted: int = 0,
|
| 377 |
+
fallback_to_greedy: int = 0,
|
| 378 |
+
) -> None:
|
| 379 |
+
"""Loga uma chamada de inferência."""
|
| 380 |
+
self.inference_stats["total_tokens"] += n_tokens
|
| 381 |
+
self.inference_stats["total_time_ms"] += elapsed_ms
|
| 382 |
+
self.inference_stats["medusa_accepted"] += medusa_accepted
|
| 383 |
+
self.inference_stats["fallback_to_greedy"] += fallback_to_greedy
|
| 384 |
+
self.inference_stats["n_calls"] += 1
|
| 385 |
+
|
| 386 |
+
def get_inference_throughput(self) -> Dict[str, float]:
|
| 387 |
+
"""Retorna throughput de inferência."""
|
| 388 |
+
s = self.inference_stats
|
| 389 |
+
if s["total_time_ms"] == 0:
|
| 390 |
+
return {"tokens_per_s": 0.0, "ms_per_token": 0.0, "medusa_accept_rate": 0.0}
|
| 391 |
+
total_s = s["total_time_ms"] / 1000.0
|
| 392 |
+
tps = s["total_tokens"] / total_s if total_s > 0 else 0.0
|
| 393 |
+
mspt = s["total_time_ms"] / s["total_tokens"] if s["total_tokens"] > 0 else 0.0
|
| 394 |
+
accept_rate = s["medusa_accepted"] / max(1, s["total_tokens"])
|
| 395 |
+
return {
|
| 396 |
+
"tokens_per_s": tps,
|
| 397 |
+
"ms_per_token": mspt,
|
| 398 |
+
"medusa_accept_rate": accept_rate,
|
| 399 |
+
"n_calls": s["n_calls"],
|
| 400 |
+
"total_tokens": s["total_tokens"],
|
| 401 |
+
}
|
| 402 |
+
|
| 403 |
+
# ----------------------------------------------------------------------
|
| 404 |
+
# Evolution tracking (componentes específicos)
|
| 405 |
+
# ----------------------------------------------------------------------
|
| 406 |
+
|
| 407 |
+
def log_cyclic_reasoning(self, stats: Dict[str, Any]) -> None:
|
| 408 |
+
"""Loga estatísticas do CyclicReasoning."""
|
| 409 |
+
stats["step"] = self._global_step
|
| 410 |
+
stats["timestamp"] = time.time()
|
| 411 |
+
self.evolution.cyclic_reasoning_stats.append(stats)
|
| 412 |
+
|
| 413 |
+
def log_vqvae2_usage(self, stats: Dict[str, Any]) -> None:
|
| 414 |
+
"""Loga uso do codebook VQ-VAE-2."""
|
| 415 |
+
stats["step"] = self._global_step
|
| 416 |
+
stats["timestamp"] = time.time()
|
| 417 |
+
self.evolution.vqvae2_codebook_usage.append(stats)
|
| 418 |
+
|
| 419 |
+
def log_quantization(self, stats: Dict[str, Any]) -> None:
|
| 420 |
+
"""Loga estatísticas de quantização W8A8."""
|
| 421 |
+
stats["step"] = self._global_step
|
| 422 |
+
stats["timestamp"] = time.time()
|
| 423 |
+
self.evolution.quantization_stats.append(stats)
|
| 424 |
+
|
| 425 |
+
def log_ewc_penalty(self, penalty: float, num_tasks: int) -> None:
|
| 426 |
+
"""Loga evolução do penalty EWC."""
|
| 427 |
+
self.evolution.ewc_penalty_evolution.append({
|
| 428 |
+
"step": self._global_step,
|
| 429 |
+
"timestamp": time.time(),
|
| 430 |
+
"penalty": penalty,
|
| 431 |
+
"num_tasks": num_tasks,
|
| 432 |
+
})
|
| 433 |
+
|
| 434 |
+
def log_medusa_accept(self, accepted: int, total: int) -> None:
|
| 435 |
+
"""Loga taxa de aceitação Medusa."""
|
| 436 |
+
rate = accepted / max(1, total)
|
| 437 |
+
self.evolution.medusa_accept_rate.append({
|
| 438 |
+
"step": self._global_step,
|
| 439 |
+
"timestamp": time.time(),
|
| 440 |
+
"accepted": accepted,
|
| 441 |
+
"total": total,
|
| 442 |
+
"rate": rate,
|
| 443 |
+
})
|
| 444 |
+
|
| 445 |
+
def register_component(self, name: str, active: bool = True) -> None:
|
| 446 |
+
"""Registra que um componente está ativo."""
|
| 447 |
+
if name in self.component_status:
|
| 448 |
+
self.component_status[name] = active
|
| 449 |
+
|
| 450 |
+
# ----------------------------------------------------------------------
|
| 451 |
+
# System tracking
|
| 452 |
+
# ----------------------------------------------------------------------
|
| 453 |
+
|
| 454 |
+
def _track_system(self) -> None:
|
| 455 |
+
"""Rastreia métricas de sistema."""
|
| 456 |
+
if not self.track_system:
|
| 457 |
+
return
|
| 458 |
+
try:
|
| 459 |
+
import psutil
|
| 460 |
+
cpu = psutil.cpu_percent(interval=None)
|
| 461 |
+
mem = psutil.virtual_memory()
|
| 462 |
+
sm = SystemMetrics(
|
| 463 |
+
timestamp=time.time(),
|
| 464 |
+
cpu_percent=cpu,
|
| 465 |
+
memory_mb=mem.used / (1024 * 1024),
|
| 466 |
+
n_active_threads=psutil.cpu_count() or 0,
|
| 467 |
+
)
|
| 468 |
+
# GPU memory (se disponível)
|
| 469 |
+
if torch.cuda.is_available():
|
| 470 |
+
try:
|
| 471 |
+
gpu_mem = torch.cuda.memory_allocated() / (1024 * 1024)
|
| 472 |
+
sm.gpu_memory_mb = gpu_mem
|
| 473 |
+
except Exception:
|
| 474 |
+
pass
|
| 475 |
+
self.system_history.append(sm)
|
| 476 |
+
except Exception as e:
|
| 477 |
+
logger.debug("System tracking falhou: %s", e)
|
| 478 |
+
|
| 479 |
+
# ----------------------------------------------------------------------
|
| 480 |
+
# Exportação
|
| 481 |
+
# ----------------------------------------------------------------------
|
| 482 |
+
|
| 483 |
+
def export_report(self, output_path: Optional[Union[str, Path]] = None) -> Path:
|
| 484 |
+
"""Exporta relatório completo em JSON.
|
| 485 |
+
|
| 486 |
+
Args:
|
| 487 |
+
output_path: caminho do arquivo (default: output_dir/monitor_report.json)
|
| 488 |
+
|
| 489 |
+
Returns:
|
| 490 |
+
Path do arquivo salvo
|
| 491 |
+
"""
|
| 492 |
+
if output_path is None:
|
| 493 |
+
if self.output_dir is None:
|
| 494 |
+
raise ValueError("output_dir ou output_path deve ser fornecido")
|
| 495 |
+
output_path = self.output_dir / "monitor_report.json"
|
| 496 |
+
output_path = Path(output_path)
|
| 497 |
+
output_path.parent.mkdir(parents=True, exist_ok=True)
|
| 498 |
+
|
| 499 |
+
# Construir relatório
|
| 500 |
+
report = {
|
| 501 |
+
"metadata": {
|
| 502 |
+
"generated_at": time.time(),
|
| 503 |
+
"training_active": self._training_active,
|
| 504 |
+
"global_step": self._global_step,
|
| 505 |
+
"current_epoch": self._current_epoch,
|
| 506 |
+
},
|
| 507 |
+
"components": dict(self.component_status),
|
| 508 |
+
"training_summary": self._get_training_summary(),
|
| 509 |
+
"epoch_history": [asdict(e) for e in self.epoch_history],
|
| 510 |
+
"inference_stats": {
|
| 511 |
+
**self.inference_stats,
|
| 512 |
+
"throughput": self.get_inference_throughput(),
|
| 513 |
+
},
|
| 514 |
+
"evolution": {
|
| 515 |
+
"cyclic_reasoning": self.evolution.cyclic_reasoning_stats[-50:],
|
| 516 |
+
"vqvae2_codebook_usage": self.evolution.vqvae2_codebook_usage[-50:],
|
| 517 |
+
"quantization_stats": self.evolution.quantization_stats[-10:],
|
| 518 |
+
"ewc_penalty_evolution": self.evolution.ewc_penalty_evolution[-50:],
|
| 519 |
+
"medusa_accept_rate": self.evolution.medusa_accept_rate[-50:],
|
| 520 |
+
},
|
| 521 |
+
"system_metrics": [asdict(s) for s in list(self.system_history)[-50:]],
|
| 522 |
+
}
|
| 523 |
+
|
| 524 |
+
with open(output_path, "w", encoding="utf-8") as f:
|
| 525 |
+
json.dump(report, f, indent=2, default=str, ensure_ascii=False)
|
| 526 |
+
logger.info("Monitor report salvo em: %s", output_path)
|
| 527 |
+
return output_path
|
| 528 |
+
|
| 529 |
+
def export_csv(self, output_path: Optional[Union[str, Path]] = None) -> Path:
|
| 530 |
+
"""Exporta métricas de batch em CSV."""
|
| 531 |
+
if output_path is None:
|
| 532 |
+
if self.output_dir is None:
|
| 533 |
+
raise ValueError("output_dir ou output_path deve ser fornecido")
|
| 534 |
+
output_path = self.output_dir / "batch_metrics.csv"
|
| 535 |
+
output_path = Path(output_path)
|
| 536 |
+
output_path.parent.mkdir(parents=True, exist_ok=True)
|
| 537 |
+
|
| 538 |
+
if not self.batch_history:
|
| 539 |
+
logger.warning("Sem batch history para exportar")
|
| 540 |
+
return output_path
|
| 541 |
+
|
| 542 |
+
# Pegar chaves do primeiro item + extras
|
| 543 |
+
first = self.batch_history[0]
|
| 544 |
+
fieldnames = ["step", "epoch", "batch_in_epoch", "timestamp",
|
| 545 |
+
"loss", "loss_main", "loss_medusa", "loss_ewc",
|
| 546 |
+
"loss_hallucination", "loss_total", "ppl",
|
| 547 |
+
"lr_g", "lr_v", "grad_norm_g", "grad_norm_v",
|
| 548 |
+
"v_mean", "hypothesis_activations", "elapsed_ms"]
|
| 549 |
+
|
| 550 |
+
with open(output_path, "w", newline="", encoding="utf-8") as f:
|
| 551 |
+
writer = csv.DictWriter(f, fieldnames=fieldnames, extrasaction="ignore")
|
| 552 |
+
writer.writeheader()
|
| 553 |
+
for m in self.batch_history:
|
| 554 |
+
writer.writerow(asdict(m))
|
| 555 |
+
|
| 556 |
+
logger.info("CSV metrics salvo em: %s (%d rows)", output_path, len(self.batch_history))
|
| 557 |
+
return output_path
|
| 558 |
+
|
| 559 |
+
def export_markdown_summary(self, output_path: Optional[Union[str, Path]] = None) -> Path:
|
| 560 |
+
"""Exporta resumo em Markdown."""
|
| 561 |
+
if output_path is None:
|
| 562 |
+
if self.output_dir is None:
|
| 563 |
+
raise ValueError("output_dir ou output_path deve ser fornecido")
|
| 564 |
+
output_path = self.output_dir / "monitor_summary.md"
|
| 565 |
+
output_path = Path(output_path)
|
| 566 |
+
output_path.parent.mkdir(parents=True, exist_ok=True)
|
| 567 |
+
|
| 568 |
+
summary = self._get_training_summary()
|
| 569 |
+
throughput = self.get_inference_throughput()
|
| 570 |
+
|
| 571 |
+
lines = [
|
| 572 |
+
"# Monitor Report — CNN-BiGRU",
|
| 573 |
+
"",
|
| 574 |
+
f"**Generated at:** {time.strftime('%Y-%m-%d %H:%M:%S')}",
|
| 575 |
+
"",
|
| 576 |
+
"## Components Status",
|
| 577 |
+
"",
|
| 578 |
+
]
|
| 579 |
+
for comp, active in self.component_status.items():
|
| 580 |
+
icon = "[x]" if active else "[ ]"
|
| 581 |
+
lines.append(f"- {icon} {comp}")
|
| 582 |
+
|
| 583 |
+
lines.extend([
|
| 584 |
+
"",
|
| 585 |
+
"## Training Summary",
|
| 586 |
+
"",
|
| 587 |
+
f"- Total time: {summary.get('total_time_s', 0):.1f}s",
|
| 588 |
+
f"- Total epochs: {summary.get('total_epochs', 0)}",
|
| 589 |
+
f"- Total batches: {summary.get('total_batches', 0)}",
|
| 590 |
+
f"- Final loss: {summary.get('final_loss', 0):.4f}",
|
| 591 |
+
f"- Final PPL: {summary.get('final_ppl', 0):.2f}",
|
| 592 |
+
f"- Best loss: {summary.get('best_loss', 0):.4f}",
|
| 593 |
+
f"- Best PPL: {summary.get('best_ppl', 0):.2f}",
|
| 594 |
+
f"- Total hypothesis activations: {summary.get('total_hypothesis_activations', 0)}",
|
| 595 |
+
"",
|
| 596 |
+
"## Inference Throughput",
|
| 597 |
+
"",
|
| 598 |
+
f"- Tokens/s: {throughput.get('tokens_per_s', 0):.1f}",
|
| 599 |
+
f"- ms/token: {throughput.get('ms_per_token', 0):.2f}",
|
| 600 |
+
f"- Medusa accept rate: {throughput.get('medusa_accept_rate', 0):.2%}",
|
| 601 |
+
f"- Total tokens generated: {throughput.get('total_tokens', 0)}",
|
| 602 |
+
f"- Total inference calls: {throughput.get('n_calls', 0)}",
|
| 603 |
+
"",
|
| 604 |
+
"## Evolution",
|
| 605 |
+
"",
|
| 606 |
+
f"- CyclicReasoning events: {len(self.evolution.cyclic_reasoning_stats)}",
|
| 607 |
+
f"- VQ-VAE-2 usage events: {len(self.evolution.vqvae2_codebook_usage)}",
|
| 608 |
+
f"- Quantization events: {len(self.evolution.quantization_stats)}",
|
| 609 |
+
f"- EWC penalty events: {len(self.evolution.ewc_penalty_evolution)}",
|
| 610 |
+
f"- Medusa accept events: {len(self.evolution.medusa_accept_rate)}",
|
| 611 |
+
"",
|
| 612 |
+
"## Epoch History",
|
| 613 |
+
"",
|
| 614 |
+
"| Epoch | Avg Loss | Avg PPL | Batches | Hyp Act | Time(s) |",
|
| 615 |
+
"|-------|----------|---------|---------|---------|---------|",
|
| 616 |
+
])
|
| 617 |
+
for e in self.epoch_history:
|
| 618 |
+
lines.append(
|
| 619 |
+
f"| {e.epoch} | {e.avg_loss:.4f} | {e.avg_ppl:.2f} | "
|
| 620 |
+
f"{e.n_batches} | {e.n_hypothesis_activations} | {e.elapsed_s:.1f} |"
|
| 621 |
+
)
|
| 622 |
+
|
| 623 |
+
with open(output_path, "w", encoding="utf-8") as f:
|
| 624 |
+
f.write("\n".join(lines))
|
| 625 |
+
logger.info("Markdown summary salvo em: %s", output_path)
|
| 626 |
+
return output_path
|
| 627 |
+
|
| 628 |
+
def _get_training_summary(self) -> Dict[str, Any]:
|
| 629 |
+
"""Constrói sumário de treino."""
|
| 630 |
+
if not self.epoch_history:
|
| 631 |
+
return {"total_time_s": 0, "total_epochs": 0, "total_batches": self._global_step}
|
| 632 |
+
return {
|
| 633 |
+
"total_time_s": time.time() - (self._start_time or time.time()),
|
| 634 |
+
"total_epochs": len(self.epoch_history),
|
| 635 |
+
"total_batches": self._global_step,
|
| 636 |
+
"final_loss": self.epoch_history[-1].avg_loss,
|
| 637 |
+
"final_ppl": self.epoch_history[-1].avg_ppl,
|
| 638 |
+
"best_loss": min((e.avg_loss for e in self.epoch_history), default=0.0),
|
| 639 |
+
"best_ppl": min((e.avg_ppl for e in self.epoch_history), default=0.0),
|
| 640 |
+
"total_hypothesis_activations": sum(
|
| 641 |
+
e.n_hypothesis_activations for e in self.epoch_history
|
| 642 |
+
),
|
| 643 |
+
}
|
| 644 |
+
|
| 645 |
+
# ----------------------------------------------------------------------
|
| 646 |
+
# Reset
|
| 647 |
+
# ----------------------------------------------------------------------
|
| 648 |
+
|
| 649 |
+
def reset(self) -> None:
|
| 650 |
+
"""Limpa todo o histórico."""
|
| 651 |
+
self.batch_history.clear()
|
| 652 |
+
self.epoch_history.clear()
|
| 653 |
+
self.evolution = EvolutionMetrics()
|
| 654 |
+
self.system_history.clear()
|
| 655 |
+
self._global_step = 0
|
| 656 |
+
self._current_epoch = 0
|
| 657 |
+
self._current_batch = 0
|
| 658 |
+
self._start_time = None
|
| 659 |
+
self._epoch_start_time = None
|
| 660 |
+
self.inference_stats = {k: 0 if isinstance(v, int) else 0.0
|
| 661 |
+
for k, v in self.inference_stats.items()}
|
| 662 |
+
|
| 663 |
+
|
| 664 |
+
# ============================================================================
|
| 665 |
+
# Singleton instance (opcional)
|
| 666 |
+
# ============================================================================
|
| 667 |
+
|
| 668 |
+
_global_monitor: Optional[Monitor] = None
|
| 669 |
+
|
| 670 |
+
|
| 671 |
+
def get_monitor(output_dir: Optional[Union[str, Path]] = None) -> Monitor:
|
| 672 |
+
"""Retorna a instância global do monitor (singleton)."""
|
| 673 |
+
global _global_monitor
|
| 674 |
+
if _global_monitor is None:
|
| 675 |
+
_global_monitor = Monitor(output_dir=output_dir)
|
| 676 |
+
return _global_monitor
|
| 677 |
+
|
| 678 |
+
|
| 679 |
+
# ============================================================================
|
| 680 |
+
# Self-test
|
| 681 |
+
# ============================================================================
|
| 682 |
+
|
| 683 |
+
def _self_test():
|
| 684 |
+
"""Teste rápido do monitor."""
|
| 685 |
+
import tempfile
|
| 686 |
+
with tempfile.TemporaryDirectory() as tmpdir:
|
| 687 |
+
mon = Monitor(output_dir=tmpdir, log_every=1, track_system=True)
|
| 688 |
+
mon.start_training()
|
| 689 |
+
mon.start_epoch(0)
|
| 690 |
+
for i in range(5):
|
| 691 |
+
mon.log_batch({
|
| 692 |
+
"loss": 5.0 - i * 0.5,
|
| 693 |
+
"ppl": 100.0 - i * 10,
|
| 694 |
+
"lr_g": 1e-3,
|
| 695 |
+
"grad_norm": 0.5 + i * 0.1,
|
| 696 |
+
"hypothesis_activations": i,
|
| 697 |
+
"elapsed_ms": 10 + i,
|
| 698 |
+
})
|
| 699 |
+
mon.end_epoch()
|
| 700 |
+
mon.end_training()
|
| 701 |
+
|
| 702 |
+
mon.log_inference(n_tokens=50, elapsed_ms=500, medusa_accepted=10)
|
| 703 |
+
mon.log_cyclic_reasoning({"n_cycles": 3, "converged": True})
|
| 704 |
+
mon.log_vqvae2_usage({"top_usage": 0.5, "bottom_usage": 0.7})
|
| 705 |
+
mon.log_quantization({"reduction_pct": 60.0})
|
| 706 |
+
mon.log_ewc_penalty(penalty=0.001, num_tasks=1)
|
| 707 |
+
mon.log_medusa_accept(accepted=10, total=50)
|
| 708 |
+
|
| 709 |
+
report = mon.export_report()
|
| 710 |
+
csv_path = mon.export_csv()
|
| 711 |
+
md_path = mon.export_markdown_summary()
|
| 712 |
+
print(f"Report: {report}")
|
| 713 |
+
print(f"CSV: {csv_path}")
|
| 714 |
+
print(f"MD: {md_path}")
|
| 715 |
+
print(f"Throughput: {mon.get_inference_throughput()}")
|
| 716 |
+
|
| 717 |
+
|
| 718 |
+
if __name__ == "__main__":
|
| 719 |
+
_self_test()
|
| 720 |
+
|
| 721 |
+
|
| 722 |
+
__all__ = [
|
| 723 |
+
"BatchMetrics",
|
| 724 |
+
"EpochMetrics",
|
| 725 |
+
"EvolutionMetrics",
|
| 726 |
+
"SystemMetrics",
|
| 727 |
+
"Monitor",
|
| 728 |
+
"get_monitor",
|
| 729 |
+
]
|
utils/quantization.py
ADDED
|
@@ -0,0 +1,585 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""quantization.py — W8A8 Quantization via SmoothQuant para CNN-BiGRU.
|
| 2 |
+
|
| 3 |
+
Implementa quantização weight-only-8-bit + activation-8-bit (W8A8) usando
|
| 4 |
+
a técnica SmoothQuant (Xiao et al., 2023) para migrar a variância das
|
| 5 |
+
ativações para os pesos, reduzindo a perda de precisão.
|
| 6 |
+
|
| 7 |
+
==============================================================================
|
| 8 |
+
ANÁLISE MATEMÁTICA E LÓGICA — SmoothQuant + W8A8
|
| 9 |
+
==============================================================================
|
| 10 |
+
|
| 11 |
+
PROBLEMA:
|
| 12 |
+
Em modelos LLM, as ativações têm outliers em alguns canais que tornam
|
| 13 |
+
a quantização INT8 difícil. Se quantizarmos diretamente, os outliers
|
| 14 |
+
saturam os outros canais, levando a grandes erros.
|
| 15 |
+
|
| 16 |
+
SOLUÇÃO SmoothQuant:
|
| 17 |
+
Seja Y = X * W, onde X ∈ R^{B×T×d_in} e W ∈ R^{d_in×d_out}.
|
| 18 |
+
|
| 19 |
+
1. Para cada canal de entrada i, computa o máximo absoluto:
|
| 20 |
+
s_i = max|X_i| / max|W_i| (em batch, suavizado por alpha)
|
| 21 |
+
|
| 22 |
+
Ou mais precisamente:
|
| 23 |
+
s_i = (max|X_i|^alpha) / (max|W_i|^(1-alpha))
|
| 24 |
+
com alpha ∈ [0, 1] tipicamente 0.5.
|
| 25 |
+
|
| 26 |
+
2. Migra a variância: dividir X por s e multiplicar W por s:
|
| 27 |
+
X' = X / s (ativações suavizadas — outliers reduzidos)
|
| 28 |
+
W' = W * s (pesos absorvem a escala)
|
| 29 |
+
|
| 30 |
+
Como Y = X * W = (X/s) * (s*W) = X' * W', a operação matricial
|
| 31 |
+
é matematicamente equivalente.
|
| 32 |
+
|
| 33 |
+
3. Quantiza ambos para INT8 com escala por tensor ou por canal:
|
| 34 |
+
X_q = round(X' / scale_x) * scale_x
|
| 35 |
+
W_q = round(W' / scale_w) * scale_w
|
| 36 |
+
|
| 37 |
+
4. Y ≈ X_q * W_q (com erro de quantização reduzido)
|
| 38 |
+
|
| 39 |
+
W8A8:
|
| 40 |
+
- W: pesos em INT8 (8-bit weights)
|
| 41 |
+
- A: ativações em INT8 (8-bit activations)
|
| 42 |
+
- Redução de memória: ~4x (FP32 -> INT8)
|
| 43 |
+
- Speedup: 2-4x em hardware com suporte INT8 (AMX, AVX512-VNNI)
|
| 44 |
+
|
| 45 |
+
INTEGRAÇÃO COM CNN-BiGRU:
|
| 46 |
+
- Aplicado após o treino (post-training quantization)
|
| 47 |
+
- Aplicável a: Linear (atenção, FFN, classificador), Conv1d
|
| 48 |
+
- Não aplicado a: Embedding (mantém FP32 para precisão)
|
| 49 |
+
- Em CPU sem AMX, ainda economiza memória (sem speedup significativo)
|
| 50 |
+
|
| 51 |
+
==============================================================================
|
| 52 |
+
USO
|
| 53 |
+
==============================================================================
|
| 54 |
+
|
| 55 |
+
from cnn_bigru.utils.quantization import (
|
| 56 |
+
SmoothQuantizer, W8A8Config, quantize_model_w8a8
|
| 57 |
+
)
|
| 58 |
+
|
| 59 |
+
# Após treino:
|
| 60 |
+
quantizer = SmoothQuantizer(W8A8Config(alpha=0.5, n_calibration_batches=5))
|
| 61 |
+
quantized_model = quantizer.quantize(model, calibration_dataloader)
|
| 62 |
+
|
| 63 |
+
Autor: CNN-BiGRU Project
|
| 64 |
+
"""
|
| 65 |
+
from __future__ import annotations
|
| 66 |
+
|
| 67 |
+
import logging
|
| 68 |
+
import math
|
| 69 |
+
from dataclasses import dataclass, field
|
| 70 |
+
from typing import Dict, List, Optional, Tuple, Union
|
| 71 |
+
|
| 72 |
+
import torch
|
| 73 |
+
import torch.nn as nn
|
| 74 |
+
import torch.nn.functional as F
|
| 75 |
+
|
| 76 |
+
logger = logging.getLogger(__name__)
|
| 77 |
+
|
| 78 |
+
|
| 79 |
+
# ============================================================================
|
| 80 |
+
# Configuração
|
| 81 |
+
# ============================================================================
|
| 82 |
+
|
| 83 |
+
@dataclass
|
| 84 |
+
class W8A8Config:
|
| 85 |
+
"""Configuração da quantização W8A8 com SmoothQuant."""
|
| 86 |
+
# Alpha de suavização (0=só pesos, 1=só ativações, 0.5=balanceado)
|
| 87 |
+
alpha: float = 0.5
|
| 88 |
+
# Número de batches para calibração
|
| 89 |
+
n_calibration_batches: int = 5
|
| 90 |
+
# Tipo de escala: "per_tensor" ou "per_channel"
|
| 91 |
+
scale_type: str = "per_channel"
|
| 92 |
+
# Manter权重 em FP32 para quantização dinâmica (default: False = INT8 estático)
|
| 93 |
+
dynamic_weight: bool = False
|
| 94 |
+
# Quantizar embeddings (default: False)
|
| 95 |
+
quantize_embeddings: bool = False
|
| 96 |
+
# Quantizar Conv1d (default: True)
|
| 97 |
+
quantize_conv1d: bool = True
|
| 98 |
+
# Quantizar LayerNorm (default: False — sensível)
|
| 99 |
+
quantize_layernorm: bool = False
|
| 100 |
+
# Clip threshold para outliers (em desvios-padrão; None = sem clip)
|
| 101 |
+
outlier_clip_std: Optional[float] = 4.0
|
| 102 |
+
# Device para calibração
|
| 103 |
+
device: str = "cpu"
|
| 104 |
+
# Simular quantização (mantém FP32 mas aplica ruído de quantização)
|
| 105 |
+
simulate: bool = False
|
| 106 |
+
|
| 107 |
+
|
| 108 |
+
# ============================================================================
|
| 109 |
+
# SmoothQuant Calibrator
|
| 110 |
+
# ============================================================================
|
| 111 |
+
|
| 112 |
+
class SmoothQuantCalibrator:
|
| 113 |
+
"""Coleta estatísticas de ativações e pesos para SmoothQuant.
|
| 114 |
+
|
| 115 |
+
Per-corre o modelo com dados de calibração e registra max|X| e max|W|
|
| 116 |
+
para cada camada Linear/Conv1d.
|
| 117 |
+
"""
|
| 118 |
+
|
| 119 |
+
def __init__(self, config: W8A8Config):
|
| 120 |
+
self.config = config
|
| 121 |
+
self.stats: Dict[str, Dict[str, torch.Tensor]] = {}
|
| 122 |
+
|
| 123 |
+
def _hook_factory(self, name: str):
|
| 124 |
+
"""Cria um hook forward para coletar max|X|."""
|
| 125 |
+
def hook(module, input, output):
|
| 126 |
+
# input é uma tupla; input[0] é o tensor principal
|
| 127 |
+
if not isinstance(input, tuple) or len(input) == 0:
|
| 128 |
+
return
|
| 129 |
+
x = input[0]
|
| 130 |
+
if not isinstance(x, torch.Tensor):
|
| 131 |
+
return
|
| 132 |
+
# Reduzir para [d_in] (assumindo última dim = features)
|
| 133 |
+
with torch.no_grad():
|
| 134 |
+
if x.dim() >= 2:
|
| 135 |
+
# max sobre todas as dims exceto a última
|
| 136 |
+
x_flat = x.reshape(-1, x.size(-1))
|
| 137 |
+
max_x = x_flat.abs().max(dim=0).values # [d_in]
|
| 138 |
+
else:
|
| 139 |
+
max_x = x.abs() # [d_in]
|
| 140 |
+
# Acumular max (não é média — pegamos o máximo global)
|
| 141 |
+
if name not in self.stats:
|
| 142 |
+
self.stats[name] = {"max_x": max_x.clone()}
|
| 143 |
+
else:
|
| 144 |
+
self.stats[name]["max_x"] = torch.maximum(
|
| 145 |
+
self.stats[name]["max_x"], max_x
|
| 146 |
+
)
|
| 147 |
+
return hook
|
| 148 |
+
|
| 149 |
+
def calibrate(
|
| 150 |
+
self,
|
| 151 |
+
model: nn.Module,
|
| 152 |
+
dataloader,
|
| 153 |
+
n_batches: Optional[int] = None,
|
| 154 |
+
) -> Dict[str, Dict[str, torch.Tensor]]:
|
| 155 |
+
"""Coleta estatísticas via forward hooks.
|
| 156 |
+
|
| 157 |
+
Args:
|
| 158 |
+
model: modelo a calibrar
|
| 159 |
+
dataloader: iterador de batches
|
| 160 |
+
n_batches: número de batches (default: config.n_calibration_batches)
|
| 161 |
+
|
| 162 |
+
Returns:
|
| 163 |
+
dict {layer_name: {"max_x": [d_in], "max_w": [d_out]}}
|
| 164 |
+
"""
|
| 165 |
+
n_batches = n_batches or self.config.n_calibration_batches
|
| 166 |
+
device = self.config.device
|
| 167 |
+
|
| 168 |
+
# Registrar hooks em todas as camadas Linear e Conv1d
|
| 169 |
+
hooks = []
|
| 170 |
+
layer_modules = []
|
| 171 |
+
for name, module in model.named_modules():
|
| 172 |
+
if isinstance(module, (nn.Linear, nn.Conv1d)):
|
| 173 |
+
# Skip embeddings
|
| 174 |
+
if "embedding" in name.lower() and not self.config.quantize_embeddings:
|
| 175 |
+
continue
|
| 176 |
+
if isinstance(module, nn.Conv1d) and not self.config.quantize_conv1d:
|
| 177 |
+
continue
|
| 178 |
+
hook = module.register_forward_hook(self._hook_factory(name))
|
| 179 |
+
hooks.append(hook)
|
| 180 |
+
layer_modules.append((name, module))
|
| 181 |
+
|
| 182 |
+
# Forward pass em modo eval (sem gradientes)
|
| 183 |
+
model.eval()
|
| 184 |
+
model.to(device)
|
| 185 |
+
was_training = model.training
|
| 186 |
+
model.eval()
|
| 187 |
+
|
| 188 |
+
try:
|
| 189 |
+
with torch.no_grad():
|
| 190 |
+
for i, batch in enumerate(dataloader):
|
| 191 |
+
if i >= n_batches:
|
| 192 |
+
break
|
| 193 |
+
try:
|
| 194 |
+
# Tentar diferentes formatos de batch
|
| 195 |
+
if isinstance(batch, dict):
|
| 196 |
+
input_ids_a = batch.get("input_ids_a")
|
| 197 |
+
input_ids_b = batch.get("input_ids_b")
|
| 198 |
+
images = batch.get("images")
|
| 199 |
+
audios = batch.get("audios")
|
| 200 |
+
if input_ids_a is not None and input_ids_b is not None:
|
| 201 |
+
# Tentar forward do modelo multimodal
|
| 202 |
+
try:
|
| 203 |
+
model(
|
| 204 |
+
input_ids_a.to(device),
|
| 205 |
+
input_ids_b.to(device),
|
| 206 |
+
images=images.to(device) if images is not None else None,
|
| 207 |
+
audios=audios.to(device) if audios is not None else None,
|
| 208 |
+
mode="classify",
|
| 209 |
+
)
|
| 210 |
+
except Exception:
|
| 211 |
+
# Fallback: forward sem imagens/áudios
|
| 212 |
+
model(input_ids_a.to(device), input_ids_b.to(device))
|
| 213 |
+
elif isinstance(batch, (list, tuple)) and len(batch) >= 2:
|
| 214 |
+
model(batch[0].to(device), batch[1].to(device))
|
| 215 |
+
else:
|
| 216 |
+
logger.debug(f"Batch format não reconhecido: {type(batch)}")
|
| 217 |
+
except Exception as e:
|
| 218 |
+
logger.debug(f"Calibração batch {i} falhou: {e}")
|
| 219 |
+
continue
|
| 220 |
+
finally:
|
| 221 |
+
# Remover hooks
|
| 222 |
+
for hook in hooks:
|
| 223 |
+
hook.remove()
|
| 224 |
+
if was_training:
|
| 225 |
+
model.train()
|
| 226 |
+
|
| 227 |
+
# Coletar max|W| para cada camada
|
| 228 |
+
for name, module in layer_modules:
|
| 229 |
+
if name not in self.stats:
|
| 230 |
+
continue
|
| 231 |
+
w = module.weight.data
|
| 232 |
+
if isinstance(module, nn.Linear):
|
| 233 |
+
# w: [d_out, d_in] -> max sobre d_out
|
| 234 |
+
max_w = w.abs().amax(dim=0) # [d_in]
|
| 235 |
+
elif isinstance(module, nn.Conv1d):
|
| 236 |
+
# w: [out_channels, in_channels//groups, kernel_size]
|
| 237 |
+
# max sobre out_channels e kernel_size
|
| 238 |
+
max_w = w.abs().amax(dim=(0, 2)) # [in_channels]
|
| 239 |
+
else:
|
| 240 |
+
continue
|
| 241 |
+
self.stats[name]["max_w"] = max_w.clone()
|
| 242 |
+
|
| 243 |
+
logger.info(
|
| 244 |
+
"SmoothQuant calibrado: %d camadas, %d batches",
|
| 245 |
+
len(self.stats), n_batches,
|
| 246 |
+
)
|
| 247 |
+
return self.stats
|
| 248 |
+
|
| 249 |
+
|
| 250 |
+
# ============================================================================
|
| 251 |
+
# SmoothQuantizer
|
| 252 |
+
# ============================================================================
|
| 253 |
+
|
| 254 |
+
class SmoothQuantizer:
|
| 255 |
+
"""Aplica quantização W8A8 com SmoothQuant a um modelo.
|
| 256 |
+
|
| 257 |
+
Args:
|
| 258 |
+
config: configuração W8A8
|
| 259 |
+
"""
|
| 260 |
+
|
| 261 |
+
def __init__(self, config: W8A8Config):
|
| 262 |
+
self.config = config
|
| 263 |
+
self.calibrator = SmoothQuantCalibrator(config)
|
| 264 |
+
|
| 265 |
+
def compute_scales(
|
| 266 |
+
self,
|
| 267 |
+
stats: Dict[str, Dict[str, torch.Tensor]],
|
| 268 |
+
) -> Dict[str, torch.Tensor]:
|
| 269 |
+
"""Computa fatores de escala s_i = (max_x^alpha) / (max_w^(1-alpha)).
|
| 270 |
+
|
| 271 |
+
Args:
|
| 272 |
+
stats: dict {layer_name: {"max_x": [d_in], "max_w": [d_in]}}
|
| 273 |
+
|
| 274 |
+
Returns:
|
| 275 |
+
dict {layer_name: scale [d_in]}
|
| 276 |
+
"""
|
| 277 |
+
scales = {}
|
| 278 |
+
alpha = self.config.alpha
|
| 279 |
+
eps = 1e-8
|
| 280 |
+
|
| 281 |
+
for name, s in stats.items():
|
| 282 |
+
max_x = s["max_x"].float()
|
| 283 |
+
max_w = s.get("max_w")
|
| 284 |
+
if max_w is None:
|
| 285 |
+
# Sem max_w (não calibrado), usa só max_x
|
| 286 |
+
scale = torch.ones_like(max_x)
|
| 287 |
+
else:
|
| 288 |
+
max_w = max_w.float().to(max_x.device)
|
| 289 |
+
# s = (max_x^alpha) / (max_w^(1-alpha))
|
| 290 |
+
# Adiciona eps para evitar divisão por zero
|
| 291 |
+
scale = (max_x.clamp(min=eps).pow(alpha) /
|
| 292 |
+
max_w.clamp(min=eps).pow(1 - alpha))
|
| 293 |
+
# Clipa outliers se configurado
|
| 294 |
+
if self.config.outlier_clip_std is not None:
|
| 295 |
+
mean = scale.mean()
|
| 296 |
+
std = scale.std()
|
| 297 |
+
threshold = self.config.outlier_clip_std * std
|
| 298 |
+
scale = scale.clamp(min=mean - threshold, max=mean + threshold)
|
| 299 |
+
# Normaliza para média 1 (preserva escala global)
|
| 300 |
+
scale = scale * (scale.numel() / scale.sum().clamp(min=eps))
|
| 301 |
+
scales[name] = scale
|
| 302 |
+
|
| 303 |
+
return scales
|
| 304 |
+
|
| 305 |
+
def quantize_tensor_per_channel(
|
| 306 |
+
self,
|
| 307 |
+
tensor: torch.Tensor,
|
| 308 |
+
scale: Optional[torch.Tensor] = None,
|
| 309 |
+
n_bits: int = 8,
|
| 310 |
+
axis: int = -1,
|
| 311 |
+
) -> torch.Tensor:
|
| 312 |
+
"""Quantiza um tensor para INT8 (per-channel ou per-tensor).
|
| 313 |
+
|
| 314 |
+
Args:
|
| 315 |
+
tensor: tensor a quantizar
|
| 316 |
+
scale: escala por canal (se None, calcula automaticamente)
|
| 317 |
+
Se fornecido, deve ter shape broadcastable com `tensor`.
|
| 318 |
+
Para per-row em [d_out, d_in], use scale shape [d_out, 1].
|
| 319 |
+
n_bits: número de bits (default 8)
|
| 320 |
+
axis: eixo para per-channel (apenas quando scale=None e é 1D)
|
| 321 |
+
|
| 322 |
+
Returns:
|
| 323 |
+
tensor quantizado (dequantizado para FP32 para uso em forward)
|
| 324 |
+
"""
|
| 325 |
+
qmax = 2 ** (n_bits - 1) - 1 # 127 para INT8 simétrico
|
| 326 |
+
qmin = -qmax
|
| 327 |
+
|
| 328 |
+
if scale is None:
|
| 329 |
+
# Per-tensor
|
| 330 |
+
max_abs = tensor.abs().max()
|
| 331 |
+
scale = max_abs.clamp(min=1e-8) / qmax
|
| 332 |
+
# Quantize-dequantize
|
| 333 |
+
q = torch.round(tensor / scale).clamp(qmin, qmax)
|
| 334 |
+
return q * scale
|
| 335 |
+
else:
|
| 336 |
+
# Per-channel: scale deve ser broadcastable com tensor
|
| 337 |
+
scale = scale.to(tensor.device)
|
| 338 |
+
# Se scale é 1D, expandimos para o eixo
|
| 339 |
+
if scale.dim() == 1:
|
| 340 |
+
shape = [1] * tensor.dim()
|
| 341 |
+
shape[axis] = scale.size(0)
|
| 342 |
+
scale_b = scale.view(shape)
|
| 343 |
+
else:
|
| 344 |
+
# scale já tem shape broadcastable
|
| 345 |
+
scale_b = scale
|
| 346 |
+
q = torch.round(tensor / scale_b).clamp(qmin, qmax)
|
| 347 |
+
return q * scale_b
|
| 348 |
+
|
| 349 |
+
def quantize_model(
|
| 350 |
+
self,
|
| 351 |
+
model: nn.Module,
|
| 352 |
+
dataloader=None,
|
| 353 |
+
) -> nn.Module:
|
| 354 |
+
"""Aplica quantização W8A8 ao modelo.
|
| 355 |
+
|
| 356 |
+
Args:
|
| 357 |
+
model: modelo a quantizar
|
| 358 |
+
dataloader: dados de calibração (necessário para SmoothQuant estático)
|
| 359 |
+
|
| 360 |
+
Returns:
|
| 361 |
+
modelo quantizado (parâmetros substituídos por versões INT8 simuladas)
|
| 362 |
+
"""
|
| 363 |
+
config = self.config
|
| 364 |
+
device = config.device
|
| 365 |
+
model.to(device)
|
| 366 |
+
|
| 367 |
+
# 1. Calibrar se dataloader fornecido
|
| 368 |
+
if dataloader is not None:
|
| 369 |
+
stats = self.calibrator.calibrate(model, dataloader)
|
| 370 |
+
scales = self.compute_scales(stats)
|
| 371 |
+
else:
|
| 372 |
+
stats = {}
|
| 373 |
+
scales = {}
|
| 374 |
+
|
| 375 |
+
# 2. Aplicar SmoothQuant + quantização W8A8 a cada camada
|
| 376 |
+
n_quantized = 0
|
| 377 |
+
n_skipped = 0
|
| 378 |
+
with torch.no_grad():
|
| 379 |
+
for name, module in model.named_modules():
|
| 380 |
+
if not isinstance(module, (nn.Linear, nn.Conv1d)):
|
| 381 |
+
continue
|
| 382 |
+
if "embedding" in name.lower() and not config.quantize_embeddings:
|
| 383 |
+
n_skipped += 1
|
| 384 |
+
continue
|
| 385 |
+
if isinstance(module, nn.Conv1d) and not config.quantize_conv1d:
|
| 386 |
+
n_skipped += 1
|
| 387 |
+
continue
|
| 388 |
+
|
| 389 |
+
# Aplicar SmoothQuant: W' = W * s (migrar escala dos pesos)
|
| 390 |
+
scale = scales.get(name)
|
| 391 |
+
w = module.weight.data
|
| 392 |
+
if scale is not None:
|
| 393 |
+
# Suavizar pesos: multiplicar pela escala
|
| 394 |
+
if isinstance(module, nn.Linear):
|
| 395 |
+
# w: [d_out, d_in], scale: [d_in]
|
| 396 |
+
w_smoothed = w * scale.to(w.device).unsqueeze(0)
|
| 397 |
+
elif isinstance(module, nn.Conv1d):
|
| 398 |
+
# w: [out_ch, in_ch//groups, kernel_size], scale: [in_ch]
|
| 399 |
+
w_smoothed = w * scale.to(w.device).view(1, -1, 1)
|
| 400 |
+
else:
|
| 401 |
+
w_smoothed = w
|
| 402 |
+
else:
|
| 403 |
+
w_smoothed = w
|
| 404 |
+
|
| 405 |
+
# Quantizar pesos para INT8 (simulado — dequantiza de volta)
|
| 406 |
+
if config.scale_type == "per_channel":
|
| 407 |
+
if isinstance(module, nn.Linear):
|
| 408 |
+
# Per-output-channel para pesos
|
| 409 |
+
w_scale = w_smoothed.abs().amax(dim=1) / 127.0
|
| 410 |
+
w_scale = w_scale.clamp(min=1e-8)
|
| 411 |
+
w_q = self.quantize_tensor_per_channel(
|
| 412 |
+
w_smoothed, scale=w_scale.unsqueeze(1), axis=1,
|
| 413 |
+
)
|
| 414 |
+
elif isinstance(module, nn.Conv1d):
|
| 415 |
+
w_scale = w_smoothed.abs().amax(dim=(1, 2)) / 127.0
|
| 416 |
+
w_scale = w_scale.clamp(min=1e-8)
|
| 417 |
+
w_q = self.quantize_tensor_per_channel(
|
| 418 |
+
w_smoothed, scale=w_scale.view(-1, 1, 1), axis=0,
|
| 419 |
+
)
|
| 420 |
+
else:
|
| 421 |
+
# Per-tensor
|
| 422 |
+
w_q = self.quantize_tensor_per_channel(w_smoothed)
|
| 423 |
+
|
| 424 |
+
module.weight.data = w_q.to(module.weight.dtype)
|
| 425 |
+
n_quantized += 1
|
| 426 |
+
|
| 427 |
+
logger.info(
|
| 428 |
+
"SmoothQuant W8A8 aplicado: %d camadas quantizadas, %d ignoradas",
|
| 429 |
+
n_quantized, n_skipped,
|
| 430 |
+
)
|
| 431 |
+
|
| 432 |
+
# Marcar modelo como quantizado (para uso futuro)
|
| 433 |
+
model._is_quantized_w8a8 = True # type: ignore
|
| 434 |
+
model._quantization_config = config # type: ignore
|
| 435 |
+
|
| 436 |
+
return model
|
| 437 |
+
|
| 438 |
+
@staticmethod
|
| 439 |
+
def is_quantized(model: nn.Module) -> bool:
|
| 440 |
+
"""Verifica se um modelo foi quantizado."""
|
| 441 |
+
return getattr(model, "_is_quantized_w8a8", False)
|
| 442 |
+
|
| 443 |
+
|
| 444 |
+
# ============================================================================
|
| 445 |
+
# Helper: quantize modelo inteiro
|
| 446 |
+
# ============================================================================
|
| 447 |
+
|
| 448 |
+
def quantize_model_w8a8(
|
| 449 |
+
model: nn.Module,
|
| 450 |
+
dataloader=None,
|
| 451 |
+
alpha: float = 0.5,
|
| 452 |
+
n_calibration_batches: int = 5,
|
| 453 |
+
device: str = "cpu",
|
| 454 |
+
**kwargs,
|
| 455 |
+
) -> nn.Module:
|
| 456 |
+
"""Atalho para quantizar um modelo W8A8.
|
| 457 |
+
|
| 458 |
+
Args:
|
| 459 |
+
model: modelo a quantizar
|
| 460 |
+
dataloader: dados de calibração (None para quantização dinâmica)
|
| 461 |
+
alpha: fator SmoothQuant (0.5 default)
|
| 462 |
+
n_calibration_batches: batches para calibração
|
| 463 |
+
device: device para calibração
|
| 464 |
+
**kwargs: outros parâmetros de W8A8Config
|
| 465 |
+
|
| 466 |
+
Returns:
|
| 467 |
+
modelo quantizado
|
| 468 |
+
"""
|
| 469 |
+
config = W8A8Config(
|
| 470 |
+
alpha=alpha,
|
| 471 |
+
n_calibration_batches=n_calibration_batches,
|
| 472 |
+
device=device,
|
| 473 |
+
**kwargs,
|
| 474 |
+
)
|
| 475 |
+
quantizer = SmoothQuantizer(config)
|
| 476 |
+
return quantizer.quantize_model(model, dataloader=dataloader)
|
| 477 |
+
|
| 478 |
+
|
| 479 |
+
# ============================================================================
|
| 480 |
+
# Helper: estimar redução de memória
|
| 481 |
+
# ============================================================================
|
| 482 |
+
|
| 483 |
+
def estimate_memory_savings(model: nn.Module) -> Dict[str, float]:
|
| 484 |
+
"""Estima redução de memória após quantização W8A8.
|
| 485 |
+
|
| 486 |
+
Como a quantização é SIMULADA (dequantiza de volta para FP32 para uso em
|
| 487 |
+
forward), a memória real não muda. Esta função reporta a redução POTENCIAL
|
| 488 |
+
se os pesos fossem armazenados como INT8 real.
|
| 489 |
+
|
| 490 |
+
Args:
|
| 491 |
+
model: modelo (preferencialmente quantizado)
|
| 492 |
+
|
| 493 |
+
Returns:
|
| 494 |
+
dict com tamanhos em MB
|
| 495 |
+
"""
|
| 496 |
+
is_quantized = getattr(model, "_is_quantized_w8a8", False)
|
| 497 |
+
fp32_bytes = 0 # embeddings + biases (sempre FP32)
|
| 498 |
+
quantizable_bytes = 0 # pesos que seriam INT8
|
| 499 |
+
|
| 500 |
+
for name, param in model.named_parameters():
|
| 501 |
+
n = param.numel()
|
| 502 |
+
if "embedding" in name.lower():
|
| 503 |
+
# Embeddings mantêm FP32
|
| 504 |
+
fp32_bytes += n * 4
|
| 505 |
+
elif "bias" in name.lower():
|
| 506 |
+
# Biases mantêm FP32 (típico em W8A8)
|
| 507 |
+
fp32_bytes += n * 4
|
| 508 |
+
elif param.dim() >= 2:
|
| 509 |
+
# Pesos de matrizes (Linear, Conv) — quantizáveis
|
| 510 |
+
quantizable_bytes += n
|
| 511 |
+
else:
|
| 512 |
+
# Outros 1D — mantêm FP32
|
| 513 |
+
fp32_bytes += n * 4
|
| 514 |
+
|
| 515 |
+
# Se quantizado: pesos seriam INT8 (1 byte cada)
|
| 516 |
+
# Se não quantizado: pesos seriam FP32 (4 bytes cada)
|
| 517 |
+
if is_quantized:
|
| 518 |
+
int8_bytes = quantizable_bytes * 1 # INT8
|
| 519 |
+
else:
|
| 520 |
+
int8_bytes = quantizable_bytes * 4 # FP32
|
| 521 |
+
|
| 522 |
+
fp32_mb = fp32_bytes / (1024 ** 2)
|
| 523 |
+
int8_mb = int8_bytes / (1024 ** 2)
|
| 524 |
+
total_mb = fp32_mb + int8_mb
|
| 525 |
+
# Sem quantização, tudo seria FP32
|
| 526 |
+
no_quant_mb = (fp32_bytes + quantizable_bytes * 4) / (1024 ** 2)
|
| 527 |
+
reduction = (1 - total_mb / no_quant_mb) * 100 if no_quant_mb > 0 else 0
|
| 528 |
+
|
| 529 |
+
return {
|
| 530 |
+
"fp32_mb": fp32_mb,
|
| 531 |
+
"int8_mb": int8_mb,
|
| 532 |
+
"total_mb": total_mb,
|
| 533 |
+
"no_quant_mb": no_quant_mb,
|
| 534 |
+
"reduction_pct": reduction,
|
| 535 |
+
"is_quantized": is_quantized,
|
| 536 |
+
}
|
| 537 |
+
|
| 538 |
+
|
| 539 |
+
# ============================================================================
|
| 540 |
+
# Self-test
|
| 541 |
+
# ============================================================================
|
| 542 |
+
|
| 543 |
+
def _self_test():
|
| 544 |
+
"""Teste rápido da quantização W8A8."""
|
| 545 |
+
torch.manual_seed(42)
|
| 546 |
+
|
| 547 |
+
# Modelo simples para teste
|
| 548 |
+
class TinyModel(nn.Module):
|
| 549 |
+
def __init__(self):
|
| 550 |
+
super().__init__()
|
| 551 |
+
self.fc1 = nn.Linear(32, 64)
|
| 552 |
+
self.fc2 = nn.Linear(64, 32)
|
| 553 |
+
self.embedding = nn.Embedding(100, 32)
|
| 554 |
+
|
| 555 |
+
def forward(self, x):
|
| 556 |
+
return self.fc2(torch.relu(self.fc1(x)))
|
| 557 |
+
|
| 558 |
+
model = TinyModel()
|
| 559 |
+
print(f"Antes: {sum(p.numel() for p in model.parameters())} params")
|
| 560 |
+
|
| 561 |
+
# Quantizar sem dataloader (dinâmico)
|
| 562 |
+
quantized = quantize_model_w8a8(model, dataloader=None, alpha=0.5)
|
| 563 |
+
print(f"Quantizado: {SmoothQuantizer.is_quantized(quantized)}")
|
| 564 |
+
|
| 565 |
+
# Verificar forward ainda funciona
|
| 566 |
+
x = torch.randn(2, 32)
|
| 567 |
+
out = quantized(x)
|
| 568 |
+
print(f"Output shape: {out.shape}")
|
| 569 |
+
|
| 570 |
+
# Estimar savings
|
| 571 |
+
savings = estimate_memory_savings(quantized)
|
| 572 |
+
print(f"Memory savings: {savings}")
|
| 573 |
+
|
| 574 |
+
|
| 575 |
+
if __name__ == "__main__":
|
| 576 |
+
_self_test()
|
| 577 |
+
|
| 578 |
+
|
| 579 |
+
__all__ = [
|
| 580 |
+
"W8A8Config",
|
| 581 |
+
"SmoothQuantCalibrator",
|
| 582 |
+
"SmoothQuantizer",
|
| 583 |
+
"quantize_model_w8a8",
|
| 584 |
+
"estimate_memory_savings",
|
| 585 |
+
]
|
utils/vqvae2.py
ADDED
|
@@ -0,0 +1,570 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""vqvae2.py — VQ-VAE-2 Hierárquico para compressão multimodal.
|
| 2 |
+
|
| 3 |
+
Implementa Vector Quantized Variational Autoencoder 2 (Roy et al., 2019),
|
| 4 |
+
uma versão hierárquica do VQ-VAE que usa dois níveis de quantização:
|
| 5 |
+
- Top level: captura padrões semânticos globais (low-resolution)
|
| 6 |
+
- Bottom level: captura detalhes locais (high-resolution)
|
| 7 |
+
|
| 8 |
+
==============================================================================
|
| 9 |
+
ANÁLISE MATEMÁTICA E LÓGICA
|
| 10 |
+
==============================================================================
|
| 11 |
+
|
| 12 |
+
ARQUITETURA:
|
| 13 |
+
Input x ∈ R^{B×C×H×W}
|
| 14 |
+
|
|
| 15 |
+
v
|
| 16 |
+
[Bottom Encoder] -> z_b ∈ R^{B×d_b×H×W}
|
| 17 |
+
|
|
| 18 |
+
v
|
| 19 |
+
[Top Encoder] -> z_t ∈ R^{B×d_t×H/2×W/2}
|
| 20 |
+
|
|
| 21 |
+
v
|
| 22 |
+
[Top Quantizer] -> z_t_q (codebook lookup)
|
| 23 |
+
|
|
| 24 |
+
v
|
| 25 |
+
[Top Decoder] -> z_t_dec ∈ R^{B×d_b×H×W}
|
| 26 |
+
|
|
| 27 |
+
v
|
| 28 |
+
[Bottom Quantizer with z_t_dec] -> z_b_q
|
| 29 |
+
|
|
| 30 |
+
v
|
| 31 |
+
[Bottom Decoder] -> x_recon
|
| 32 |
+
|
| 33 |
+
LOSS:
|
| 34 |
+
L = L_recon + beta * L_commit_top + beta * L_commit_bottom
|
| 35 |
+
+ gamma * L_codebook_top + gamma * L_codebook_bottom
|
| 36 |
+
|
| 37 |
+
onde:
|
| 38 |
+
L_recon = MSE(x, x_recon) ou L1
|
| 39 |
+
L_commit_k = ||sg(z_k) - e_k||^2 (commit loss)
|
| 40 |
+
L_codebook_k = ||z_k - sg(e_k)||^2 (codebook loss)
|
| 41 |
+
sg = stop-gradient
|
| 42 |
+
|
| 43 |
+
Update via EMA (Exponential Moving Average) — alternativo ao gradiente:
|
| 44 |
+
n_k^{t+1} = gamma_ema * n_k^t + (1 - gamma_ema) * sum_k 1[assigned]
|
| 45 |
+
m_k^{t+1} = gamma_ema * m_k^t + (1 - gamma_ema) * sum_k z_k
|
| 46 |
+
e_k = m_k / n_k
|
| 47 |
+
|
| 48 |
+
CODEBOOK:
|
| 49 |
+
Top codebook: K_t vetores de dimensão d_t (tipicamente 512-1024 vetores)
|
| 50 |
+
Bottom codebook: K_b vetores de dimensão d_b (tipicamente 512-1024 vetores)
|
| 51 |
+
|
| 52 |
+
==============================================================================
|
| 53 |
+
INTEGRAÇÃO COM CNN-BiGRU
|
| 54 |
+
==============================================================================
|
| 55 |
+
|
| 56 |
+
O VQ-VAE-2 é usado para comprimir representações multimodais (imagem, áudio)
|
| 57 |
+
em códigos discretos que podem ser processados pelo backbone CNN-BiGRU:
|
| 58 |
+
|
| 59 |
+
1. Imagem -> VQ-VAE-2 -> códigos top + bottom (tokens discretos)
|
| 60 |
+
2. Códigos -> embedding lookup -> [B, T, D]
|
| 61 |
+
3. [B, T, D] -> CooperativeCNNBiGRU (como se fosse texto tokenizado)
|
| 62 |
+
|
| 63 |
+
Isto permite que o modelo trate imagem/áudio da mesma forma que texto,
|
| 64 |
+
unificando o pipeline multimodal.
|
| 65 |
+
|
| 66 |
+
Autor: CNN-BiGRU Project
|
| 67 |
+
"""
|
| 68 |
+
from __future__ import annotations
|
| 69 |
+
|
| 70 |
+
import logging
|
| 71 |
+
import math
|
| 72 |
+
from dataclasses import dataclass, field
|
| 73 |
+
from typing import Dict, List, Optional, Tuple
|
| 74 |
+
|
| 75 |
+
import torch
|
| 76 |
+
import torch.nn as nn
|
| 77 |
+
import torch.nn.functional as F
|
| 78 |
+
|
| 79 |
+
logger = logging.getLogger(__name__)
|
| 80 |
+
|
| 81 |
+
|
| 82 |
+
# ============================================================================
|
| 83 |
+
# Configuração
|
| 84 |
+
# ============================================================================
|
| 85 |
+
|
| 86 |
+
@dataclass
|
| 87 |
+
class VQVAE2Config:
|
| 88 |
+
"""Configuração do VQ-VAE-2 hierárquico."""
|
| 89 |
+
# Dimensões de entrada (assumindo features [B, C, H, W])
|
| 90 |
+
in_channels: int = 1
|
| 91 |
+
# Dimensões latentes
|
| 92 |
+
bottom_channels: int = 64 # d_b
|
| 93 |
+
top_channels: int = 32 # d_t
|
| 94 |
+
# Tamanhos dos codebooks
|
| 95 |
+
n_bottom_codes: int = 512 # K_b
|
| 96 |
+
n_top_codes: int = 512 # K_t
|
| 97 |
+
# Beta (commit loss weight)
|
| 98 |
+
commitment_cost: float = 0.25
|
| 99 |
+
# Gamma (codebook loss weight — apenas para grad update, não EMA)
|
| 100 |
+
codebook_cost: float = 0.0
|
| 101 |
+
# Usar EMA para atualizar codebook (default: True)
|
| 102 |
+
use_ema: bool = True
|
| 103 |
+
ema_decay: float = 0.99
|
| 104 |
+
# Usar residual connection (z + delta)
|
| 105 |
+
use_residual: bool = False
|
| 106 |
+
# Número de camadas de down/upsampling
|
| 107 |
+
n_downsample: int = 2
|
| 108 |
+
# Base channels do encoder/decoder
|
| 109 |
+
hidden_channels: int = 64
|
| 110 |
+
# Device
|
| 111 |
+
device: str = "cpu"
|
| 112 |
+
|
| 113 |
+
|
| 114 |
+
# ============================================================================
|
| 115 |
+
# Residual Blocks (para encoder/decoder)
|
| 116 |
+
# ============================================================================
|
| 117 |
+
|
| 118 |
+
def _adaptive_groups(channels: int, max_groups: int = 8) -> int:
|
| 119 |
+
"""Computa número de grupos para GroupNorm que divide channels."""
|
| 120 |
+
if channels <= 0:
|
| 121 |
+
return 1
|
| 122 |
+
g = min(max_groups, channels)
|
| 123 |
+
while channels % g != 0 and g > 1:
|
| 124 |
+
g -= 1
|
| 125 |
+
return max(1, g)
|
| 126 |
+
|
| 127 |
+
|
| 128 |
+
class ResidualBlock(nn.Module):
|
| 129 |
+
"""Bloco residual simples com Conv2d + GroupNorm + SiLU."""
|
| 130 |
+
|
| 131 |
+
def __init__(self, channels: int, hidden_channels: Optional[int] = None):
|
| 132 |
+
super().__init__()
|
| 133 |
+
h = hidden_channels or channels
|
| 134 |
+
self.conv1 = nn.Conv2d(channels, h, kernel_size=3, padding=1)
|
| 135 |
+
self.conv2 = nn.Conv2d(h, channels, kernel_size=3, padding=1)
|
| 136 |
+
# GroupNorm adaptativo por número de canais
|
| 137 |
+
self.norm1 = nn.GroupNorm(_adaptive_groups(channels), channels)
|
| 138 |
+
# norm2 aplica-se à saída de conv1 (h canais)
|
| 139 |
+
self.norm2 = nn.GroupNorm(_adaptive_groups(h), h)
|
| 140 |
+
|
| 141 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 142 |
+
h = F.silu(self.norm1(x))
|
| 143 |
+
h = self.conv1(h)
|
| 144 |
+
h = F.silu(self.norm2(h))
|
| 145 |
+
h = self.conv2(h)
|
| 146 |
+
return x + h
|
| 147 |
+
|
| 148 |
+
|
| 149 |
+
# ============================================================================
|
| 150 |
+
# Encoder / Decoder (Conv2d-based)
|
| 151 |
+
# ============================================================================
|
| 152 |
+
|
| 153 |
+
class ConvEncoder(nn.Module):
|
| 154 |
+
"""Encoder convolucional: reduz resolução espacial por fator 2^n_downsample."""
|
| 155 |
+
|
| 156 |
+
def __init__(
|
| 157 |
+
self,
|
| 158 |
+
in_channels: int,
|
| 159 |
+
out_channels: int,
|
| 160 |
+
hidden_channels: int = 64,
|
| 161 |
+
n_downsample: int = 2,
|
| 162 |
+
):
|
| 163 |
+
super().__init__()
|
| 164 |
+
layers = []
|
| 165 |
+
c_in = in_channels
|
| 166 |
+
for i in range(n_downsample):
|
| 167 |
+
c_out = hidden_channels * (2 ** i)
|
| 168 |
+
layers.extend([
|
| 169 |
+
nn.Conv2d(c_in, c_out, kernel_size=4, stride=2, padding=1),
|
| 170 |
+
nn.GroupNorm(_adaptive_groups(c_out), c_out),
|
| 171 |
+
nn.SiLU(),
|
| 172 |
+
])
|
| 173 |
+
c_in = c_out
|
| 174 |
+
# Residual blocks para refinar
|
| 175 |
+
layers.append(ResidualBlock(c_in, out_channels))
|
| 176 |
+
# Projeta para out_channels
|
| 177 |
+
layers.append(nn.Conv2d(c_in, out_channels, kernel_size=1))
|
| 178 |
+
self.net = nn.Sequential(*layers)
|
| 179 |
+
|
| 180 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 181 |
+
return self.net(x)
|
| 182 |
+
|
| 183 |
+
|
| 184 |
+
class ConvDecoder(nn.Module):
|
| 185 |
+
"""Decoder convolucional: aumenta resolução espacial por fator 2^n_upsample."""
|
| 186 |
+
|
| 187 |
+
def __init__(
|
| 188 |
+
self,
|
| 189 |
+
in_channels: int,
|
| 190 |
+
out_channels: int,
|
| 191 |
+
hidden_channels: int = 64,
|
| 192 |
+
n_upsample: int = 2,
|
| 193 |
+
):
|
| 194 |
+
super().__init__()
|
| 195 |
+
layers = []
|
| 196 |
+
# Primeiro projeta de in_channels para hidden_channels * 2^(n-1)
|
| 197 |
+
c_start = hidden_channels * (2 ** (n_upsample - 1))
|
| 198 |
+
layers.append(nn.Conv2d(in_channels, c_start, kernel_size=1))
|
| 199 |
+
layers.append(ResidualBlock(c_start, c_start))
|
| 200 |
+
|
| 201 |
+
c_in = c_start
|
| 202 |
+
for i in range(n_upsample - 1, -1, -1):
|
| 203 |
+
c_out = hidden_channels * (2 ** max(0, i - 1)) if i > 0 else out_channels
|
| 204 |
+
layers.extend([
|
| 205 |
+
nn.ConvTranspose2d(c_in, c_out, kernel_size=4, stride=2, padding=1),
|
| 206 |
+
nn.GroupNorm(_adaptive_groups(c_out), c_out) if i > 0 else nn.Identity(),
|
| 207 |
+
nn.SiLU() if i > 0 else nn.Identity(),
|
| 208 |
+
])
|
| 209 |
+
c_in = c_out
|
| 210 |
+
# Camada final 1x1 para ajustar canais
|
| 211 |
+
layers.append(nn.Conv2d(c_in, out_channels, kernel_size=1))
|
| 212 |
+
self.net = nn.Sequential(*layers)
|
| 213 |
+
|
| 214 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 215 |
+
return self.net(x)
|
| 216 |
+
|
| 217 |
+
|
| 218 |
+
# ============================================================================
|
| 219 |
+
# Vector Quantizer (com EMA opcional)
|
| 220 |
+
# ============================================================================
|
| 221 |
+
|
| 222 |
+
class VectorQuantizer(nn.Module):
|
| 223 |
+
"""Quantizador por codebook com Straight-Through Estimator (STE).
|
| 224 |
+
|
| 225 |
+
Args:
|
| 226 |
+
n_codes: número de códigos no codebook (K)
|
| 227 |
+
dim: dimensão de cada código (d)
|
| 228 |
+
commitment_cost: peso da commitment loss
|
| 229 |
+
use_ema: se True, atualiza codebook via EMA em vez de gradiente
|
| 230 |
+
ema_decay: decay do EMA
|
| 231 |
+
"""
|
| 232 |
+
|
| 233 |
+
def __init__(
|
| 234 |
+
self,
|
| 235 |
+
n_codes: int,
|
| 236 |
+
dim: int,
|
| 237 |
+
commitment_cost: float = 0.25,
|
| 238 |
+
use_ema: bool = True,
|
| 239 |
+
ema_decay: float = 0.99,
|
| 240 |
+
):
|
| 241 |
+
super().__init__()
|
| 242 |
+
self.n_codes = n_codes
|
| 243 |
+
self.dim = dim
|
| 244 |
+
self.commitment_cost = commitment_cost
|
| 245 |
+
self.use_ema = use_ema
|
| 246 |
+
self.ema_decay = ema_decay
|
| 247 |
+
|
| 248 |
+
# Codebook: [K, d]
|
| 249 |
+
codebook = torch.randn(n_codes, dim) * 0.01
|
| 250 |
+
if use_ema:
|
| 251 |
+
self.register_buffer("codebook", codebook)
|
| 252 |
+
self.register_buffer("ema_count", torch.zeros(n_codes))
|
| 253 |
+
self.register_buffer("ema_weight", codebook.clone())
|
| 254 |
+
else:
|
| 255 |
+
self.codebook = nn.Parameter(codebook)
|
| 256 |
+
|
| 257 |
+
def forward(
|
| 258 |
+
self,
|
| 259 |
+
z: torch.Tensor,
|
| 260 |
+
) -> Tuple[torch.Tensor, torch.Tensor, Dict[str, torch.Tensor]]:
|
| 261 |
+
"""Quantiza z para o codebook mais próximo.
|
| 262 |
+
|
| 263 |
+
Args:
|
| 264 |
+
z: [B, d, H, W] ou [B, d, T]
|
| 265 |
+
|
| 266 |
+
Returns:
|
| 267 |
+
z_q: quantizado [B, d, H, W]
|
| 268 |
+
indices: [B, H, W] índices do codebook
|
| 269 |
+
loss_dict: dict com commitment loss, codebook loss, etc.
|
| 270 |
+
"""
|
| 271 |
+
# Reshape z para [N, d] onde N = B * H * W
|
| 272 |
+
original_shape = z.shape
|
| 273 |
+
if z.dim() == 4:
|
| 274 |
+
B, d, H, W = z.shape
|
| 275 |
+
z_flat = z.permute(0, 2, 3, 1).reshape(-1, d) # [B*H*W, d]
|
| 276 |
+
elif z.dim() == 3:
|
| 277 |
+
B, d, T = z.shape
|
| 278 |
+
z_flat = z.permute(0, 2, 1).reshape(-1, d) # [B*T, d]
|
| 279 |
+
else:
|
| 280 |
+
z_flat = z
|
| 281 |
+
|
| 282 |
+
# Distância euclidiana ao codebook
|
| 283 |
+
# ||z - e||^2 = ||z||^2 - 2*z·e + ||e||^2
|
| 284 |
+
dist = (
|
| 285 |
+
z_flat.pow(2).sum(dim=-1, keepdim=True)
|
| 286 |
+
- 2 * z_flat @ self.codebook.t()
|
| 287 |
+
+ self.codebook.pow(2).sum(dim=-1, keepdim=False).unsqueeze(0)
|
| 288 |
+
) # [N, K]
|
| 289 |
+
|
| 290 |
+
# Encontrar código mais próximo
|
| 291 |
+
indices = dist.argmin(dim=-1) # [N]
|
| 292 |
+
one_hot = F.one_hot(indices, self.n_codes).float() # [N, K]
|
| 293 |
+
|
| 294 |
+
# Quantizado
|
| 295 |
+
z_q = one_hot @ self.codebook # [N, d]
|
| 296 |
+
|
| 297 |
+
# Commitment loss
|
| 298 |
+
commit_loss = F.mse_loss(z_flat, z_q.detach())
|
| 299 |
+
# Codebook loss (apenas se não usar EMA)
|
| 300 |
+
if not self.use_ema:
|
| 301 |
+
codebook_loss = F.mse_loss(z_q, z_flat.detach())
|
| 302 |
+
else:
|
| 303 |
+
codebook_loss = torch.zeros(1, device=z.device)
|
| 304 |
+
|
| 305 |
+
# EMA update
|
| 306 |
+
if self.use_ema and self.training:
|
| 307 |
+
with torch.no_grad():
|
| 308 |
+
# Conta quantas vezes cada código foi usado
|
| 309 |
+
code_count = one_hot.sum(dim=0) # [K]
|
| 310 |
+
# Soma dos z_flat para cada código
|
| 311 |
+
code_sum = one_hot.t() @ z_flat # [K, d]
|
| 312 |
+
|
| 313 |
+
# EMA
|
| 314 |
+
self.ema_count.mul_(self.ema_decay).add_(
|
| 315 |
+
code_count, alpha=1 - self.ema_decay
|
| 316 |
+
)
|
| 317 |
+
self.ema_weight.mul_(self.ema_decay).add_(
|
| 318 |
+
code_sum, alpha=1 - self.ema_decay
|
| 319 |
+
)
|
| 320 |
+
# Atualiza codebook
|
| 321 |
+
self.codebook.copy_(self.ema_weight / self.ema_count.unsqueeze(-1).clamp(min=1e-8))
|
| 322 |
+
|
| 323 |
+
# Straight-Through Estimator: z_q = z + (z_q - z).detach()
|
| 324 |
+
z_q_st = z_flat + (z_q - z_flat).detach()
|
| 325 |
+
|
| 326 |
+
# Reshape de volta
|
| 327 |
+
if z.dim() == 4:
|
| 328 |
+
z_q_st = z_q_st.view(B, H, W, d).permute(0, 3, 1, 2).contiguous()
|
| 329 |
+
indices = indices.view(B, H, W)
|
| 330 |
+
elif z.dim() == 3:
|
| 331 |
+
z_q_st = z_q_st.view(B, T, d).permute(0, 2, 1).contiguous()
|
| 332 |
+
indices = indices.view(B, T)
|
| 333 |
+
else:
|
| 334 |
+
z_q_st = z_q_st.view(original_shape)
|
| 335 |
+
|
| 336 |
+
# Taxa de uso do codebook (para diagnóstico)
|
| 337 |
+
with torch.no_grad():
|
| 338 |
+
used = (one_hot.sum(dim=0) > 0).float().sum()
|
| 339 |
+
usage = used / self.n_codes
|
| 340 |
+
|
| 341 |
+
loss_dict = {
|
| 342 |
+
"commit_loss": commit_loss * self.commitment_cost,
|
| 343 |
+
"codebook_loss": codebook_loss,
|
| 344 |
+
"usage": usage.detach(),
|
| 345 |
+
}
|
| 346 |
+
|
| 347 |
+
return z_q_st, indices, loss_dict
|
| 348 |
+
|
| 349 |
+
|
| 350 |
+
# ============================================================================
|
| 351 |
+
# VQ-VAE-2 Hierárquico
|
| 352 |
+
# ============================================================================
|
| 353 |
+
|
| 354 |
+
class VQVAE2(nn.Module):
|
| 355 |
+
"""VQ-VAE-2 Hierárquico.
|
| 356 |
+
|
| 357 |
+
Estrutura:
|
| 358 |
+
x -> BottomEnc -> z_b -> TopEnc -> z_t -> TopQ -> z_t_q
|
| 359 |
+
|
|
| 360 |
+
v
|
| 361 |
+
TopDec -> z_t_dec
|
| 362 |
+
|
|
| 363 |
+
v
|
| 364 |
+
z_b + z_t_dec -> BottomQ -> z_b_q -> BottomDec -> x_recon
|
| 365 |
+
"""
|
| 366 |
+
|
| 367 |
+
def __init__(self, config: VQVAE2Config):
|
| 368 |
+
super().__init__()
|
| 369 |
+
self.config = config
|
| 370 |
+
|
| 371 |
+
# Bottom encoder/decoder
|
| 372 |
+
self.bottom_encoder = ConvEncoder(
|
| 373 |
+
in_channels=config.in_channels,
|
| 374 |
+
out_channels=config.bottom_channels,
|
| 375 |
+
hidden_channels=config.hidden_channels,
|
| 376 |
+
n_downsample=config.n_downsample,
|
| 377 |
+
)
|
| 378 |
+
self.bottom_decoder = ConvDecoder(
|
| 379 |
+
in_channels=config.bottom_channels,
|
| 380 |
+
out_channels=config.in_channels,
|
| 381 |
+
hidden_channels=config.hidden_channels,
|
| 382 |
+
n_upsample=config.n_downsample,
|
| 383 |
+
)
|
| 384 |
+
|
| 385 |
+
# Top encoder/decoder (operam sobre z_b)
|
| 386 |
+
self.top_encoder = ConvEncoder(
|
| 387 |
+
in_channels=config.bottom_channels,
|
| 388 |
+
out_channels=config.top_channels,
|
| 389 |
+
hidden_channels=config.hidden_channels,
|
| 390 |
+
n_downsample=1, # downsample adicional por 2
|
| 391 |
+
)
|
| 392 |
+
self.top_decoder = ConvDecoder(
|
| 393 |
+
in_channels=config.top_channels,
|
| 394 |
+
out_channels=config.bottom_channels,
|
| 395 |
+
hidden_channels=config.hidden_channels,
|
| 396 |
+
n_upsample=1,
|
| 397 |
+
)
|
| 398 |
+
|
| 399 |
+
# Quantizadores
|
| 400 |
+
self.top_quantizer = VectorQuantizer(
|
| 401 |
+
n_codes=config.n_top_codes,
|
| 402 |
+
dim=config.top_channels,
|
| 403 |
+
commitment_cost=config.commitment_cost,
|
| 404 |
+
use_ema=config.use_ema,
|
| 405 |
+
ema_decay=config.ema_decay,
|
| 406 |
+
)
|
| 407 |
+
self.bottom_quantizer = VectorQuantizer(
|
| 408 |
+
n_codes=config.n_bottom_codes,
|
| 409 |
+
dim=config.bottom_channels,
|
| 410 |
+
commitment_cost=config.commitment_cost,
|
| 411 |
+
use_ema=config.use_ema,
|
| 412 |
+
ema_decay=config.ema_decay,
|
| 413 |
+
)
|
| 414 |
+
|
| 415 |
+
def forward(
|
| 416 |
+
self,
|
| 417 |
+
x: torch.Tensor,
|
| 418 |
+
) -> Dict[str, torch.Tensor]:
|
| 419 |
+
"""Forward pass completo.
|
| 420 |
+
|
| 421 |
+
Args:
|
| 422 |
+
x: [B, C, H, W] entrada (imagem, espectrograma, etc.)
|
| 423 |
+
|
| 424 |
+
Returns:
|
| 425 |
+
dict com:
|
| 426 |
+
x_recon: [B, C, H, W] reconstrução
|
| 427 |
+
top_indices: [B, H_t, W_t] índices top
|
| 428 |
+
bottom_indices: [B, H_b, W_b] índices bottom
|
| 429 |
+
loss: loss total (recon + commit + codebook)
|
| 430 |
+
loss_dict: breakdown
|
| 431 |
+
"""
|
| 432 |
+
# 1. Bottom encoder
|
| 433 |
+
z_b = self.bottom_encoder(x) # [B, d_b, H_b, W_b]
|
| 434 |
+
|
| 435 |
+
# 2. Top encoder (sobre z_b)
|
| 436 |
+
z_t = self.top_encoder(z_b) # [B, d_t, H_t, W_t]
|
| 437 |
+
|
| 438 |
+
# 3. Top quantizer
|
| 439 |
+
z_t_q, top_indices, top_loss = self.top_quantizer(z_t)
|
| 440 |
+
|
| 441 |
+
# 4. Top decoder (reconstroi z_b aproximado do top)
|
| 442 |
+
z_t_dec = self.top_decoder(z_t_q) # [B, d_b, H_b, W_b]
|
| 443 |
+
|
| 444 |
+
# 5. Bottom quantizer sobre (z_b + z_t_dec) ou (z_b - z_t_dec)
|
| 445 |
+
if self.config.use_residual:
|
| 446 |
+
z_b_residual = z_b - z_t_dec
|
| 447 |
+
z_b_q, bottom_indices, bottom_loss = self.bottom_quantizer(z_b_residual)
|
| 448 |
+
z_b_full = z_b_q + z_t_dec # reconstrói
|
| 449 |
+
else:
|
| 450 |
+
z_b_q, bottom_indices, bottom_loss = self.bottom_quantizer(z_b)
|
| 451 |
+
z_b_full = z_b_q + z_t_dec # combina top + bottom
|
| 452 |
+
|
| 453 |
+
# 6. Bottom decoder
|
| 454 |
+
x_recon = self.bottom_decoder(z_b_full) # [B, C, H, W]
|
| 455 |
+
|
| 456 |
+
# 7. Loss
|
| 457 |
+
recon_loss = F.mse_loss(x_recon, x)
|
| 458 |
+
total_loss = (
|
| 459 |
+
recon_loss
|
| 460 |
+
+ top_loss["commit_loss"]
|
| 461 |
+
+ bottom_loss["commit_loss"]
|
| 462 |
+
+ top_loss["codebook_loss"]
|
| 463 |
+
+ bottom_loss["codebook_loss"]
|
| 464 |
+
)
|
| 465 |
+
|
| 466 |
+
return {
|
| 467 |
+
"x_recon": x_recon,
|
| 468 |
+
"top_indices": top_indices,
|
| 469 |
+
"bottom_indices": bottom_indices,
|
| 470 |
+
"loss": total_loss,
|
| 471 |
+
"loss_dict": {
|
| 472 |
+
"recon_loss": recon_loss.detach(),
|
| 473 |
+
"top_commit": top_loss["commit_loss"].detach(),
|
| 474 |
+
"bottom_commit": bottom_loss["commit_loss"].detach(),
|
| 475 |
+
"top_usage": top_loss["usage"],
|
| 476 |
+
"bottom_usage": bottom_loss["usage"],
|
| 477 |
+
},
|
| 478 |
+
}
|
| 479 |
+
|
| 480 |
+
def encode(self, x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
| 481 |
+
"""Apenas encoder (para uso em pipeline multimodal).
|
| 482 |
+
|
| 483 |
+
Returns:
|
| 484 |
+
(top_indices, bottom_indices) — códigos discretos
|
| 485 |
+
"""
|
| 486 |
+
z_b = self.bottom_encoder(x)
|
| 487 |
+
z_t = self.top_encoder(z_b)
|
| 488 |
+
_, top_idx, _ = self.top_quantizer(z_t)
|
| 489 |
+
z_t_dec = self.top_decoder(self.top_quantizer.codebook[top_idx].permute(
|
| 490 |
+
0, 3, 1, 2
|
| 491 |
+
) if top_idx.dim() == 3 else self.top_quantizer.codebook[top_idx])
|
| 492 |
+
# Simplificado: apenas retorna top_idx e bottom_idx brutos
|
| 493 |
+
if self.config.use_residual:
|
| 494 |
+
z_b_resid = z_b - z_t_dec
|
| 495 |
+
_, bottom_idx, _ = self.bottom_quantizer(z_b_resid)
|
| 496 |
+
else:
|
| 497 |
+
_, bottom_idx, _ = self.bottom_quantizer(z_b)
|
| 498 |
+
return top_idx, bottom_idx
|
| 499 |
+
|
| 500 |
+
def decode_codes(
|
| 501 |
+
self,
|
| 502 |
+
top_indices: torch.Tensor,
|
| 503 |
+
bottom_indices: torch.Tensor,
|
| 504 |
+
) -> torch.Tensor:
|
| 505 |
+
"""Decodifica a partir de códigos discretos."""
|
| 506 |
+
# Top
|
| 507 |
+
z_t_q = self.top_quantizer.codebook[top_indices] # [B, H_t, W_t, d_t]
|
| 508 |
+
# Adicionar dim de canal se necessário
|
| 509 |
+
if z_t_q.dim() == 4:
|
| 510 |
+
z_t_q = z_t_q.permute(0, 3, 1, 2).contiguous()
|
| 511 |
+
z_t_dec = self.top_decoder(z_t_q)
|
| 512 |
+
|
| 513 |
+
# Bottom
|
| 514 |
+
z_b_q = self.bottom_quantizer.codebook[bottom_indices]
|
| 515 |
+
if z_b_q.dim() == 4:
|
| 516 |
+
z_b_q = z_b_q.permute(0, 3, 1, 2).contiguous()
|
| 517 |
+
z_b_full = z_b_q + z_t_dec
|
| 518 |
+
|
| 519 |
+
# Decoder
|
| 520 |
+
x_recon = self.bottom_decoder(z_b_full)
|
| 521 |
+
return x_recon
|
| 522 |
+
|
| 523 |
+
|
| 524 |
+
# ============================================================================
|
| 525 |
+
# Self-test
|
| 526 |
+
# ============================================================================
|
| 527 |
+
|
| 528 |
+
def _self_test():
|
| 529 |
+
"""Teste rápido do VQ-VAE-2."""
|
| 530 |
+
torch.manual_seed(42)
|
| 531 |
+
config = VQVAE2Config(
|
| 532 |
+
in_channels=1,
|
| 533 |
+
bottom_channels=16,
|
| 534 |
+
top_channels=8,
|
| 535 |
+
n_bottom_codes=64,
|
| 536 |
+
n_top_codes=64,
|
| 537 |
+
n_downsample=2,
|
| 538 |
+
hidden_channels=16,
|
| 539 |
+
use_ema=True,
|
| 540 |
+
commitment_cost=0.25,
|
| 541 |
+
)
|
| 542 |
+
vqvae = VQVAE2(config)
|
| 543 |
+
print(f"VQ-VAE-2 params: {sum(p.numel() for p in vqvae.parameters())}")
|
| 544 |
+
|
| 545 |
+
# Forward
|
| 546 |
+
x = torch.randn(2, 1, 16, 16)
|
| 547 |
+
out = vqvae(x)
|
| 548 |
+
print(f"Recon: {out['x_recon'].shape}")
|
| 549 |
+
print(f"Top indices: {out['top_indices'].shape}")
|
| 550 |
+
print(f"Bottom indices: {out['bottom_indices'].shape}")
|
| 551 |
+
print(f"Loss: {out['loss'].item():.4f}")
|
| 552 |
+
print(f"Loss dict: {out['loss_dict']}")
|
| 553 |
+
|
| 554 |
+
# Test encode/decode
|
| 555 |
+
top_idx, bot_idx = vqvae.encode(x)
|
| 556 |
+
print(f"Encode: top {top_idx.shape}, bot {bot_idx.shape}")
|
| 557 |
+
|
| 558 |
+
|
| 559 |
+
if __name__ == "__main__":
|
| 560 |
+
_self_test()
|
| 561 |
+
|
| 562 |
+
|
| 563 |
+
__all__ = [
|
| 564 |
+
"VQVAE2Config",
|
| 565 |
+
"ResidualBlock",
|
| 566 |
+
"ConvEncoder",
|
| 567 |
+
"ConvDecoder",
|
| 568 |
+
"VectorQuantizer",
|
| 569 |
+
"VQVAE2",
|
| 570 |
+
]
|