PowerMachine commited on
Commit
6da5316
·
verified ·
1 Parent(s): 214d542

v3.0: 9 novos módulos (CyclicReasoning, Medusa, NLG, NLP, MMA, VQVAE2, W8A8, LongContext 1M, Monitor) + bug fixes (attempt 1)

Browse files
__init__.py CHANGED
@@ -1,5 +1,5 @@
1
  """
2
- CNN-BiGRU — Modelo Multimodal Cooperativo Autoaprendível (v2.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,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
- NOVO 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,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__ = "2.0.0"
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, com:
 
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 de 50 samples
8
 
9
- Para o projeto CNN-BiGRU, usamos um gerador sintético multi-modal padrão
10
- quando a rede ou HF não está acessível, garantindo que o teste de 50 amostras
11
- sempre execute.
 
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
- context_window.init_cache(batch_size=1, device=input_a.device)
 
 
 
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 (mantém janela)
291
- # Se context_window ativo, usa-o; senão, trunca para max_seq
292
- if context_window is not None:
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 * 2 * 2, gru_hidden), # 256*2 = 512 (fused A+B do premissa+passo)
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' (v2.0).
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
- V2.0 melhorias:
11
- - Usa `upload_folder` (single call) em vez de `upload_file` em loop
12
- - Filtra arquivos por extensão e por padrões (ignora caches, logs, etc.)
13
- - Mantém árvore de pastas e subpastas conforme existente localmente
14
- - Log mais detalhado de sucesso/falha
 
 
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"v2.0: EWC + Context Window + RoPE + TransformerBlock + bug fixes (attempt {attempt})",
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[:10]:
172
  rel = f.relative_to(REPO_ROOT)
173
  logger.info(f" - {rel}")
174
- if len(files_to_upload) > 10:
175
- logger.info(f" ... e mais {len(files_to_upload) - 10} arquivos")
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
- # 4. Remove HF_TOKEN do ambiente (requisito do usuário)
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
+ ]