v3.0: reorganiza arquivos sob cnn_bigru/ (preserva árvore de pastas)
Browse files- cnn_bigru/README.md +276 -0
- cnn_bigru/__init__.py +327 -0
- cnn_bigru/data/__init__.py +0 -0
- cnn_bigru/data/streaming_dataset.py +352 -0
- cnn_bigru/docs/MATH_ANALYSIS.md +819 -0
- cnn_bigru/inference/__init__.py +0 -0
- cnn_bigru/inference/inference.py +450 -0
- cnn_bigru/losses/__init__.py +0 -0
- cnn_bigru/losses/losses.py +271 -0
- cnn_bigru/models/__init__.py +0 -0
- cnn_bigru/models/context_window.py +593 -0
- cnn_bigru/models/cooperative_bigru.py +519 -0
- cnn_bigru/models/cyclic_reasoning.py +411 -0
- cnn_bigru/models/generator_verifier.py +364 -0
- cnn_bigru/models/medusa_heads.py +409 -0
- cnn_bigru/models/multimodal_attention.py +470 -0
- cnn_bigru/models/multimodal_encoders.py +141 -0
- cnn_bigru/models/multimodal_model.py +214 -0
- cnn_bigru/models/nlg.py +457 -0
- cnn_bigru/models/nlp.py +654 -0
- cnn_bigru/models/rope.py +196 -0
- cnn_bigru/models/transformer_block.py +401 -0
- cnn_bigru/requirements.txt +29 -0
- cnn_bigru/scripts/push_to_hf.py +294 -0
- cnn_bigru/tests/__init__.py +0 -0
- cnn_bigru/tests/test_500_samples.py +1035 -0
- cnn_bigru/tests/test_50_samples.py +832 -0
- cnn_bigru/tokenizer/__init__.py +0 -0
- cnn_bigru/tokenizer/bbpe_tokenizer.py +173 -0
- cnn_bigru/training/__init__.py +0 -0
- cnn_bigru/training/auto_learner.py +301 -0
- cnn_bigru/training/hypothesis_controller.py +290 -0
- cnn_bigru/training/trainer.py +672 -0
- cnn_bigru/utils/__init__.py +0 -0
- cnn_bigru/utils/ewc.py +382 -0
- cnn_bigru/utils/memory_optimizer.py +139 -0
- cnn_bigru/utils/monitoring.py +729 -0
- cnn_bigru/utils/quantization.py +585 -0
- cnn_bigru/utils/semantic_embeddings.py +160 -0
- cnn_bigru/utils/vqvae2.py +570 -0
- cnn_bigru/utils/xeon_runtime.py +176 -0
cnn_bigru/README.md
ADDED
|
@@ -0,0 +1,276 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# CNN-BiGRU — Modelo Multimodal Cooperativo Autoaprendível (v2.0)
|
| 2 |
+
|
| 3 |
+
Implementação completa de um LLM multimodal com arquitetura **CNN-BiGRU cooperativa**,
|
| 4 |
+
inspirada nos projetos de referência:
|
| 5 |
+
|
| 6 |
+
- `PowerMachine/gru-ring-v13-9-2/xavante_work/`
|
| 7 |
+
- `PowerMachine/BiGRU_T_version/src/bigru_t/`
|
| 8 |
+
- `dados.txt` (pseudocódigo matemático completo)
|
| 9 |
+
|
| 10 |
+
O projeto é **diferente** dos anteriores: agrupa todos os componentes multimodais
|
| 11 |
+
(texto + imagem + áudio) num único pipeline cooperativo, com divisão por competências.
|
| 12 |
+
|
| 13 |
+
---
|
| 14 |
+
|
| 15 |
+
## Novidades v2.0
|
| 16 |
+
|
| 17 |
+
Esta versão adiciona os **módulos faltantes** identificados na análise matemática
|
| 18 |
+
(`docs/MATH_ANALYSIS.md`):
|
| 19 |
+
|
| 20 |
+
### Novos Módulos
|
| 21 |
+
|
| 22 |
+
| Módulo | Função | Arquivo |
|
| 23 |
+
|--------|--------|---------|
|
| 24 |
+
| **EWC** (Elastic Weight Consolidation) | Aprendizado contínuo sem catastrophic forgetting via matriz de Fisher | `utils/ewc.py` |
|
| 25 |
+
| **Context Window** | Janela deslizante com cache KV para sequências longas (Sink+Sliding / StreamingLLM) | `models/context_window.py` |
|
| 26 |
+
| **RoPE** (Rotary Position Embeddings) | Codificação posicional rotacional (atenção relativa) | `models/rope.py` |
|
| 27 |
+
| **TransformerBlock** | CausalSelfAttention + FFN com Pre-LN (decoder-only, estilo GPT) | `models/transformer_block.py` |
|
| 28 |
+
|
| 29 |
+
### Bug Fixes Críticos
|
| 30 |
+
|
| 31 |
+
- **`cooperative_bigru.py`**: `SelfAttentionSummary` agora é uma self-attention TRUE
|
| 32 |
+
(CLS-token + multi-head + learned Q/K/V), não mais mean-pool.
|
| 33 |
+
- **`hypothesis_controller.py`**: gate logic corrigido — `x_proj` agora é uma projeção
|
| 34 |
+
LINEAR REAL do input (não mais `combined`, o que tornava o gate inefetivo).
|
| 35 |
+
- **`generator_verifier.py`**: `enc_out` agora usa a sequência temporal REAL do encoder
|
| 36 |
+
(não mais `fused.unsqueeze(1).expand(...)` que repetia o mesmo vetor).
|
| 37 |
+
- **`trainer.py`**: verificador REALMENTE chamado (não proxy); saída da hipótese
|
| 38 |
+
REALMENTE usada (não só flag); EWC integrado à perda total.
|
| 39 |
+
- **`losses.py`**: todos os `torch.tensor(0.0)` agora usam `device=` correto.
|
| 40 |
+
- **`auto_learner.py`**: `_spec_u` registrado como buffer (migra com `.to(device)`);
|
| 41 |
+
`l2_reg` adicionado ao `AutoLearnConfig`; `estimate_curvature` agora usa `try/finally`.
|
| 42 |
+
- **`inference.py`**: geração autoregressiva REAL via GeneratorCNNBiGRU (não single-step).
|
| 43 |
+
- **`streaming_dataset.py`**: synthetic fallback agora é ALCANÇÁVEL quando HF falha
|
| 44 |
+
silenciosamente; token HF é limpo após uso.
|
| 45 |
+
- **`multimodal_model.py`**: `attn_mask_a`/`attn_mask_b` agora são usados; weight tying opcional.
|
| 46 |
+
|
| 47 |
+
---
|
| 48 |
+
|
| 49 |
+
## Características Principais
|
| 50 |
+
|
| 51 |
+
### 1. Arquitetura Cooperativa de 3 Níveis (do `dados.txt`)
|
| 52 |
+
|
| 53 |
+
| Nível | Mecanismo | Descrição |
|
| 54 |
+
|-------|-----------|-----------|
|
| 55 |
+
| 1 | Cross-Attention Global | Ponte de cooperação pós-CNN entre redes A e B, com máscaras de padding |
|
| 56 |
+
| 2 | Célula GRU Cooperativa Passo-a-Passo | `CellStepCooperativaComAtenuacao` com porta Gated (alfa ∈ [0,1]) |
|
| 57 |
+
| 3 | Self-Attention Final + Fusão | Self-attention TRUE (CLS-token) + concatenação (128+128=256) |
|
| 58 |
+
|
| 59 |
+
### 2. Multimodal (Texto + Imagem + Áudio)
|
| 60 |
+
|
| 61 |
+
- **Texto**: Dual-stream CNN-BiGRU cooperativo (streams A e B)
|
| 62 |
+
- **Imagem**: `ImageEncoder` (CNN 2D → flatten → projeção)
|
| 63 |
+
- **Áudio**: `AudioEncoder` (CNN 1D sobre espectrograma)
|
| 64 |
+
- **Fusão**: `MultimodalFusion` com gated fusion por modalidade
|
| 65 |
+
|
| 66 |
+
### 3. Autoaprendizado (parâmetros matematicamente ajustados)
|
| 67 |
+
|
| 68 |
+
Conforme `dados.txt` seções 9-10:
|
| 69 |
+
|
| 70 |
+
- **Ajuste Dinâmico de LR**:
|
| 71 |
+
- `fator_G = exp(-κ * ||grad_G||) * (1 + v_media) / 2`
|
| 72 |
+
- `fator_V = exp(-κ * ||grad_V||) / (1 + L_V_media)`
|
| 73 |
+
- **Controle de Gradiente**: clipping L2 + spectral normalization (power iteration, buffer registrado)
|
| 74 |
+
- **Inicialização Ortogonal**: matrizes recorrentes e convolucionais (com log de fallback)
|
| 75 |
+
- **Estimativa de Curvatura**: aproximação por diferenças finas (com try/finally)
|
| 76 |
+
|
| 77 |
+
### 4. EWC (Elastic Weight Consolidation) — NOVO v2.0
|
| 78 |
+
|
| 79 |
+
- **Cálculo da diagonal da Matriz de Fisher** (empirical Fisher)
|
| 80 |
+
- **Online EWC**: $F^{\text{agg}} = \gamma F^{\text{prev}} + (1-\gamma) F^{\text{new}}$
|
| 81 |
+
- **Penalidade diferenciável**: $L_{\text{EWC}} = \sum_i \frac{\lambda}{2} F_i (\theta_i - \theta_i^*)^2$
|
| 82 |
+
- **Integração à perda total**: $L_{\text{total}} = \alpha L_G + \beta L_V + \gamma L_{AH} + \delta L_{\text{reg}} + L_{\text{EWC}}$
|
| 83 |
+
|
| 84 |
+
### 5. Context Window — NOVO v2.0
|
| 85 |
+
|
| 86 |
+
- **Cache KV** para CausalSelfAttention (speedup ~T× em inferência)
|
| 87 |
+
- **Sliding Window + Sink** (StreamingLLM): mantém $k$ tokens iniciais + últimos $L_{\max}-k$
|
| 88 |
+
- **Estratégias**: `sliding`, `sink_sliding` (default), `recompute`
|
| 89 |
+
- **Integração com inferência**: `generate_with_sampling(context_window=cw)`
|
| 90 |
+
|
| 91 |
+
### 6. RoPE (Rotary Position Embeddings) — NOVO v2.0
|
| 92 |
+
|
| 93 |
+
- Codificação posicional rotacional: $\theta_i = 10000^{-2i/d}$
|
| 94 |
+
- Propriedade relacional: $\langle \text{RoPE}(q, m), \text{RoPE}(k, n) \rangle = \langle \text{RoPE}(q, m-n), k \rangle$
|
| 95 |
+
- Cache de frequências como buffer (migra com `.to(device)`)
|
| 96 |
+
- Extensão dinâmica para sequências longas
|
| 97 |
+
|
| 98 |
+
### 7. TransformerBlock (Decoder-Only) — NOVO v2.0
|
| 99 |
+
|
| 100 |
+
- `CausalSelfAttention`: multi-head, máscara triangular inferior, QKV unificada
|
| 101 |
+
- `TransformerBlock`: Pre-LN + residual + FFN (Linear → GELU → Linear)
|
| 102 |
+
- `TransformerDecoderStack`: N camadas + LayerNorm final + LM head
|
| 103 |
+
- **Weight tying** opcional (compartilha embedding ↔ lm_head)
|
| 104 |
+
- Suporte a cache KV multicamada para geração autoregressiva
|
| 105 |
+
|
| 106 |
+
### 8. Mecanismo de Hipóteses (ativado em punições)
|
| 107 |
+
|
| 108 |
+
Quando o verificador pune um passo (v < threshold), ativa:
|
| 109 |
+
|
| 110 |
+
- **N hipóteses paralelas** com transforms lineares diferentes
|
| 111 |
+
- **Seleção via Gumbel-Softmax** (diferenciável)
|
| 112 |
+
- **Porta de ativação** baseada em v — CORRIGIDA v2.0: agora a saída das hipóteses é
|
| 113 |
+
REALMENTE usada (interpolação `gate * combined + (1-gate) * x_proj`)
|
| 114 |
+
|
| 115 |
+
### 9. Synergy Search (N tentativas de acerto)
|
| 116 |
+
|
| 117 |
+
No início de cada época, busca N combinações de hiperparâmetros (pesos das perdas)
|
| 118 |
+
e seleciona a configuração de menor perda.
|
| 119 |
+
|
| 120 |
+
### 10. Múltiplas Funções de Perda
|
| 121 |
+
|
| 122 |
+
Conforme `dados.txt` seção 8:
|
| 123 |
+
|
| 124 |
+
```
|
| 125 |
+
L_total = α * L_G_total + β * L_V + γ * L_AH + δ * L_reg + L_EWC
|
| 126 |
+
L_G_total = L_G + λ * L_penal_linear + μ * L_penal_exp
|
| 127 |
+
L_penal_linear = Σ (1 - v_t) (todos os passos)
|
| 128 |
+
L_penal_exp = Σ_{v_t < threshold} exp(γ * (1 - v_t)) (apenas graves)
|
| 129 |
+
L_V = BCE(v_t, rotulo_real)
|
| 130 |
+
L_AH = Σ (1 - τ_t) (anti-alucinação fuzzy)
|
| 131 |
+
L_reg = L2 + curvatura
|
| 132 |
+
L_EWC = Σ (λ/2) * F * (θ - θ*)² [NOVO v2.0]
|
| 133 |
+
```
|
| 134 |
+
|
| 135 |
+
### 11. Camada Anti-Alucinação (Lógica Fuzzy de Łukasiewicz)
|
| 136 |
+
|
| 137 |
+
```python
|
| 138 |
+
NOT(x) = 1 - sigmoid(x)
|
| 139 |
+
AND(x,y) = relu(sigmoid(x) + sigmoid(y) - 1)
|
| 140 |
+
OR(x,y) = min(1, sigmoid(x) + sigmoid(y))
|
| 141 |
+
IMP(x,y) = min(1, 1 - sigmoid(x) + sigmoid(y))
|
| 142 |
+
```
|
| 143 |
+
|
| 144 |
+
### 12. Inferência Avançada
|
| 145 |
+
|
| 146 |
+
- Temperatura, Top-K, Top-P (Nucleus Sampling)
|
| 147 |
+
- Presence Penalty, Frequency Penalty
|
| 148 |
+
- **Geração autoregressiva REAL via GeneratorCNNBiGRU** (v2.0)
|
| 149 |
+
- **Cache KV opcional** via ContextWindowManager (v2.0)
|
| 150 |
+
- Avaliação de Perplexidade (PPL) com generator real
|
| 151 |
+
|
| 152 |
+
---
|
| 153 |
+
|
| 154 |
+
## Estrutura por Competências
|
| 155 |
+
|
| 156 |
+
```
|
| 157 |
+
cnn_bigru/
|
| 158 |
+
├── tokenizer/ # Competência: tokenização
|
| 159 |
+
│ └── bbpe_tokenizer.py # BBPE byte-level tokenizer
|
| 160 |
+
├── utils/ # Competência: utilitários
|
| 161 |
+
│ ├── xeon_runtime.py # Ativação Intel Xeon (AVX512+AMX+IPEX)
|
| 162 |
+
│ ├── memory_optimizer.py # Otimização de memória (AMP, GC, offload)
|
| 163 |
+
│ ├── semantic_embeddings.py # Embeddings semânticos (Sinkhorn OT + IB)
|
| 164 |
+
│ └── ewc.py # NOVO v2.0: Elastic Weight Consolidation
|
| 165 |
+
├── data/ # Competência: carregamento de dados
|
| 166 |
+
│ └── streaming_dataset.py # Dataset streaming multimodal (com fallback)
|
| 167 |
+
├── models/ # Competência: modelos neurais
|
| 168 |
+
│ ├── cooperative_bigru.py # Núcleo CNN-BiGRU cooperativo (self-attn TRUE)
|
| 169 |
+
│ ├── multimodal_encoders.py # Encoders de imagem/áudio + fusão
|
| 170 |
+
│ ├── multimodal_model.py # Modelo multimodal completo (weight tying opc.)
|
| 171 |
+
│ ├── generator_verifier.py # Gerador + Verificador + Anti-alucinação
|
| 172 |
+
│ ├── rope.py # NOVO v2.0: Rotary Position Embeddings
|
| 173 |
+
│ ├── transformer_block.py # NOVO v2.0: CausalSelfAttention + TransformerBlock
|
| 174 |
+
│ └── context_window.py # NOVO v2.0: Sliding window + KV cache
|
| 175 |
+
├── losses/ # Competência: funções de perda
|
| 176 |
+
│ └── losses.py # Multi-loss + EWC penalty (device-safe)
|
| 177 |
+
├── training/ # Competência: treinamento
|
| 178 |
+
│ ├── auto_learner.py # Ajuste dinâmico LR + spectral norm (buffer)
|
| 179 |
+
│ ├── hypothesis_controller.py # Hipóteses + synergy search (gate fix)
|
| 180 |
+
│ └── trainer.py # Loop cooperativo + EWC + verificador real
|
| 181 |
+
├── inference/ # Competência: inferência
|
| 182 |
+
│ └── inference.py # Sampling + PPL + generator real + context window
|
| 183 |
+
├── tests/ # Competência: testes
|
| 184 |
+
│ └── test_50_samples.py # Teste de 50 amostras (com novos módulos v2.0)
|
| 185 |
+
├── scripts/ # Competência: scripts auxiliares
|
| 186 |
+
│ └── push_to_hf.py # Upload ao HF (upload_folder + retry)
|
| 187 |
+
├── docs/ # Competência: documentação
|
| 188 |
+
│ └── MATH_ANALYSIS.md # NOVO v2.0: análise matemática de EWC + Context Window
|
| 189 |
+
└── requirements.txt
|
| 190 |
+
```
|
| 191 |
+
|
| 192 |
+
---
|
| 193 |
+
|
| 194 |
+
## Como Executar
|
| 195 |
+
|
| 196 |
+
### Requisitos
|
| 197 |
+
|
| 198 |
+
```bash
|
| 199 |
+
pip install torch tokenizers huggingface_hub datasets psutil tqdm numpy
|
| 200 |
+
```
|
| 201 |
+
|
| 202 |
+
### Teste de 50 Amostras (com novos módulos v2.0)
|
| 203 |
+
|
| 204 |
+
```bash
|
| 205 |
+
python -m cnn_bigru.tests.test_50_samples
|
| 206 |
+
```
|
| 207 |
+
|
| 208 |
+
O teste executa:
|
| 209 |
+
1. Ativação do runtime Xeon
|
| 210 |
+
2. Treino do BBPE tokenizer
|
| 211 |
+
3. Criação do dataset streaming (50 amostras)
|
| 212 |
+
4. Instanciação dos modelos
|
| 213 |
+
5. **Teste isolado dos novos módulos v2.0**:
|
| 214 |
+
- EWC (Fisher, consolidate, penalty)
|
| 215 |
+
- Context Window (cache KV, eviction)
|
| 216 |
+
- RoPE (preservação de norma)
|
| 217 |
+
- TransformerBlock (forward, cache, weight tying)
|
| 218 |
+
5'. Treinamento cooperativo com EWC + synergy + hipóteses + auto-learn
|
| 219 |
+
6. Inferência com generator + context window
|
| 220 |
+
7. Avaliação de perplexidade com generator
|
| 221 |
+
8. Relatório de erros lógicos ou falhas
|
| 222 |
+
|
| 223 |
+
### Resultado Esperado
|
| 224 |
+
|
| 225 |
+
```
|
| 226 |
+
✓ NENHUM ERRO CRÍTICO encontrado
|
| 227 |
+
Amostras processadas: 50
|
| 228 |
+
Synergy attempts: 3
|
| 229 |
+
Hypothesis activations: variável
|
| 230 |
+
EWC active: True
|
| 231 |
+
Context Window active: True
|
| 232 |
+
Used generator for inference: True
|
| 233 |
+
STATUS: SUCESSO
|
| 234 |
+
```
|
| 235 |
+
|
| 236 |
+
---
|
| 237 |
+
|
| 238 |
+
## Análise Matemática
|
| 239 |
+
|
| 240 |
+
A inclusão de EWC e Context Window foi analisada matematicamente em
|
| 241 |
+
[`docs/MATH_ANALYSIS.md`](docs/MATH_ANALYSIS.md), cobrendo:
|
| 242 |
+
|
| 243 |
+
1. **EWC**: formulação da Fisher Information, integração na perda, garantia de Laplace
|
| 244 |
+
2. **Context Window**: complexidade computacional, políticas de evicção (sliding, sink+sliding, recompute)
|
| 245 |
+
3. **Demais elementos faltantes**: RoPE, CausalSelfAttention, Self-Attention final, etc.
|
| 246 |
+
|
| 247 |
+
---
|
| 248 |
+
|
| 249 |
+
## Referências
|
| 250 |
+
|
| 251 |
+
- `dados.txt` (arquivo de entrada com pseudocódigo completo)
|
| 252 |
+
- `huggingface.co/PowerMachine/gru-ring-v13-9-2/tree/main/xavante_work/`
|
| 253 |
+
- `huggingface.co/PowerMachine/BiGRU_T_version/blob/main/src/bigru_t/tokenizer/bbpe_tokenizer.py`
|
| 254 |
+
- `huggingface.co/PowerMachine/gru-ring-v13-9-2/blob/main/xavante_work/xavante/utils/semantic_embeddings.py`
|
| 255 |
+
- `huggingface.co/PowerMachine/gru-ring-v13-9-2/blob/main/xavante_work/xavante/utils/memory_optimizer.py`
|
| 256 |
+
- `huggingface.co/PowerMachine/gru-ring-v13-9-2/tree/main/xavante_work/xavante/training`
|
| 257 |
+
- `huggingface.co/PowerMachine/BiGRU_T_version/blob/main/src/bigru_t/utils/xeon_runtime.py`
|
| 258 |
+
- `huggingface.co/PowerMachine/gru-ring-v13-9-2/blob/main/xavante_work/flexnet/streaming_datasets_v13_9.py`
|
| 259 |
+
|
| 260 |
+
---
|
| 261 |
+
|
| 262 |
+
## Notas de Segurança
|
| 263 |
+
|
| 264 |
+
- **HF_TOKEN**: usado apenas para push inicial ao repositório. Após o push, o token
|
| 265 |
+
é removido do ambiente e dos scripts. Os scripts publicados não contêm nenhum
|
| 266 |
+
token embutido.
|
| 267 |
+
- **Exceções**: todos os módulos usam `try/except` com logging adequado para
|
| 268 |
+
garantir robustez. O treinamento continua mesmo se alguns batches falharem.
|
| 269 |
+
- **Curvatura**: `estimate_curvature` agora usa `try/finally` para garantir
|
| 270 |
+
restauração dos parâmetros mesmo em caso de exceção.
|
| 271 |
+
|
| 272 |
+
---
|
| 273 |
+
|
| 274 |
+
## Licença
|
| 275 |
+
|
| 276 |
+
MIT
|
cnn_bigru/__init__.py
ADDED
|
@@ -0,0 +1,327 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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)
|
| 6 |
+
- Encoders multimodais (imagem + áudio) com fusão gated
|
| 7 |
+
- Gerador + Verificador + Camada anti-alucinação (lógica fuzzy de Łukasiewicz)
|
| 8 |
+
- Auto-aprendizado: ajuste dinâmico de LR, spectral norm, orthogonal init
|
| 9 |
+
- Hipóteses ativadas por punições (Gumbel-Softmax diferenciável)
|
| 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)
|
| 17 |
+
- TransformerBlock (CausalSelfAttention + FFN)
|
| 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
|
| 39 |
+
__all__ = [
|
| 40 |
+
"__version__",
|
| 41 |
+
"__author__",
|
| 42 |
+
# Tokenizer
|
| 43 |
+
"BBPETokenizer",
|
| 44 |
+
# Utils
|
| 45 |
+
"MemoryOptimizer",
|
| 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",
|
| 60 |
+
# Data
|
| 61 |
+
"MultimodalStreamingDataset",
|
| 62 |
+
"collate_multimodal",
|
| 63 |
+
# Models
|
| 64 |
+
"CooperativeCNNBiGRU",
|
| 65 |
+
"MultimodalCNNBiGRU",
|
| 66 |
+
"ImageEncoder",
|
| 67 |
+
"AudioEncoder",
|
| 68 |
+
"MultimodalFusion",
|
| 69 |
+
"GeneratorCNNBiGRU",
|
| 70 |
+
"VerifierCNNBiGRU",
|
| 71 |
+
"AntiHallucinationLayer",
|
| 72 |
+
"RotaryPositionEmbedding",
|
| 73 |
+
"CausalSelfAttention",
|
| 74 |
+
"TransformerBlock",
|
| 75 |
+
"TransformerDecoderStack",
|
| 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",
|
| 103 |
+
# Training
|
| 104 |
+
"TrainerConfig",
|
| 105 |
+
"CooperativeTrainer",
|
| 106 |
+
"AutoLearnConfig",
|
| 107 |
+
"AutoLearner",
|
| 108 |
+
"HypothesisConfig",
|
| 109 |
+
"HypothesisController",
|
| 110 |
+
"SynergySearcher",
|
| 111 |
+
# Inference
|
| 112 |
+
"generate_with_sampling",
|
| 113 |
+
"evaluate_perplexity",
|
| 114 |
+
]
|
| 115 |
+
|
| 116 |
+
|
| 117 |
+
# Lazy imports — apenas quando acessado
|
| 118 |
+
def __getattr__(name: str):
|
| 119 |
+
if name in ("BBPETokenizer",):
|
| 120 |
+
from .tokenizer.bbpe_tokenizer import BBPETokenizer
|
| 121 |
+
return BBPETokenizer
|
| 122 |
+
if name in ("MemoryOptimizer",):
|
| 123 |
+
from .utils.memory_optimizer import MemoryOptimizer
|
| 124 |
+
return MemoryOptimizer
|
| 125 |
+
if name in ("SemanticEmbedder",):
|
| 126 |
+
from .utils.semantic_embeddings import SemanticEmbedder
|
| 127 |
+
return SemanticEmbedder
|
| 128 |
+
if name in ("EWCConfig", "EWCState"):
|
| 129 |
+
from .utils.ewc import EWCConfig, EWCState
|
| 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
|
| 173 |
+
if name == "MultimodalCNNBiGRU":
|
| 174 |
+
from .models.multimodal_model import MultimodalCNNBiGRU
|
| 175 |
+
return MultimodalCNNBiGRU
|
| 176 |
+
if name in ("ImageEncoder", "AudioEncoder", "MultimodalFusion"):
|
| 177 |
+
from .models.multimodal_encoders import ImageEncoder, AudioEncoder, MultimodalFusion
|
| 178 |
+
if name == "ImageEncoder":
|
| 179 |
+
return ImageEncoder
|
| 180 |
+
if name == "AudioEncoder":
|
| 181 |
+
return AudioEncoder
|
| 182 |
+
return MultimodalFusion
|
| 183 |
+
if name in ("GeneratorCNNBiGRU", "VerifierCNNBiGRU", "AntiHallucinationLayer"):
|
| 184 |
+
from .models.generator_verifier import (
|
| 185 |
+
GeneratorCNNBiGRU,
|
| 186 |
+
VerifierCNNBiGRU,
|
| 187 |
+
AntiHallucinationLayer,
|
| 188 |
+
)
|
| 189 |
+
if name == "GeneratorCNNBiGRU":
|
| 190 |
+
return GeneratorCNNBiGRU
|
| 191 |
+
if name == "VerifierCNNBiGRU":
|
| 192 |
+
return VerifierCNNBiGRU
|
| 193 |
+
return AntiHallucinationLayer
|
| 194 |
+
if name == "RotaryPositionEmbedding":
|
| 195 |
+
from .models.rope import RotaryPositionEmbedding
|
| 196 |
+
return RotaryPositionEmbedding
|
| 197 |
+
if name in ("CausalSelfAttention", "TransformerBlock", "TransformerDecoderStack"):
|
| 198 |
+
from .models.transformer_block import (
|
| 199 |
+
CausalSelfAttention,
|
| 200 |
+
TransformerBlock,
|
| 201 |
+
TransformerDecoderStack,
|
| 202 |
+
)
|
| 203 |
+
if name == "CausalSelfAttention":
|
| 204 |
+
return CausalSelfAttention
|
| 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":
|
| 212 |
+
return 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":
|
| 303 |
+
return TrainerConfig
|
| 304 |
+
return CooperativeTrainer
|
| 305 |
+
if name in ("AutoLearnConfig", "AutoLearner"):
|
| 306 |
+
from .training.auto_learner import AutoLearnConfig, AutoLearner
|
| 307 |
+
if name == "AutoLearnConfig":
|
| 308 |
+
return AutoLearnConfig
|
| 309 |
+
return AutoLearner
|
| 310 |
+
if name in ("HypothesisConfig", "HypothesisController", "SynergySearcher"):
|
| 311 |
+
from .training.hypothesis_controller import (
|
| 312 |
+
HypothesisConfig,
|
| 313 |
+
HypothesisController,
|
| 314 |
+
SynergySearcher,
|
| 315 |
+
)
|
| 316 |
+
if name == "HypothesisConfig":
|
| 317 |
+
return HypothesisConfig
|
| 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":
|
| 325 |
+
return generate_with_sampling
|
| 326 |
+
return evaluate_perplexity
|
| 327 |
+
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|
cnn_bigru/data/__init__.py
ADDED
|
File without changes
|
cnn_bigru/data/streaming_dataset.py
ADDED
|
@@ -0,0 +1,352 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
+
|
| 18 |
+
import logging
|
| 19 |
+
import os
|
| 20 |
+
import random
|
| 21 |
+
from dataclasses import dataclass, field
|
| 22 |
+
from typing import Any, Dict, Iterator, List, Optional, Sequence
|
| 23 |
+
|
| 24 |
+
import numpy as np
|
| 25 |
+
import torch
|
| 26 |
+
|
| 27 |
+
logger = logging.getLogger(__name__)
|
| 28 |
+
|
| 29 |
+
|
| 30 |
+
@dataclass
|
| 31 |
+
class MultimodalSample:
|
| 32 |
+
"""Amostra multimodal: texto + imagem (HxWxC float) + áudio (spec TxF)."""
|
| 33 |
+
sample_id: int
|
| 34 |
+
text_a: str # Stream A (e.g. pergunta/título)
|
| 35 |
+
text_b: str # Stream B (e.g. contexto/corpo)
|
| 36 |
+
image: Optional[np.ndarray] = None # [H, W, C] float32 in [0,1]
|
| 37 |
+
audio: Optional[np.ndarray] = None # [T, F] spectrogram float32
|
| 38 |
+
label: int = 0
|
| 39 |
+
metadata: Dict[str, Any] = field(default_factory=dict)
|
| 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",
|
| 52 |
+
"TucanoBR/GigaVerbo",
|
| 53 |
+
"nvidia/OpenMathReasoning",
|
| 54 |
+
"MathLLMs/MathVision",
|
| 55 |
+
"nvidia/OpenMathInstruct-2",
|
| 56 |
+
"dominguesm/restore-punctuation-ptbr-dataset",
|
| 57 |
+
"carolina-c4ai/corpus-carolina",
|
| 58 |
+
]
|
| 59 |
+
|
| 60 |
+
|
| 61 |
+
def _extract_field(sample: Dict[str, Any], candidates: List[str]) -> Optional[str]:
|
| 62 |
+
for f in candidates:
|
| 63 |
+
if f in sample:
|
| 64 |
+
v = sample[f]
|
| 65 |
+
if isinstance(v, str) and v.strip():
|
| 66 |
+
return v
|
| 67 |
+
if isinstance(v, list):
|
| 68 |
+
parts = []
|
| 69 |
+
for m in v:
|
| 70 |
+
if isinstance(m, dict):
|
| 71 |
+
c = m.get("content", "")
|
| 72 |
+
if isinstance(c, str) and c.strip():
|
| 73 |
+
parts.append(c)
|
| 74 |
+
elif isinstance(m, str) and m.strip():
|
| 75 |
+
parts.append(m)
|
| 76 |
+
if parts:
|
| 77 |
+
return "\n".join(parts)
|
| 78 |
+
return None
|
| 79 |
+
|
| 80 |
+
|
| 81 |
+
def _try_load_hf_streaming(
|
| 82 |
+
dataset_name: str,
|
| 83 |
+
split: str = "train",
|
| 84 |
+
hf_token: Optional[str] = None,
|
| 85 |
+
):
|
| 86 |
+
"""Tenta carregar dataset HF em streaming. Retorna None se falhar."""
|
| 87 |
+
try:
|
| 88 |
+
from datasets import load_dataset
|
| 89 |
+
ds = load_dataset(dataset_name, split=split, streaming=True, token=hf_token)
|
| 90 |
+
logger.info("Streaming OK: %s[%s]", dataset_name, split)
|
| 91 |
+
return ds
|
| 92 |
+
except Exception as e:
|
| 93 |
+
logger.info("Streaming falhou para %s: %s", dataset_name, str(e)[:120])
|
| 94 |
+
return None
|
| 95 |
+
|
| 96 |
+
|
| 97 |
+
def synthetic_multimodal_stream(
|
| 98 |
+
n_samples: int,
|
| 99 |
+
seed: int = 42,
|
| 100 |
+
image_size: tuple = (28, 28, 1),
|
| 101 |
+
audio_shape: tuple = (64, 40),
|
| 102 |
+
vocab_texts: Optional[Sequence[str]] = None,
|
| 103 |
+
) -> Iterator[MultimodalSample]:
|
| 104 |
+
"""Gera amostras sintéticas multimodais determinísticas.
|
| 105 |
+
|
| 106 |
+
Usado quando o dataset HF está indisponível (offline, rate-limit, etc.).
|
| 107 |
+
Garante que o teste de 50 amostras sempre execute.
|
| 108 |
+
"""
|
| 109 |
+
rng = random.Random(seed)
|
| 110 |
+
np_rng = np.random.RandomState(seed)
|
| 111 |
+
|
| 112 |
+
if vocab_texts is None:
|
| 113 |
+
vocab_texts = [
|
| 114 |
+
"o modelo aprende padrões locais com convoluções",
|
| 115 |
+
"a ponte de cooperação troca informações entre fluxos",
|
| 116 |
+
"grus capturam dependências temporais bidirecionais",
|
| 117 |
+
"atenção cruzada reduz perplexidade em textos longos",
|
| 118 |
+
"a fusão multimodal combina texto imagem e áudio",
|
| 119 |
+
"penalidades evitam repetições viciosas na geração",
|
| 120 |
+
"a camada anti-alucinação usa lógica fuzzy de lukasiewicz",
|
| 121 |
+
"o verificador classifica passos com sigmoid binária",
|
| 122 |
+
"ajustes dinâmicos de lr controlam explosão de gradiente",
|
| 123 |
+
"normalização espectral estabiliza o treinamento",
|
| 124 |
+
]
|
| 125 |
+
|
| 126 |
+
n_vocab = len(vocab_texts)
|
| 127 |
+
for i in range(n_samples):
|
| 128 |
+
text_a = vocab_texts[rng.randrange(n_vocab)]
|
| 129 |
+
text_b = vocab_texts[rng.randrange(n_vocab)]
|
| 130 |
+
# Imagem sintética: padrões estruturados simples
|
| 131 |
+
img = np_rng.rand(*image_size).astype(np.float32)
|
| 132 |
+
# Adiciona um padrão que depende do índice (para discriminação)
|
| 133 |
+
img[i % image_size[0], :, :] = 1.0
|
| 134 |
+
# Áudio sintético: espectrograma
|
| 135 |
+
aud = np_rng.rand(*audio_shape).astype(np.float32) * 0.5
|
| 136 |
+
# Espectro com picos determinísticos
|
| 137 |
+
aud[i % audio_shape[0], :] += 0.5
|
| 138 |
+
yield MultimodalSample(
|
| 139 |
+
sample_id=i,
|
| 140 |
+
text_a=text_a,
|
| 141 |
+
text_b=text_b,
|
| 142 |
+
image=img,
|
| 143 |
+
audio=aud,
|
| 144 |
+
label=i % 3,
|
| 145 |
+
metadata={"source": "synthetic", "idx": i},
|
| 146 |
+
)
|
| 147 |
+
|
| 148 |
+
|
| 149 |
+
class MultimodalStreamingDataset(torch.utils.data.IterableDataset):
|
| 150 |
+
"""Dataset streaming multimodal. Suporta HF + fallback sintético.
|
| 151 |
+
|
| 152 |
+
Args:
|
| 153 |
+
n_samples: número total de amostras a produzir.
|
| 154 |
+
hf_datasets: lista de datasets HF para tentar (em ordem).
|
| 155 |
+
hf_token: token HF (será limpo após uso).
|
| 156 |
+
use_synthetic_fallback: se True, usa sintético quando HF falha.
|
| 157 |
+
seed: seed para reprodutibilidade.
|
| 158 |
+
"""
|
| 159 |
+
|
| 160 |
+
def __init__(
|
| 161 |
+
self,
|
| 162 |
+
n_samples: int = 50,
|
| 163 |
+
hf_datasets: Optional[List[str]] = None,
|
| 164 |
+
hf_token: Optional[str] = None,
|
| 165 |
+
use_synthetic_fallback: bool = True,
|
| 166 |
+
seed: int = 42,
|
| 167 |
+
image_size: tuple = (28, 28, 1),
|
| 168 |
+
audio_shape: tuple = (64, 40),
|
| 169 |
+
):
|
| 170 |
+
super().__init__()
|
| 171 |
+
self.n_samples = n_samples
|
| 172 |
+
self.hf_datasets = hf_datasets or DEFAULT_DATASETS
|
| 173 |
+
# NOTA: o hf_token NÃO é persistido como atributo de instância
|
| 174 |
+
# para evitar que seja exposto em dumps/logs. Em vez disso, é
|
| 175 |
+
# passado como parâmetro local durante a iteração.
|
| 176 |
+
self._hf_token = hf_token # private, limpo após iter
|
| 177 |
+
self.use_synthetic_fallback = use_synthetic_fallback
|
| 178 |
+
self.seed = seed
|
| 179 |
+
self.image_size = image_size
|
| 180 |
+
self.audio_shape = audio_shape
|
| 181 |
+
|
| 182 |
+
@property
|
| 183 |
+
def hf_token(self) -> Optional[str]:
|
| 184 |
+
"""Retorna o token HF atual (ou None se já limpo)."""
|
| 185 |
+
return getattr(self, "_hf_token", None)
|
| 186 |
+
|
| 187 |
+
def clear_hf_token(self) -> None:
|
| 188 |
+
"""Limpa o token HF da memória da instância (boa prática de segurança)."""
|
| 189 |
+
self._hf_token = None
|
| 190 |
+
|
| 191 |
+
def _try_hf(self) -> Iterator[Dict[str, Any]]:
|
| 192 |
+
"""Tenta carregar amostras HF. Retorna iterator vazio se falhar."""
|
| 193 |
+
for ds_name in self.hf_datasets:
|
| 194 |
+
ds = _try_load_hf_streaming(ds_name, split="train", hf_token=self._hf_token)
|
| 195 |
+
if ds is None:
|
| 196 |
+
continue
|
| 197 |
+
count = 0
|
| 198 |
+
text_candidates = ["text", "content", "question", "problem", "input",
|
| 199 |
+
"conversa", "description", "prompt", "instruction"]
|
| 200 |
+
label_candidates = ["answer", "response", "output", "solution",
|
| 201 |
+
"punctuated", "restored"]
|
| 202 |
+
for raw in ds:
|
| 203 |
+
if count >= self.n_samples:
|
| 204 |
+
break
|
| 205 |
+
try:
|
| 206 |
+
text_a = _extract_field(raw, text_candidates) or ""
|
| 207 |
+
text_b = _extract_field(raw, label_candidates) or ""
|
| 208 |
+
if len(text_a) < 5:
|
| 209 |
+
text_a = "pergunta de exemplo sobre o tema"
|
| 210 |
+
if len(text_b) < 5:
|
| 211 |
+
text_b = "resposta de exemplo para contexto"
|
| 212 |
+
# Image/audio placeholder: geramos sintéticos para manter multimodal
|
| 213 |
+
img = np.random.rand(*self.image_size).astype(np.float32) * 0.5
|
| 214 |
+
aud = np.random.rand(*self.audio_shape).astype(np.float32) * 0.5
|
| 215 |
+
yield {
|
| 216 |
+
"sample_id": count,
|
| 217 |
+
"text_a": text_a[:200], # truncado para teste
|
| 218 |
+
"text_b": text_b[:200],
|
| 219 |
+
"image": img,
|
| 220 |
+
"audio": aud,
|
| 221 |
+
"label": count % 3,
|
| 222 |
+
"metadata": {"source": ds_name, "idx": count},
|
| 223 |
+
}
|
| 224 |
+
count += 1
|
| 225 |
+
except Exception as e:
|
| 226 |
+
logger.debug("skip sample %d from %s: %s", count, ds_name, e)
|
| 227 |
+
continue
|
| 228 |
+
if count > 0:
|
| 229 |
+
logger.info("HF yield %d samples de %s", count, ds_name)
|
| 230 |
+
return
|
| 231 |
+
|
| 232 |
+
def __iter__(self) -> Iterator[Dict[str, Any]]:
|
| 233 |
+
"""Itera sobre as amostras: tenta HF primeiro, depois fallback sintético.
|
| 234 |
+
|
| 235 |
+
CORREÇÃO v2.0: O bug original retornava imediatamente após `yield from
|
| 236 |
+
self._try_hf()` mesmo se HF não tivesse produzido nenhuma amostra,
|
| 237 |
+
tornando o fallback sintético INACESSÍVEL quando HF falhava silenciosamente.
|
| 238 |
+
Agora rastreamos o número de amostras produzidas e fazemos fallback se zero.
|
| 239 |
+
"""
|
| 240 |
+
n_yielded = 0
|
| 241 |
+
|
| 242 |
+
# Tentar HF primeiro
|
| 243 |
+
if self.hf_datasets:
|
| 244 |
+
try:
|
| 245 |
+
for sample in self._try_hf():
|
| 246 |
+
yield sample
|
| 247 |
+
n_yielded += 1
|
| 248 |
+
if n_yielded >= self.n_samples:
|
| 249 |
+
return
|
| 250 |
+
except Exception as e:
|
| 251 |
+
logger.warning(f"HF streaming falhou ({e}) — usando fallback sintético")
|
| 252 |
+
|
| 253 |
+
# Se HF não produziu amostras suficientes, usar fallback sintético
|
| 254 |
+
if n_yielded < self.n_samples and self.use_synthetic_fallback:
|
| 255 |
+
if n_yielded == 0:
|
| 256 |
+
logger.info("Nenhuma amostra HF produzida — usando 100%% sintético")
|
| 257 |
+
else:
|
| 258 |
+
logger.info(f"HF produziu apenas {n_yielded}/{self.n_samples} — completando com sintético")
|
| 259 |
+
remaining = self.n_samples - n_yielded
|
| 260 |
+
for s in synthetic_multimodal_stream(
|
| 261 |
+
n_samples=remaining,
|
| 262 |
+
seed=self.seed + n_yielded, # seed diferente para variar
|
| 263 |
+
image_size=self.image_size,
|
| 264 |
+
audio_shape=self.audio_shape,
|
| 265 |
+
):
|
| 266 |
+
yield {
|
| 267 |
+
"sample_id": s.sample_id + n_yielded, # offset para não colidir
|
| 268 |
+
"text_a": s.text_a,
|
| 269 |
+
"text_b": s.text_b,
|
| 270 |
+
"image": s.image,
|
| 271 |
+
"audio": s.audio,
|
| 272 |
+
"label": s.label,
|
| 273 |
+
"metadata": {**s.metadata, "source": "synthetic_fallback"},
|
| 274 |
+
}
|
| 275 |
+
n_yielded += 1
|
| 276 |
+
if n_yielded >= self.n_samples:
|
| 277 |
+
# Limpa o token HF após uso completo (boa prática de segurança)
|
| 278 |
+
self.clear_hf_token()
|
| 279 |
+
return
|
| 280 |
+
# Limpa o token HF após uso completo (boa prática de segurança)
|
| 281 |
+
self.clear_hf_token()
|
| 282 |
+
|
| 283 |
+
|
| 284 |
+
def collate_multimodal(
|
| 285 |
+
batch: List[Dict[str, Any]],
|
| 286 |
+
tokenizer,
|
| 287 |
+
max_len: int = 64,
|
| 288 |
+
) -> Dict[str, torch.Tensor]:
|
| 289 |
+
"""Cola um batch de amostras multimodais em tensores.
|
| 290 |
+
|
| 291 |
+
Returns dict com:
|
| 292 |
+
input_ids_a: [B, T] (stream A)
|
| 293 |
+
input_ids_b: [B, T] (stream B)
|
| 294 |
+
attn_mask_a: [B, T]
|
| 295 |
+
attn_mask_b: [B, T]
|
| 296 |
+
images: [B, C, H, W]
|
| 297 |
+
audios: [B, 1, T, F]
|
| 298 |
+
labels: [B]
|
| 299 |
+
"""
|
| 300 |
+
texts_a = [b["text_a"] for b in batch]
|
| 301 |
+
texts_b = [b["text_b"] for b in batch]
|
| 302 |
+
ids_a = tokenizer.encode_batch(texts_a, add_special=True)
|
| 303 |
+
ids_b = tokenizer.encode_batch(texts_b, add_special=True)
|
| 304 |
+
|
| 305 |
+
pad_id = tokenizer.pad_id
|
| 306 |
+
|
| 307 |
+
def _pad(seqs, max_len):
|
| 308 |
+
out = []
|
| 309 |
+
masks = []
|
| 310 |
+
for s in seqs:
|
| 311 |
+
s = s[:max_len]
|
| 312 |
+
n = len(s)
|
| 313 |
+
padded = s + [pad_id] * (max_len - n)
|
| 314 |
+
mask = [1] * n + [0] * (max_len - n)
|
| 315 |
+
out.append(padded)
|
| 316 |
+
masks.append(mask)
|
| 317 |
+
return out, masks
|
| 318 |
+
|
| 319 |
+
padded_a, masks_a = _pad(ids_a, max_len)
|
| 320 |
+
padded_b, masks_b = _pad(ids_b, max_len)
|
| 321 |
+
|
| 322 |
+
images = np.stack([b["image"] for b in batch]) # [B, H, W, C]
|
| 323 |
+
# converte para [B, C, H, W]
|
| 324 |
+
if images.ndim == 4:
|
| 325 |
+
images = np.transpose(images, (0, 3, 1, 2))
|
| 326 |
+
else:
|
| 327 |
+
images = images[:, None, :, :] # add channel dim
|
| 328 |
+
|
| 329 |
+
audios = np.stack([b["audio"] for b in batch]) # [B, T, F]
|
| 330 |
+
audios = audios[:, None, :, :] # [B, 1, T, F]
|
| 331 |
+
|
| 332 |
+
labels = np.array([b["label"] for b in batch], dtype=np.int64)
|
| 333 |
+
|
| 334 |
+
return {
|
| 335 |
+
"input_ids_a": torch.tensor(padded_a, dtype=torch.long),
|
| 336 |
+
"input_ids_b": torch.tensor(padded_b, dtype=torch.long),
|
| 337 |
+
"attn_mask_a": torch.tensor(masks_a, dtype=torch.float),
|
| 338 |
+
"attn_mask_b": torch.tensor(masks_b, dtype=torch.float),
|
| 339 |
+
"images": torch.tensor(images, dtype=torch.float32),
|
| 340 |
+
"audios": torch.tensor(audios, dtype=torch.float32),
|
| 341 |
+
"labels": torch.tensor(labels, dtype=torch.long),
|
| 342 |
+
"sample_ids": [b["sample_id"] for b in batch],
|
| 343 |
+
}
|
| 344 |
+
|
| 345 |
+
|
| 346 |
+
__all__ = [
|
| 347 |
+
"MultimodalSample",
|
| 348 |
+
"MultimodalStreamingDataset",
|
| 349 |
+
"collate_multimodal",
|
| 350 |
+
"DEFAULT_DATASETS",
|
| 351 |
+
"synthetic_multimodal_stream",
|
| 352 |
+
]
|
cnn_bigru/docs/MATH_ANALYSIS.md
ADDED
|
@@ -0,0 +1,819 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Análise Matemática — EWC, Context Window e Elementos Faltantes
|
| 2 |
+
|
| 3 |
+
**Projeto:** CNN-BiGRU Multimodal Cooperativo Autoaprendível
|
| 4 |
+
**Versão:** 2.0 (aprimoramento com EWC + Context Window + RoPE + Transformer Decoder)
|
| 5 |
+
**Data:** 2026-08-11
|
| 6 |
+
|
| 7 |
+
---
|
| 8 |
+
|
| 9 |
+
## 1. Visão Geral
|
| 10 |
+
|
| 11 |
+
Este documento registra a **análise matemática e lógica** da inclusão de:
|
| 12 |
+
|
| 13 |
+
1. **EWC (Elastic Weight Consolidation)** — aprendizado contínuo sem catastrophic forgetting
|
| 14 |
+
2. **Context Window** — janela deslizante com cache KV para sequências longas
|
| 15 |
+
3. **RoPE (Rotary Position Embeddings)** — codificação posicional rotacional
|
| 16 |
+
4. **CausalSelfAttention / TransformerBlock** — bloco decodificador causal
|
| 17 |
+
5. **Self-Attention final layer** — camada de auto-atenção verdadeira ao final de cada bloco (dados.txt linha 470)
|
| 18 |
+
|
| 19 |
+
Também lista demais elementos faltantes identificados na análise do `dados.txt`.
|
| 20 |
+
|
| 21 |
+
---
|
| 22 |
+
|
| 23 |
+
## 2. Análise Matemática do EWC (Elastic Weight Consolidation)
|
| 24 |
+
|
| 25 |
+
### 2.1. Fundamentação Teórica
|
| 26 |
+
|
| 27 |
+
O EWC (Kirkpatrick et al., 2017) regulariza o treinamento contínuo penalizando
|
| 28 |
+
mudanças grandes nos parâmetros $\theta$ que foram importantes para tarefas
|
| 29 |
+
anteriores. A importância é estimada pela **diagonal da matriz de informação
|
| 30 |
+
de Fisher** $\mathcal{I}$.
|
| 31 |
+
|
| 32 |
+
Para uma tarefa $\mathcal{T}_A$ treinada primeiro, a loss para a próxima tarefa
|
| 33 |
+
$\mathcal{T}_B$ é:
|
| 34 |
+
|
| 35 |
+
$$
|
| 36 |
+
\mathcal{L}_{\text{total}}(\theta) \;=\; \mathcal{L}_{B}(\theta) \;+\; \sum_{i}
|
| 37 |
+
\frac{\lambda_{\text{EWC}}}{2} \, \mathcal{I}_{i}^{A} \, \bigl(\theta_{i} -
|
| 38 |
+
\theta_{A,i}^{*}\bigr)^{2}
|
| 39 |
+
$$
|
| 40 |
+
|
| 41 |
+
onde:
|
| 42 |
+
- $\theta_{A}^{*}$ é o ótimo da tarefa $\mathcal{T}_A$ (congelado após treino de $A$)
|
| 43 |
+
- $\mathcal{I}_{i}^{A}$ é a entrada diagonal $i$-ésima da matriz de Fisher para $A$
|
| 44 |
+
- $\lambda_{\text{EWC}}$ é o peso da regularização
|
| 45 |
+
|
| 46 |
+
### 2.2. Estimativa da Matriz de Fisher (Diagonal)
|
| 47 |
+
|
| 48 |
+
Para modelos com saída probabilística $p(y \mid x; \theta)$, a diagonal da
|
| 49 |
+
Fisher Information Matrix pode ser estimada **em uma única passagem** como a
|
| 50 |
+
média do quadrado dos gradientes pontuais:
|
| 51 |
+
|
| 52 |
+
$$
|
| 53 |
+
\mathcal{I}_{i}^{A} \;\approx\; \frac{1}{N} \sum_{n=1}^{N}
|
| 54 |
+
\Bigl[\,\nabla_{\theta_{i}} \log p\bigl(y_{n}^{*} \mid x_{n}; \theta_{A}^{*}\bigr)
|
| 55 |
+
\,\Bigr]^{2}
|
| 56 |
+
$$
|
| 57 |
+
|
| 58 |
+
onde $y_{n}^{*} = \arg\max_{y} p(y \mid x_n; \theta_A^*)$ (rótulo predito pelo
|
| 59 |
+
próprio modelo — *empirical Fisher*).
|
| 60 |
+
|
| 61 |
+
**Justificativa matemática:** A Fisher Information é a variância esperada do
|
| 62 |
+
gradiente da log-verossimilhança; sob a hipótese de que $y_n^*$ é o rótulo
|
| 63 |
+
verdadeiro (maximização), $\mathbb{E}[(\nabla \log p)^2] = -\mathbb{E}[\nabla^2
|
| 64 |
+
\log p]$, que é a definição da Fisher Information. A versão *empirical* usa
|
| 65 |
+
$y^*$ no lugar de $y$ verdadeiro, que é viesada porém consistente para modelos
|
| 66 |
+
bem-treinados.
|
| 67 |
+
|
| 68 |
+
### 2.3. Integração na Função de Perda do CNN-BiGRU
|
| 69 |
+
|
| 70 |
+
A perda total do projeto (conforme `dados.txt` seção 8) é:
|
| 71 |
+
|
| 72 |
+
$$
|
| 73 |
+
\mathcal{L}_{\text{total}} \;=\; \alpha \, \mathcal{L}_{G}^{\text{total}}
|
| 74 |
+
\;+\; \beta \, \mathcal{L}_{V} \;+\; \gamma \, \mathcal{L}_{AH} \;+\; \delta \,
|
| 75 |
+
\mathcal{L}_{\text{reg}}
|
| 76 |
+
$$
|
| 77 |
+
|
| 78 |
+
A inclusão do EWC adiciona um novo termo $\mathcal{L}_{\text{EWC}}$ à
|
| 79 |
+
regularização $\mathcal{L}_{\text{reg}}$:
|
| 80 |
+
|
| 81 |
+
$$
|
| 82 |
+
\mathcal{L}_{\text{reg}}^{\text{new}} \;=\; \mathcal{L}_{\text{reg}}^{\text{old}}
|
| 83 |
+
\;+\; \mathcal{L}_{\text{EWC}}
|
| 84 |
+
\;=\; \lambda_{L2} \sum_i \theta_i^2 \;+\;
|
| 85 |
+
\kappa \, \text{tr}(\hat{H}) \;+\; \sum_i \frac{\lambda_{\text{EWC}}}{2} \,
|
| 86 |
+
F_i \, (\theta_i - \theta_i^*)^2
|
| 87 |
+
$$
|
| 88 |
+
|
| 89 |
+
**Implementação diferenciável:** a penalidade é uma forma quadrática nos
|
| 90 |
+
parâmetros atuais, portanto $\nabla_{\theta} \mathcal{L}_{\text{EWC}} = \lambda_{\text{EWC}}
|
| 91 |
+
\cdot F \odot (\theta - \theta^*)$, onde $\odot$ é produto Hadamard.
|
| 92 |
+
|
| 93 |
+
### 2.4. Pontos de Integração no Pipeline
|
| 94 |
+
|
| 95 |
+
| Etapa | Ação | Módulo |
|
| 96 |
+
|-------|------|--------|
|
| 97 |
+
| Após treinar tarefa A | Calcular $F$ e armazenar $\theta^*$ | `utils/ewc.py:compute_fisher` |
|
| 98 |
+
| A cada step da tarefa B | Somar $\mathcal{L}_{\text{EWC}}$ à perda | `losses/losses.py` |
|
| 99 |
+
| Treino multitarefa | Manter lista $[(\theta^*_k, F_k)]$ e somar EWC de todas as tarefas anteriores | `utils/ewc.py:EWCState` |
|
| 100 |
+
|
| 101 |
+
### 2.5. Análise de Estabilidade
|
| 102 |
+
|
| 103 |
+
- **Limit $\lambda_{\text{EWC}} \to 0$:** EWC desligado → catastrophic forgetting total
|
| 104 |
+
- **Limit $\lambda_{\text{EWC}} \to \infty$:** parâmetros congelados → sem aprendizado novo
|
| 105 |
+
- **Compromisso prático:** $\lambda_{\text{EWC}} \in [100, 10000]$ tipicamente
|
| 106 |
+
recomendado na literatura
|
| 107 |
+
|
| 108 |
+
**Garantia matemática:** sob a hipótese de que $\mathcal{L}_B$ é fortemente
|
| 109 |
+
convexa em $\theta^*$ e $F \succ 0$, o EWC é equivalente a uma aproximação
|
| 110 |
+
de Laplace da posterior $p(\theta \mid \mathcal{D}_A)$, garantindo que a solução
|
| 111 |
+
ótima permaneça próxima de $\theta^*$ ao longo do treino de $\mathcal{T}_B$.
|
| 112 |
+
|
| 113 |
+
### 2.6. Riscos e Mitigação
|
| 114 |
+
|
| 115 |
+
| Risco | Mitigação |
|
| 116 |
+
|-------|-----------|
|
| 117 |
+
| $F$ é ruidosa para pequenos batches | Usar `n_samples_fisher ≥ 500` e média móvel exponencial |
|
| 118 |
+
| Memória: armazenar $F$ igual ao tamanho do modelo | Usar **diagonal apenas** (mesmo tamanho do modelo, viável) |
|
| 119 |
+
| $\theta^*$ fica defasado após muitas tarefas | Implementar **online EWC**: $F^{\text{new}} = \gamma F^{\text{old}} + F^{\text{current}}$ |
|
| 120 |
+
| EWC pode entrar em conflito com spectral norm | Aplicar EWC **apenas** aos pesos não-normalizados espectralmente |
|
| 121 |
+
|
| 122 |
+
---
|
| 123 |
+
|
| 124 |
+
## 3. Análise Matemática do Context Window
|
| 125 |
+
|
| 126 |
+
### 3.1. Motivação
|
| 127 |
+
|
| 128 |
+
O `dados.txt` seção 7 (Transformer Decoder) define `max_seq_len` como
|
| 129 |
+
comprimento máximo de contexto. Sem um mecanismo de **janela deslizante**, o
|
| 130 |
+
modelo:
|
| 131 |
+
|
| 132 |
+
1. Não consegue processar entradas maiores que `max_seq_len`
|
| 133 |
+
2. Recomputa toda a sequência a cada token gerado → $O(T^2)$ em inferência
|
| 134 |
+
3. Não consegue manter coerência em diálogos longos
|
| 135 |
+
|
| 136 |
+
### 3.2. Formulação Matemática
|
| 137 |
+
|
| 138 |
+
Seja $x_{1:T}$ uma sequência de comprimento $T > L_{\max}$ (janela máxima).
|
| 139 |
+
Definimos a janela deslizante $\mathcal{W}_t$ no passo $t$ como:
|
| 140 |
+
|
| 141 |
+
$$
|
| 142 |
+
\mathcal{W}_t \;=\; x_{\max(1,\,t-L_{\max}+1)\,:\,t}
|
| 143 |
+
$$
|
| 144 |
+
|
| 145 |
+
A saída do modelo no passo $t$ depende apenas de $\mathcal{W}_t$:
|
| 146 |
+
|
| 147 |
+
$$
|
| 148 |
+
h_t \;=\; f_\theta\bigl(\mathcal{W}_t\bigr)
|
| 149 |
+
$$
|
| 150 |
+
|
| 151 |
+
### 3.3. Cache KV para Atenção Eficiente
|
| 152 |
+
|
| 153 |
+
Para a CausalSelfAttention, o custo de inferência autoregressiva sem cache é:
|
| 154 |
+
|
| 155 |
+
$$
|
| 156 |
+
\text{FLOPs}_{\text{no cache}}(T) \;=\; O\bigl(T^2 \cdot d\bigr)
|
| 157 |
+
$$
|
| 158 |
+
|
| 159 |
+
Com cache dos pares $(K, V)$ dos passos anteriores, cada novo token $t$ apenas:
|
| 160 |
+
|
| 161 |
+
1. Calcula $q_t, k_t, v_t$ para o novo token (custo $O(d)$)
|
| 162 |
+
2. Concatena com cache: $K_{1:t} = [K_{1:t-1}; k_t]$, $V_{1:t} = [V_{1:t-1}; v_t]$
|
| 163 |
+
3. Calcula atenção apenas do novo $q_t$ contra todo $K_{1:t}$ (custo $O(t \cdot d)$)
|
| 164 |
+
|
| 165 |
+
$$
|
| 166 |
+
\text{FLOPs}_{\text{with cache}}(T) \;=\; O\bigl(T \cdot d\bigr)
|
| 167 |
+
$$
|
| 168 |
+
|
| 169 |
+
**Speedup teórico:** fator $T$. Para $T = 512$ e $d = 256$, speedup ~512×.
|
| 170 |
+
|
| 171 |
+
### 3.4. Política de Evicção (Chunked Cache)
|
| 172 |
+
|
| 173 |
+
Quando $|\mathcal{W}_t| > L_{\max}$, aplicamos uma política de evicção. As três
|
| 174 |
+
estratégias analisadas:
|
| 175 |
+
|
| 176 |
+
#### 3.4.1. Sliding Window Pura
|
| 177 |
+
Mantém apenas os últimos $L_{\max}$ tokens. Custo: $O(L_{\max})$ por step.
|
| 178 |
+
Desvantagem: perde informação histórica relevante.
|
| 179 |
+
|
| 180 |
+
#### 3.4.2. Sink + Sliding (StreamingLLM)
|
| 181 |
+
Mantém os primeiros $k$ tokens ("sink", tipicamente BOS + system prompt) e os
|
| 182 |
+
últimos $L_{\max} - k$ tokens. Custo: $O(L_{\max})$. Vantagem: preserva
|
| 183 |
+
informação estrutural inicial.
|
| 184 |
+
|
| 185 |
+
$$
|
| 186 |
+
\mathcal{W}_t \;=\; \bigl[\,x_{1:k}\,;\, x_{t-(L_{\max}-k)+1\,:\,t}\,\bigr]
|
| 187 |
+
$$
|
| 188 |
+
|
| 189 |
+
#### 3.4.3. Attention Recomputation (sem cache)
|
| 190 |
+
Recomputa a atenção sobre toda a sequência a cada step. Custo: $O(T^2)$.
|
| 191 |
+
Vantagem: máxima precisão. Desvantagem: proibitivo para $T$ grande.
|
| 192 |
+
|
| 193 |
+
**Implementação escolhida:** Sink + Sliding (StreamingLLM), com $k=4$ tokens
|
| 194 |
+
sink configurável. Esta estratégia tem garantia teórica de estabilidade para
|
| 195 |
+
$T \to \infty$ (Xiao et al., 2023).
|
| 196 |
+
|
| 197 |
+
### 3.5. Integração com CNN-BiGRU Cooperativo
|
| 198 |
+
|
| 199 |
+
O CNN-BiGRU cooperativo é **bidirecional** (processa $x_{1:T}$ em ambos os
|
| 200 |
+
sentidos). Isto quebra a hipótese causal. Para preservar a bidirecionalidade:
|
| 201 |
+
|
| 202 |
+
- **Modo encoder (treino/classificação):** sem cache, processa a janela inteira
|
| 203 |
+
$\mathcal{W}_t$ com BiGRU normal (forward + backward).
|
| 204 |
+
- **Modo decoder (geração):** usa GRU unidirecional + CausalSelfAttention com
|
| 205 |
+
cache KV. Esta é a configuração usada para geração autoregressiva.
|
| 206 |
+
|
| 207 |
+
### 3.6. Análise de Complexidade
|
| 208 |
+
|
| 209 |
+
| Componente | Sem Context Window | Com Context Window |
|
| 210 |
+
|------------|-------------------|-------------------|
|
| 211 |
+
| Geração de $T$ tokens | $O(T^2 d)$ | $O(T \cdot L_{\max} \cdot d)$ |
|
| 212 |
+
| Memória de cache | $O(T d)$ | $O(L_{\max} d)$ |
|
| 213 |
+
| Precisão em longas sequências | Falha para $T > L_{\max}$ | Mantém coerência via sink |
|
| 214 |
+
|
| 215 |
+
### 3.7. Riscos e Mitigação
|
| 216 |
+
|
| 217 |
+
| Risco | Mitigação |
|
| 218 |
+
|-------|-----------|
|
| 219 |
+
| Perda de contexto ao evictar | Estratégia sink + sliding |
|
| 220 |
+
| Cache desatualizado após fine-tuning | Invalidar cache quando pesos mudam |
|
| 221 |
+
| Incompatibilidade com BiGRU (bidirecional) | Usar GRU unidirecional no decoder |
|
| 222 |
+
| Memória cresce com $L_{\max}$ | Limitar $L_{\max} \leq 1024$ para CPU |
|
| 223 |
+
|
| 224 |
+
---
|
| 225 |
+
|
| 226 |
+
## 4. Demais Elementos Faltantes Identificados
|
| 227 |
+
|
| 228 |
+
### 4.1. RoPE (Rotary Position Embeddings)
|
| 229 |
+
|
| 230 |
+
**Status no projeto atual:** AUSENTE. A codificação posicional atual depende
|
| 231 |
+
apenas da convolução 1D (local) e da recorrência GRU (implícita).
|
| 232 |
+
|
| 233 |
+
**Formulação matemática (Su et al., 2021):**
|
| 234 |
+
|
| 235 |
+
Para um par $(q, k)$ em uma dada posição $(m, n)$, RoPE aplica uma rotação
|
| 236 |
+
no plano 2D para cada par de dimensões $(q_{2i}, q_{2i+1})$:
|
| 237 |
+
|
| 238 |
+
$$
|
| 239 |
+
\begin{pmatrix} q'_{2i} \\ q'_{2i+1} \end{pmatrix}
|
| 240 |
+
\;=\;
|
| 241 |
+
\begin{pmatrix} \cos(m\theta_i) & -\sin(m\theta_i) \\ \sin(m\theta_i) & \cos(m\theta_i) \end{pmatrix}
|
| 242 |
+
\begin{pmatrix} q_{2i} \\ q_{2i+1} \end{pmatrix}
|
| 243 |
+
$$
|
| 244 |
+
|
| 245 |
+
com $\theta_i = 10000^{-2i/d}$. A propriedade chave é:
|
| 246 |
+
|
| 247 |
+
$$
|
| 248 |
+
\langle \text{RoPE}(q, m), \text{RoPE}(k, n) \rangle \;=\; \langle
|
| 249 |
+
\text{RoPE}(q, m-n), k \rangle
|
| 250 |
+
$$
|
| 251 |
+
|
| 252 |
+
→ a atenção depende apenas da diferença posicional relativa, não absoluta.
|
| 253 |
+
|
| 254 |
+
**Integração:** Substitui `position_embedding = nn.Embedding(max_seq_len, embed_dim)`
|
| 255 |
+
no TransformerBlock. Aplica-se a `q` e `k` antes do produto escalar.
|
| 256 |
+
|
| 257 |
+
### 4.2. CausalSelfAttention / TransformerBlock
|
| 258 |
+
|
| 259 |
+
**Status no projeto atual:** AUSENTE. Apenas Bahdanau attention existe no
|
| 260 |
+
`generator_verifier.py`, e ela é não-causal (encoder-decoder).
|
| 261 |
+
|
| 262 |
+
**Especificação (dados.txt seções 6-7):**
|
| 263 |
+
|
| 264 |
+
```python
|
| 265 |
+
class CausalSelfAttention(nn.Module):
|
| 266 |
+
# Multi-head com máscara triangular inferior
|
| 267 |
+
# QKV projection unificada
|
| 268 |
+
# Scaled dot-product: (q @ k.T) / sqrt(d_head)
|
| 269 |
+
# Máscara causal: triu(ones(T,T)) == 0 → -inf
|
| 270 |
+
# Softmax + dropout
|
| 271 |
+
|
| 272 |
+
class TransformerBlock(nn.Module):
|
| 273 |
+
# Pre-LN: x = x + attn(ln1(x))
|
| 274 |
+
# x = x + ffn(ln2(x))
|
| 275 |
+
# FFN: Linear(d, 4d) → GELU → Linear(4d, d) → Dropout
|
| 276 |
+
```
|
| 277 |
+
|
| 278 |
+
**Integração:** Como camada opcional no `CooperativeCNNBiGRU` para agregar
|
| 279 |
+
informação global antes da fusão final.
|
| 280 |
+
|
| 281 |
+
### 4.3. Self-Attention Final Layer (dados.txt linha 470)
|
| 282 |
+
|
| 283 |
+
**Status no projeto atual:** PARCIAL. Existe `SelfAttentionSummary` mas usa
|
| 284 |
+
`x.mean(dim=1)` como query — não é uma self-attention verdadeira.
|
| 285 |
+
|
| 286 |
+
**Correção:** substituir por uma self-attention com `query = Linear(d, d)(x)`
|
| 287 |
+
(verdadeira projeção de query, não média).
|
| 288 |
+
|
| 289 |
+
### 4.4. Gradient Checkpointing
|
| 290 |
+
|
| 291 |
+
**Status no projeto atual:** Declarado em `memory_optimizer.py` mas NUNCA
|
| 292 |
+
chamado. Deve ser ativado no `CooperativeTrainer` para reduzir consumo de
|
| 293 |
+
VRAM em modelos grandes.
|
| 294 |
+
|
| 295 |
+
### 4.5. Weight Tying
|
| 296 |
+
|
| 297 |
+
**Status no projeto atual:** AUSENTE. O `dados.txt` linha 632 explicitamente
|
| 298 |
+
recomenda: `self.token_embedding.weight = self.lm_head.weight`. Deve ser
|
| 299 |
+
opcional no `MultimodalCNNBiGRU`.
|
| 300 |
+
|
| 301 |
+
### 4.6. AMP (Automatic Mixed Precision)
|
| 302 |
+
|
| 303 |
+
**Status no projeto atual:** PARCIAL — `memory_optimizer.amp_context` retorna
|
| 304 |
+
`nullcontext()` para CPU. Como o `xeon_runtime.py` detecta AMX/AVX512_BF16,
|
| 305 |
+
deve habilitar `bfloat16` autocast em CPU também (PyTorch 1.10+).
|
| 306 |
+
|
| 307 |
+
### 4.7. Online EWC
|
| 308 |
+
|
| 309 |
+
**Status no projeto atual:** AUSENTE. Após múltiplas tarefas, a soma de
|
| 310 |
+
$F_k$ explode em memória. Online EWC mantém apenas uma $F$ agregada:
|
| 311 |
+
|
| 312 |
+
$$
|
| 313 |
+
F^{\text{agg}} \;=\; \gamma \, F^{\text{prev}} \;+\; (1-\gamma) \, F^{\text{new}}
|
| 314 |
+
$$
|
| 315 |
+
|
| 316 |
+
### 4.8. Tokenizer: add_tokens dinâmico
|
| 317 |
+
|
| 318 |
+
**Status no projeto atual:** AUSENTE. O `BBPETokenizer` não suporta adicionar
|
| 319 |
+
novos tokens após o treino. Útil para fine-tuning com vocabulário especializado.
|
| 320 |
+
|
| 321 |
+
### 4.9. Verificador Realmente Acionado
|
| 322 |
+
|
| 323 |
+
**Status no projeto atual:** BUG. `trainer.py` instancia `VerifierCNNBiGRU`
|
| 324 |
+
mas **NUNCA chama seu forward** — usa um proxy `sigmoid(fused.mean(-1))`.
|
| 325 |
+
|
| 326 |
+
**Correção:** Chamar `verifier(premissas_a, passo_a, premissas_b, passo_b)`
|
| 327 |
+
com tokens reais gerados pelo `GeneratorCNNBiGRU`.
|
| 328 |
+
|
| 329 |
+
### 4.10. Hipótese com Porta Funcional
|
| 330 |
+
|
| 331 |
+
**Status no projeto atual:** BUG. `hypothesis_controller.py` calcula `gate`
|
| 332 |
+
mas `output = gate*combined + (1-gate)*combined = combined` (sempre).
|
| 333 |
+
|
| 334 |
+
**Correção:** `x_proj = hypothesis_layer(x)` (projeção real), `output = gate *
|
| 335 |
+
x_proj + (1-gate) * x` (interpolação entre a hipótese e o input original).
|
| 336 |
+
|
| 337 |
+
### 4.11. Geração Autoregressiva Real
|
| 338 |
+
|
| 339 |
+
**Status no projeto atual:** BUG. `inference.py` chama
|
| 340 |
+
`model(..., mode="generate")` que retorna `[B, V]` (single-step) — não há
|
| 341 |
+
autoregressão real.
|
| 342 |
+
|
| 343 |
+
**Correção:** Implementar loop que:
|
| 344 |
+
1. Processa prompt inicial
|
| 345 |
+
2. A cada passo: extrai último logit, aplica sampling, anexa novo token,
|
| 346 |
+
atualiza cache KV, repete
|
| 347 |
+
|
| 348 |
+
### 4.12. Multi-loss com Device Consistency
|
| 349 |
+
|
| 350 |
+
**Status no projeto atual:** BUG. `losses.py` usa `torch.tensor(0.0)` sem
|
| 351 |
+
`device=`, quebrando em GPU.
|
| 352 |
+
|
| 353 |
+
**Correção:** Sempre criar tensores zero no device correto, ou usar
|
| 354 |
+
`torch.zeros((), device=device)`.
|
| 355 |
+
|
| 356 |
+
### 4.13. Inicialização Ortogonal Robusta
|
| 357 |
+
|
| 358 |
+
**Status no projeto atual:** PARCIAL — `_orthogonal_init` faz fallback
|
| 359 |
+
silencioso para `xavier_uniform_` em caso de exceção.
|
| 360 |
+
|
| 361 |
+
**Correção:** Registrar o erro em log, propagar exceção em modo debug.
|
| 362 |
+
|
| 363 |
+
### 4.14. Penalidade Exponencial Conforme Spec
|
| 364 |
+
|
| 365 |
+
**Status no projeto atual:** DESVIO SEMÂNTICO. `losses.py` aplica
|
| 366 |
+
$\sum \exp(\gamma(1-v))$ **apenas** para $v < \text{threshold}$, mas o
|
| 367 |
+
`dados.txt` seção 8.2 diz "SE v_t < THRESHOLD_ERR: L_exp_penal += ..." (que é
|
| 368 |
+
o comportamento atual — OK na verdade). Mas a soma linear
|
| 369 |
+
$\sum (1-v_t)$ deve ser sobre TODOS os passos, não apenas os penalizados.
|
| 370 |
+
|
| 371 |
+
**Revisão:** está correto conforme spec; apenas documentar claramente.
|
| 372 |
+
|
| 373 |
+
### 4.15. Package `__init__.py`
|
| 374 |
+
|
| 375 |
+
**Status no projeto atual:** VAZIO. Deve exportar principais classes,
|
| 376 |
+
definir `__version__`, configurar logging.
|
| 377 |
+
|
| 378 |
+
---
|
| 379 |
+
|
| 380 |
+
## 5. Resumo dos Elementos a Implementar
|
| 381 |
+
|
| 382 |
+
| ID | Elemento | Prioridade | Complexidade | Impacto |
|
| 383 |
+
|----|----------|-----------|--------------|---------|
|
| 384 |
+
| E1 | EWC (Fisher + penalty) | Alta | Média | Mitiga catastrophic forgetting |
|
| 385 |
+
| E2 | Context Window (sliding + KV cache) | Alta | Alta | Suporte a sequências longas |
|
| 386 |
+
| E3 | RoPE | Média | Baixa | Codificação posicional relativa |
|
| 387 |
+
| E4 | CausalSelfAttention + TransformerBlock | Alta | Média | Decoder autoregressivo |
|
| 388 |
+
| E5 | Self-Attention final layer correta | Alta | Baixa | Resumo temporal conforme spec |
|
| 389 |
+
| E6 | Gradient Checkpointing | Baixa | Baixa | Reduz VRAM |
|
| 390 |
+
| E7 | Weight Tying | Baixa | Baixa | Reduz parâmetros |
|
| 391 |
+
| E8 | AMP em CPU (bfloat16) | Média | Baixa | Speedup em Xeon |
|
| 392 |
+
| E9 | Online EWC | Baixa | Baixa | Suporte a múltiplas tarefas |
|
| 393 |
+
| E10 | Verificador real no trainer | Alta | Média | Funcionalidade core |
|
| 394 |
+
| E11 | Hipótese com porta funcional | Alta | Baixa | Funcionalidade core |
|
| 395 |
+
| E12 | Geração autoregressiva real | Alta | Alta | Inferência correta |
|
| 396 |
+
| E13 | Multi-loss device-safe | Alta | Baixa | Correção de bug |
|
| 397 |
+
| E14 | Orthogonal init robusto | Média | Baixa | Estabilidade |
|
| 398 |
+
| E15 | `__init__.py` com exports | Baixa | Baixa | Usabilidade |
|
| 399 |
+
|
| 400 |
+
---
|
| 401 |
+
|
| 402 |
+
## 6. Plano de Implementação
|
| 403 |
+
|
| 404 |
+
### Fase 1 — Novos Módulos
|
| 405 |
+
1. `cnn_bigru/utils/ewc.py` — EWC completo (offline + online)
|
| 406 |
+
2. `cnn_bigru/models/context_window.py` — Context window com KV cache
|
| 407 |
+
3. `cnn_bigru/models/rope.py` — RoPE
|
| 408 |
+
4. `cnn_bigru/models/transformer_block.py` — CausalSelfAttention + TransformerBlock
|
| 409 |
+
|
| 410 |
+
### Fase 2 — Correções de Bugs Críticos
|
| 411 |
+
5. `cnn_bigru/models/cooperative_bigru.py` — Self-Attention final correta
|
| 412 |
+
6. `cnn_bigru/models/multimodal_model.py` — Fix dimensão d_text + weight tying
|
| 413 |
+
7. `cnn_bigru/training/hypothesis_controller.py` — Fix gate logic
|
| 414 |
+
8. `cnn_bigru/models/generator_verifier.py` — Fix enc_out repetition
|
| 415 |
+
9. `cnn_bigru/training/trainer.py` — Integrar verificador + hipótese + EWC + context window
|
| 416 |
+
10. `cnn_bigru/losses/losses.py` — Device-safe + EWC term
|
| 417 |
+
11. `cnn_bigru/training/auto_learner.py` — Buffer registration + l2_reg + try/finally
|
| 418 |
+
12. `cnn_bigru/inference/inference.py` — Geração autoregressiva real
|
| 419 |
+
13. `cnn_bigru/data/streaming_dataset.py` — Synthetic fallback reachability
|
| 420 |
+
14. `cnn_bigru/utils/semantic_embeddings.py` — Integrar no pipeline
|
| 421 |
+
|
| 422 |
+
### Fase 3 — Embalagem
|
| 423 |
+
15. `cnn_bigru/__init__.py` — Exports + versão
|
| 424 |
+
16. `cnn_bigru/scripts/push_to_hf.py` — `upload_folder` + filtro robusto
|
| 425 |
+
|
| 426 |
+
### Fase 4 — Validação
|
| 427 |
+
17. `cnn_bigru/tests/test_50_samples.py` — Testar todos os novos módulos
|
| 428 |
+
18. Executar teste, capturar erros, iterar
|
| 429 |
+
|
| 430 |
+
### Fase 5 — Publicação
|
| 431 |
+
19. Push para HF `CNN-BiGRU` (overwrite outdated)
|
| 432 |
+
20. Deletar `HF_TOKEN` do ambiente
|
| 433 |
+
|
| 434 |
+
---
|
| 435 |
+
|
| 436 |
+
## 7. Referências
|
| 437 |
+
|
| 438 |
+
- Kirkpatrick, J. et al. **"Overcoming catastrophic forgetting in neural networks"**. PNAS 2017.
|
| 439 |
+
- Xiao, G. et al. **"Efficient Streaming Language Models with Attention Sinks"**. 2023.
|
| 440 |
+
- Su, J. et al. **"RoFormer: Enhanced Transformer with Rotary Position Embedding"**. 2021.
|
| 441 |
+
- Vaswani, A. et al. **"Attention Is All You Need"**. NeurIPS 2017.
|
| 442 |
+
- `dados.txt` — pseudocódigo completo do projeto CNN-BiGRU cooperativo.
|
| 443 |
+
|
| 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.**
|
cnn_bigru/inference/__init__.py
ADDED
|
File without changes
|
cnn_bigru/inference/inference.py
ADDED
|
@@ -0,0 +1,450 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""inference.py — Inferência com Temperatura, Top-K, Top-P e Penalidades.
|
| 2 |
+
|
| 3 |
+
Implementa as técnicas descritas em dados.txt (seção sobre geração avançada):
|
| 4 |
+
1. Temperatura: ajusta entropia das previsões
|
| 5 |
+
2. Top-K: filtra os K tokens mais prováveis
|
| 6 |
+
3. Top-P (Nucleus Sampling): limite dinâmico por probabilidade acumulada
|
| 7 |
+
4. Presence Penalty: pune se já apareceu
|
| 8 |
+
5. Frequency Penalty: pune proporcionalmente à frequência
|
| 9 |
+
|
| 10 |
+
VERSÃO 2.0:
|
| 11 |
+
- Usa GeneratorCNNBiGRU (quando disponível) para geração autoregressiva REAL
|
| 12 |
+
- Integra ContextWindowManager para sequências longas (sliding window)
|
| 13 |
+
- Fallback para MultimodalCNNBiGRU.lm_head quando generator não disponível
|
| 14 |
+
- PPL evaluation correta usando GeneratorCNNBiGRU (não single-step broadcast)
|
| 15 |
+
"""
|
| 16 |
+
from __future__ import annotations
|
| 17 |
+
|
| 18 |
+
import logging
|
| 19 |
+
import math
|
| 20 |
+
from typing import Any, Dict, List, Optional, Tuple
|
| 21 |
+
|
| 22 |
+
import torch
|
| 23 |
+
import torch.nn as nn
|
| 24 |
+
import torch.nn.functional as F
|
| 25 |
+
|
| 26 |
+
from ..models.context_window import ContextWindowManager, ContextWindowConfig
|
| 27 |
+
from ..models.generator_verifier import GeneratorCNNBiGRU
|
| 28 |
+
from ..models.multimodal_model import MultimodalCNNBiGRU
|
| 29 |
+
|
| 30 |
+
logger = logging.getLogger(__name__)
|
| 31 |
+
|
| 32 |
+
|
| 33 |
+
# ============================================================================
|
| 34 |
+
# Aplicar penalidades, temperatura, top-k, top-p
|
| 35 |
+
# ============================================================================
|
| 36 |
+
|
| 37 |
+
def apply_sampling_filters(
|
| 38 |
+
logits: torch.Tensor,
|
| 39 |
+
token_counts: Dict[int, int],
|
| 40 |
+
temperature: float,
|
| 41 |
+
top_k: int,
|
| 42 |
+
top_p: float,
|
| 43 |
+
presence_penalty: float,
|
| 44 |
+
frequency_penalty: float,
|
| 45 |
+
) -> Tuple[torch.Tensor, Dict[str, int]]:
|
| 46 |
+
"""Aplica presença/frequência penalidades + temperatura + top-k + top-p.
|
| 47 |
+
|
| 48 |
+
Args:
|
| 49 |
+
logits: [B, V] (será modificado in-place via clone)
|
| 50 |
+
token_counts: dict {token_id: count} para penalidades
|
| 51 |
+
temperature: ajusta entropia (0 = greedy)
|
| 52 |
+
top_k: filtra K mais prováveis
|
| 53 |
+
top_p: limite de probabilidade acumulada
|
| 54 |
+
presence_penalty: pune token se já apareceu
|
| 55 |
+
frequency_penalty: pune proporcionalmente à frequência
|
| 56 |
+
|
| 57 |
+
Returns:
|
| 58 |
+
(filtered_logits, metrics_dict)
|
| 59 |
+
"""
|
| 60 |
+
metrics = {
|
| 61 |
+
"tokens_after_topk": 0,
|
| 62 |
+
"tokens_after_topp": 0,
|
| 63 |
+
}
|
| 64 |
+
logits = logits.clone()
|
| 65 |
+
|
| 66 |
+
# Penalidades de repetição (direct logit modification)
|
| 67 |
+
if presence_penalty != 0.0 or frequency_penalty != 0.0:
|
| 68 |
+
for tok_id, count in token_counts.items():
|
| 69 |
+
if 0 <= tok_id < logits.size(-1):
|
| 70 |
+
logits[:, tok_id] -= (presence_penalty + count * frequency_penalty)
|
| 71 |
+
|
| 72 |
+
# Temperatura
|
| 73 |
+
if temperature > 0.0:
|
| 74 |
+
logits = logits / temperature
|
| 75 |
+
# Se temperature == 0, chamador deve usar greedy (argmax)
|
| 76 |
+
|
| 77 |
+
# Top-K
|
| 78 |
+
if top_k > 0 and top_k < logits.size(-1):
|
| 79 |
+
k = min(top_k, logits.size(-1))
|
| 80 |
+
top_vals, _ = torch.topk(logits, k, dim=-1)
|
| 81 |
+
thresh = top_vals[:, -1:]
|
| 82 |
+
logits[logits < thresh] = float("-inf")
|
| 83 |
+
metrics["tokens_after_topk"] = k
|
| 84 |
+
|
| 85 |
+
# Top-P (Nucleus Sampling)
|
| 86 |
+
if top_p < 1.0:
|
| 87 |
+
sorted_logits, sorted_idx = torch.sort(logits, descending=True, dim=-1)
|
| 88 |
+
sorted_probs = F.softmax(sorted_logits, dim=-1)
|
| 89 |
+
cum_probs = torch.cumsum(sorted_probs, dim=-1)
|
| 90 |
+
|
| 91 |
+
# Remove tokens cuja probabilidade acumulada ultrapassa top_p
|
| 92 |
+
# (mantém o primeiro token mesmo se já ultrapassar)
|
| 93 |
+
remove = cum_probs > top_p
|
| 94 |
+
remove[:, 1:] = remove[:, :-1].clone()
|
| 95 |
+
remove[:, 0] = False
|
| 96 |
+
|
| 97 |
+
indices_to_remove = remove.scatter(1, sorted_idx, remove)
|
| 98 |
+
logits[indices_to_remove] = float("-inf")
|
| 99 |
+
metrics["tokens_after_topp"] = int((~indices_to_remove[0]).sum().item())
|
| 100 |
+
|
| 101 |
+
return logits, metrics
|
| 102 |
+
|
| 103 |
+
|
| 104 |
+
# ============================================================================
|
| 105 |
+
# Geração Autoregressiva (V2.0 — CORRIGIDA)
|
| 106 |
+
# ============================================================================
|
| 107 |
+
|
| 108 |
+
@torch.no_grad()
|
| 109 |
+
def generate_with_sampling(
|
| 110 |
+
model: nn.Module,
|
| 111 |
+
tokenizer,
|
| 112 |
+
prompt_a: str,
|
| 113 |
+
prompt_b: str = "",
|
| 114 |
+
images: Optional[torch.Tensor] = None,
|
| 115 |
+
audios: Optional[torch.Tensor] = None,
|
| 116 |
+
max_new_tokens: int = 32,
|
| 117 |
+
temperature: float = 0.7,
|
| 118 |
+
top_k: int = 40,
|
| 119 |
+
top_p: float = 0.9,
|
| 120 |
+
presence_penalty: float = 0.3,
|
| 121 |
+
frequency_penalty: float = 0.3,
|
| 122 |
+
device: str = "cpu",
|
| 123 |
+
generator: Optional[GeneratorCNNBiGRU] = None,
|
| 124 |
+
context_window: Optional[ContextWindowManager] = None,
|
| 125 |
+
) -> Dict[str, Any]:
|
| 126 |
+
"""Gera texto com amostragem avançada.
|
| 127 |
+
|
| 128 |
+
VERSÃO 2.0:
|
| 129 |
+
- Se `generator` (GeneratorCNNBiGRU) for fornecido, usa geração
|
| 130 |
+
autoregressiva REAL com o decoder GRU + Bahdanau attention.
|
| 131 |
+
- Caso contrário, fallback para MultimodalCNNBiGRU em modo "generate"
|
| 132 |
+
(single-step — NÃO é autoregressivo real, apenas para testes).
|
| 133 |
+
- Se `context_window` for fornecido, aplica sliding window para
|
| 134 |
+
suportar sequências longas (> max_seq_len do modelo).
|
| 135 |
+
|
| 136 |
+
Args:
|
| 137 |
+
model: MultimodalCNNBiGRU
|
| 138 |
+
tokenizer: BBPETokenizer
|
| 139 |
+
prompt_a, prompt_b: textos de entrada
|
| 140 |
+
images, audios: tensores opcionais
|
| 141 |
+
max_new_tokens: máximo de tokens a gerar
|
| 142 |
+
temperature: ajusta entropia (0 = greedy)
|
| 143 |
+
top_k: filtra K mais prováveis
|
| 144 |
+
top_p: limite de probabilidade acumulada
|
| 145 |
+
presence_penalty: pune token se já apareceu
|
| 146 |
+
frequency_penalty: pune proporcionalmente à frequência
|
| 147 |
+
device: dispositivo de computação
|
| 148 |
+
generator: GeneratorCNNBiGRU opcional (recomendado para geração real)
|
| 149 |
+
context_window: gerenciador de janela de contexto opcional
|
| 150 |
+
|
| 151 |
+
Returns:
|
| 152 |
+
dict com:
|
| 153 |
+
text: texto gerado
|
| 154 |
+
token_ids: lista de IDs gerados
|
| 155 |
+
metrics: dicionário de métricas
|
| 156 |
+
used_generator: bool indicando se usou generator real
|
| 157 |
+
"""
|
| 158 |
+
model.eval()
|
| 159 |
+
if generator is not None:
|
| 160 |
+
generator.eval()
|
| 161 |
+
|
| 162 |
+
# Tokeniza prompts
|
| 163 |
+
ids_a = tokenizer.encode(prompt_a, add_special=True)
|
| 164 |
+
ids_b = tokenizer.encode(prompt_b, add_special=True) if prompt_b else [tokenizer.pad_id] * len(ids_a)
|
| 165 |
+
|
| 166 |
+
# Pad para mesmo comprimento
|
| 167 |
+
max_len = max(len(ids_a), len(ids_b), 8)
|
| 168 |
+
pad_id = tokenizer.pad_id
|
| 169 |
+
ids_a = ids_a + [pad_id] * (max_len - len(ids_a))
|
| 170 |
+
ids_b = ids_b + [pad_id] * (max_len - len(ids_b))
|
| 171 |
+
|
| 172 |
+
input_a = torch.tensor([ids_a], dtype=torch.long, device=device)
|
| 173 |
+
input_b = torch.tensor([ids_b], dtype=torch.long, device=device)
|
| 174 |
+
|
| 175 |
+
# Prepara imagens/áudios
|
| 176 |
+
if images is not None:
|
| 177 |
+
images = images.unsqueeze(0).to(device) if images.dim() == 3 else images.to(device)
|
| 178 |
+
if audios is not None:
|
| 179 |
+
audios = audios.unsqueeze(0).to(device) if audios.dim() == 3 else audios.to(device)
|
| 180 |
+
|
| 181 |
+
# Histórico de contagem para penalidades
|
| 182 |
+
token_counts: Dict[int, int] = {}
|
| 183 |
+
generated_ids: List[int] = []
|
| 184 |
+
|
| 185 |
+
eos_id = tokenizer.eos_id
|
| 186 |
+
bos_id = tokenizer.bos_id
|
| 187 |
+
|
| 188 |
+
metrics = {
|
| 189 |
+
"steps": 0,
|
| 190 |
+
"fallback_to_greedy": 0,
|
| 191 |
+
"tokens_after_topk": 0,
|
| 192 |
+
"tokens_after_topp": 0,
|
| 193 |
+
"used_generator": generator is not None,
|
| 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
|
| 210 |
+
|
| 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:
|
| 221 |
+
# GERAÇÃO AUTOREGRESSIVA REAL via GeneratorCNNBiGRU
|
| 222 |
+
# O generator tem decoder GRU + Bahdanau attention
|
| 223 |
+
# Passamos target_proof=None para que ele gere livremente
|
| 224 |
+
# mas limitamos max_len ao tamanho atual + 1
|
| 225 |
+
gen_out = generator(
|
| 226 |
+
input_a, input_b,
|
| 227 |
+
target_proof=None,
|
| 228 |
+
bos_id=bos_id,
|
| 229 |
+
eos_id=eos_id,
|
| 230 |
+
max_len=1, # gera 1 token por vez
|
| 231 |
+
)
|
| 232 |
+
logits = gen_out["logits"] # [B, 1, V]
|
| 233 |
+
logits = logits[:, -1, :] # [1, V]
|
| 234 |
+
else:
|
| 235 |
+
# FALLBACK: MultimodalCNNBiGRU em modo "generate" (single-step)
|
| 236 |
+
out = model(
|
| 237 |
+
input_a, input_b,
|
| 238 |
+
images=images, audios=audios,
|
| 239 |
+
mode="generate",
|
| 240 |
+
)
|
| 241 |
+
logits = out["logits"] # [B, V] ou [B, T, V]
|
| 242 |
+
if logits.dim() == 3:
|
| 243 |
+
logits = logits[:, -1, :]
|
| 244 |
+
logits = logits.clone() # [1, V]
|
| 245 |
+
except Exception as e:
|
| 246 |
+
logger.warning(f"Forward falhou no step {step}: {e}")
|
| 247 |
+
break
|
| 248 |
+
|
| 249 |
+
# --- Greedy (temperature == 0) ---
|
| 250 |
+
if temperature <= 0.0:
|
| 251 |
+
next_token = torch.argmax(logits, dim=-1, keepdim=True)
|
| 252 |
+
tok_id = next_token.item()
|
| 253 |
+
if tok_id == eos_id:
|
| 254 |
+
break
|
| 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 ---
|
| 264 |
+
logits, filter_metrics = apply_sampling_filters(
|
| 265 |
+
logits, token_counts,
|
| 266 |
+
temperature=temperature,
|
| 267 |
+
top_k=top_k,
|
| 268 |
+
top_p=top_p,
|
| 269 |
+
presence_penalty=presence_penalty,
|
| 270 |
+
frequency_penalty=frequency_penalty,
|
| 271 |
+
)
|
| 272 |
+
metrics["tokens_after_topk"] = filter_metrics["tokens_after_topk"]
|
| 273 |
+
metrics["tokens_after_topp"] = filter_metrics["tokens_after_topp"]
|
| 274 |
+
|
| 275 |
+
# --- Amostragem Multinomial ---
|
| 276 |
+
# Substituir -inf por um valor muito baixo (evita NaN em softmax)
|
| 277 |
+
logits = torch.where(torch.isinf(logits), torch.full_like(logits, -1e9), logits)
|
| 278 |
+
probs = F.softmax(logits, dim=-1)
|
| 279 |
+
# Lidar com probs NaN (todas -inf)
|
| 280 |
+
if torch.isnan(probs).any() or probs.sum() == 0:
|
| 281 |
+
logger.warning(f"NaN/zero probs no step {step}; usando argmax")
|
| 282 |
+
next_token = torch.argmax(logits, dim=-1, keepdim=True)
|
| 283 |
+
else:
|
| 284 |
+
try:
|
| 285 |
+
next_token = torch.multinomial(probs, num_samples=1)
|
| 286 |
+
except Exception as e:
|
| 287 |
+
logger.warning(f"Multinomial falhou: {e}; usando argmax")
|
| 288 |
+
next_token = torch.argmax(probs, dim=-1, keepdim=True)
|
| 289 |
+
|
| 290 |
+
tok_id = next_token.item()
|
| 291 |
+
if tok_id == eos_id:
|
| 292 |
+
break
|
| 293 |
+
|
| 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)
|
| 303 |
+
|
| 304 |
+
return {
|
| 305 |
+
"text": text,
|
| 306 |
+
"token_ids": generated_ids,
|
| 307 |
+
"metrics": metrics,
|
| 308 |
+
"used_generator": use_generator,
|
| 309 |
+
}
|
| 310 |
+
|
| 311 |
+
|
| 312 |
+
# ============================================================================
|
| 313 |
+
# Avaliação de Perplexidade (V2.0 — CORRIGIDA)
|
| 314 |
+
# ============================================================================
|
| 315 |
+
|
| 316 |
+
@torch.no_grad()
|
| 317 |
+
def evaluate_perplexity(
|
| 318 |
+
model: nn.Module,
|
| 319 |
+
dataloader,
|
| 320 |
+
tokenizer,
|
| 321 |
+
device: str = "cpu",
|
| 322 |
+
max_batches: int = 10,
|
| 323 |
+
generator: Optional[GeneratorCNNBiGRU] = None,
|
| 324 |
+
) -> Dict[str, float]:
|
| 325 |
+
"""Avalia perplexidade do modelo em um dataloader.
|
| 326 |
+
|
| 327 |
+
PPL = exp(loss_media)
|
| 328 |
+
|
| 329 |
+
VERSÃO 2.0:
|
| 330 |
+
- Se `generator` for fornecido, usa o generator com teacher forcing
|
| 331 |
+
para gerar logits autoregressivos REAIS por timestep.
|
| 332 |
+
- Caso contrário, faz fallback para MultimodalCNNBiGRU (single-step
|
| 333 |
+
broadcast — métrica aproximada).
|
| 334 |
+
"""
|
| 335 |
+
model.eval()
|
| 336 |
+
if generator is not None:
|
| 337 |
+
generator.eval()
|
| 338 |
+
|
| 339 |
+
total_loss = 0.0
|
| 340 |
+
total_tokens = 0
|
| 341 |
+
pad_id = tokenizer.pad_id
|
| 342 |
+
bos_id = tokenizer.bos_id
|
| 343 |
+
eos_id = tokenizer.eos_id
|
| 344 |
+
|
| 345 |
+
n_batches = 0
|
| 346 |
+
for batch_idx, batch in enumerate(dataloader):
|
| 347 |
+
if batch_idx >= max_batches:
|
| 348 |
+
break
|
| 349 |
+
|
| 350 |
+
input_ids_a = batch["input_ids_a"].to(device)
|
| 351 |
+
input_ids_b = batch["input_ids_b"].to(device)
|
| 352 |
+
attn_mask_a = batch["attn_mask_a"].to(device)
|
| 353 |
+
attn_mask_b = batch["attn_mask_b"].to(device)
|
| 354 |
+
|
| 355 |
+
try:
|
| 356 |
+
if generator is not None:
|
| 357 |
+
# USAR GENERATOR: geração autoregressiva com teacher forcing
|
| 358 |
+
# target_proof = input_ids_b → logits por timestep
|
| 359 |
+
gen_out = generator(
|
| 360 |
+
input_ids_a, input_ids_b,
|
| 361 |
+
target_proof=input_ids_b, # teacher forcing
|
| 362 |
+
bos_id=bos_id,
|
| 363 |
+
eos_id=eos_id,
|
| 364 |
+
)
|
| 365 |
+
logits = gen_out["logits"] # [B, T, V]
|
| 366 |
+
else:
|
| 367 |
+
# FALLBACK: single-step broadcast (NÃO É autoregressivo real)
|
| 368 |
+
# Apenas para compatibilidade
|
| 369 |
+
images = batch.get("images")
|
| 370 |
+
audios = batch.get("audios")
|
| 371 |
+
if images is not None:
|
| 372 |
+
images = images.to(device)
|
| 373 |
+
if audios is not None:
|
| 374 |
+
audios = audios.to(device)
|
| 375 |
+
out = model(
|
| 376 |
+
input_ids_a, input_ids_b,
|
| 377 |
+
images=images, audios=audios,
|
| 378 |
+
attn_mask_a=attn_mask_a, attn_mask_b=attn_mask_b,
|
| 379 |
+
mode="generate",
|
| 380 |
+
)
|
| 381 |
+
# out["logits"] é [B, V] (single-step) → expandir para [B, T, V]
|
| 382 |
+
B = out["logits"].size(0)
|
| 383 |
+
T = input_ids_b.size(1)
|
| 384 |
+
V = out["logits"].size(-1)
|
| 385 |
+
logits = out["logits"].unsqueeze(1).expand(-1, T, -1)
|
| 386 |
+
|
| 387 |
+
# Calcular loss por token
|
| 388 |
+
B, T, V = logits.shape
|
| 389 |
+
loss = F.cross_entropy(
|
| 390 |
+
logits.reshape(B * T, V),
|
| 391 |
+
input_ids_b.reshape(B * T),
|
| 392 |
+
ignore_index=pad_id,
|
| 393 |
+
reduction="sum",
|
| 394 |
+
)
|
| 395 |
+
n_tokens = (input_ids_b != pad_id).sum().item()
|
| 396 |
+
total_loss += loss.item()
|
| 397 |
+
total_tokens += n_tokens
|
| 398 |
+
n_batches += 1
|
| 399 |
+
except Exception as e:
|
| 400 |
+
logger.warning(f"PPL batch {batch_idx} falhou: {e}")
|
| 401 |
+
continue
|
| 402 |
+
|
| 403 |
+
if total_tokens == 0:
|
| 404 |
+
return {"loss": float("inf"), "ppl": float("inf"), "n_batches": 0}
|
| 405 |
+
|
| 406 |
+
avg_loss = total_loss / total_tokens
|
| 407 |
+
try:
|
| 408 |
+
ppl = math.exp(min(avg_loss, 20))
|
| 409 |
+
except OverflowError:
|
| 410 |
+
ppl = float("inf")
|
| 411 |
+
|
| 412 |
+
return {
|
| 413 |
+
"loss": avg_loss,
|
| 414 |
+
"ppl": ppl,
|
| 415 |
+
"n_batches": n_batches,
|
| 416 |
+
"used_generator": generator is not None,
|
| 417 |
+
}
|
| 418 |
+
|
| 419 |
+
|
| 420 |
+
# ============================================================================
|
| 421 |
+
# Helper: criar context window padrão
|
| 422 |
+
# ============================================================================
|
| 423 |
+
|
| 424 |
+
def make_default_context_window(
|
| 425 |
+
max_window: int = 256,
|
| 426 |
+
n_sink: int = 4,
|
| 427 |
+
embed_dim: int = 64,
|
| 428 |
+
n_heads: int = 4,
|
| 429 |
+
n_layers: int = 2,
|
| 430 |
+
device: str = "cpu",
|
| 431 |
+
) -> ContextWindowManager:
|
| 432 |
+
"""Factory para criar um ContextWindowManager com defaults sensatos."""
|
| 433 |
+
config = ContextWindowConfig(
|
| 434 |
+
max_window=max_window,
|
| 435 |
+
eviction_strategy="sink_sliding",
|
| 436 |
+
n_sink_tokens=n_sink,
|
| 437 |
+
embed_dim=embed_dim,
|
| 438 |
+
n_heads=n_heads,
|
| 439 |
+
n_layers=n_layers,
|
| 440 |
+
device=device,
|
| 441 |
+
)
|
| 442 |
+
return ContextWindowManager(config)
|
| 443 |
+
|
| 444 |
+
|
| 445 |
+
__all__ = [
|
| 446 |
+
"apply_sampling_filters",
|
| 447 |
+
"generate_with_sampling",
|
| 448 |
+
"evaluate_perplexity",
|
| 449 |
+
"make_default_context_window",
|
| 450 |
+
]
|
cnn_bigru/losses/__init__.py
ADDED
|
File without changes
|
cnn_bigru/losses/losses.py
ADDED
|
@@ -0,0 +1,271 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""losses.py — Sistema de múltiplas perdas para o CNN-BiGRU cooperativo.
|
| 2 |
+
|
| 3 |
+
Implementa as perdas descritas na seção 8 de dados.txt:
|
| 4 |
+
|
| 5 |
+
L_G = CrossEntropy(logits_prova, prova_alvo)
|
| 6 |
+
L_penal = λ * Σ (1 - v_t) (penalidade linear)
|
| 7 |
+
L_exp_penal = μ * Σ exp(γ * (1 - v_t)) (penalidade exponencial)
|
| 8 |
+
L_V = BCE(v_t, rotulo_real) (perda do verificador)
|
| 9 |
+
L_AH = Σ (1 - τ_t) (perda anti-alucinação)
|
| 10 |
+
L_reg = L2 + estimativa de curvatura (regularização)
|
| 11 |
+
L_total = α * L_G_total + β * L_V + γ_loss * L_AH + δ * L_reg
|
| 12 |
+
|
| 13 |
+
onde:
|
| 14 |
+
L_G_total = L_G + λ_penal * L_penal + μ_exp * L_exp_penal
|
| 15 |
+
"""
|
| 16 |
+
from __future__ import annotations
|
| 17 |
+
|
| 18 |
+
import logging
|
| 19 |
+
import math
|
| 20 |
+
from dataclasses import dataclass, field
|
| 21 |
+
from typing import Dict, List, Optional
|
| 22 |
+
|
| 23 |
+
import torch
|
| 24 |
+
import torch.nn as nn
|
| 25 |
+
import torch.nn.functional as F
|
| 26 |
+
|
| 27 |
+
logger = logging.getLogger(__name__)
|
| 28 |
+
|
| 29 |
+
|
| 30 |
+
@dataclass
|
| 31 |
+
class LossConfig:
|
| 32 |
+
"""Configuração das perdas e pesos."""
|
| 33 |
+
# Pesos das perdas (ALPHA, BETA, GAMMA, DELTA do pseudocódigo)
|
| 34 |
+
alpha: float = 1.0 # peso de L_G_total
|
| 35 |
+
beta: float = 0.5 # peso de L_V
|
| 36 |
+
gamma_loss: float = 0.3 # peso de L_AH
|
| 37 |
+
delta: float = 0.01 # peso de L_reg
|
| 38 |
+
|
| 39 |
+
# Penalidades
|
| 40 |
+
lambda_penal: float = 0.1 # LAMBDA_PENAL
|
| 41 |
+
mu_exp_penal: float = 0.05 # MU_EXP_PENAL
|
| 42 |
+
gamma_exp: float = 1.0 # GAMMA_EXP
|
| 43 |
+
threshold_err: float = 0.5 # THRESHOLD_ERR para penalidade exponencial
|
| 44 |
+
|
| 45 |
+
# Regularização
|
| 46 |
+
l2_reg: float = 1e-5
|
| 47 |
+
|
| 48 |
+
# Curvatura
|
| 49 |
+
use_curvature: bool = True
|
| 50 |
+
curvature_eps: float = 1e-3
|
| 51 |
+
|
| 52 |
+
|
| 53 |
+
class MultiLoss(nn.Module):
|
| 54 |
+
"""Computa todas as perdas do sistema cooperativo.
|
| 55 |
+
|
| 56 |
+
Uso:
|
| 57 |
+
loss_fn = MultiLoss(cfg)
|
| 58 |
+
result = loss_fn(logits, targets, v_list, tau_list, generator, verifier)
|
| 59 |
+
loss = result['total']
|
| 60 |
+
loss.backward()
|
| 61 |
+
"""
|
| 62 |
+
|
| 63 |
+
def __init__(self, config: LossConfig, pad_idx: int = 1):
|
| 64 |
+
super().__init__()
|
| 65 |
+
self.cfg = config
|
| 66 |
+
self.pad_idx = pad_idx
|
| 67 |
+
self.ce = nn.CrossEntropyLoss(ignore_index=pad_idx, reduction='mean')
|
| 68 |
+
self.bce = nn.BCELoss(reduction='mean')
|
| 69 |
+
|
| 70 |
+
def generation_loss(
|
| 71 |
+
self,
|
| 72 |
+
logits: torch.Tensor, # [B, T, V]
|
| 73 |
+
targets: torch.Tensor, # [B, T]
|
| 74 |
+
) -> torch.Tensor:
|
| 75 |
+
"""L_G: CrossEntropy ignorando padding."""
|
| 76 |
+
B, T, V = logits.shape
|
| 77 |
+
return self.ce(logits.reshape(B * T, V), targets.reshape(B * T))
|
| 78 |
+
|
| 79 |
+
def penalty_loss(
|
| 80 |
+
self,
|
| 81 |
+
v_list: List[torch.Tensor], # lista de [B, 1] probabilidades do verificador
|
| 82 |
+
) -> Dict[str, torch.Tensor]:
|
| 83 |
+
"""Calcula L_penal (linear) e L_exp_penal (exponencial).
|
| 84 |
+
|
| 85 |
+
Fórmulas (dados.txt seção 8.2):
|
| 86 |
+
L_penal = Σ_t (1 - v_t) (linear, todos os passos)
|
| 87 |
+
L_exp_penal = Σ_{t: v_t < threshold} exp(γ * (1 - v_t)) (exponencial, apenas graves)
|
| 88 |
+
"""
|
| 89 |
+
if not v_list:
|
| 90 |
+
# Retorna zero no device do próprio MultiLoss (que está em algum dispositivo)
|
| 91 |
+
device = next(self.parameters(), torch.tensor(0.0)).device if any(
|
| 92 |
+
p.requires_grad for p in self.parameters()
|
| 93 |
+
) else torch.device("cpu")
|
| 94 |
+
zero = torch.zeros((), device=device)
|
| 95 |
+
return {"linear": zero, "exp": zero}
|
| 96 |
+
|
| 97 |
+
# Stack: [T, B, 1] -> [T, B]
|
| 98 |
+
v_stack = torch.stack([v.squeeze(-1) if v.dim() == 2 else v for v in v_list], dim=0)
|
| 99 |
+
# Penalidade linear: Σ (1 - v) sobre TODOS os passos
|
| 100 |
+
l_penal = (1.0 - v_stack).sum() / max(v_stack.numel(), 1)
|
| 101 |
+
|
| 102 |
+
# Penalidade exponencial: Σ exp(γ * (1 - v)) onde v < threshold
|
| 103 |
+
mask = v_stack < self.cfg.threshold_err
|
| 104 |
+
if mask.any():
|
| 105 |
+
# exp(γ * (1 - v)) para os graves
|
| 106 |
+
exp_vals = torch.exp(self.cfg.gamma_exp * (1.0 - v_stack[mask]))
|
| 107 |
+
l_exp = exp_vals.sum() / max(v_stack.numel(), 1)
|
| 108 |
+
else:
|
| 109 |
+
l_exp = torch.zeros((), device=v_stack.device)
|
| 110 |
+
|
| 111 |
+
return {"linear": l_penal, "exp": l_exp}
|
| 112 |
+
|
| 113 |
+
def verifier_loss(
|
| 114 |
+
self,
|
| 115 |
+
v_list: List[torch.Tensor],
|
| 116 |
+
labels: List[torch.Tensor], # lista de [B, 1] com 0/1
|
| 117 |
+
) -> torch.Tensor:
|
| 118 |
+
"""L_V: BCE entre v_t e rotulo_real."""
|
| 119 |
+
if not v_list:
|
| 120 |
+
# Device-safe zero
|
| 121 |
+
device = next(self.parameters(), torch.tensor(0.0)).device if any(
|
| 122 |
+
p.requires_grad for p in self.parameters()
|
| 123 |
+
) else torch.device("cpu")
|
| 124 |
+
return torch.zeros((), device=device)
|
| 125 |
+
# Inicializar no device do primeiro v
|
| 126 |
+
first_v = v_list[0]
|
| 127 |
+
device = first_v.device
|
| 128 |
+
total = torch.zeros((), device=device)
|
| 129 |
+
for v, label in zip(v_list, labels):
|
| 130 |
+
v_flat = (v.squeeze(-1) if v.dim() == 2 else v).clamp(1e-7, 1 - 1e-7)
|
| 131 |
+
label_flat = label.squeeze(-1) if label.dim() == 2 else label
|
| 132 |
+
total = total + self.bce(v_flat, label_flat.float().to(device))
|
| 133 |
+
return total / max(len(v_list), 1)
|
| 134 |
+
|
| 135 |
+
def anti_hallucination_loss(
|
| 136 |
+
self,
|
| 137 |
+
tau_list: List[torch.Tensor], # lista de [B, 1] valores de verdade suave
|
| 138 |
+
) -> torch.Tensor:
|
| 139 |
+
"""L_AH = Σ (1 - τ_t)."""
|
| 140 |
+
if not tau_list:
|
| 141 |
+
device = next(self.parameters(), torch.tensor(0.0)).device if any(
|
| 142 |
+
p.requires_grad for p in self.parameters()
|
| 143 |
+
) else torch.device("cpu")
|
| 144 |
+
return torch.zeros((), device=device)
|
| 145 |
+
first_tau = tau_list[0]
|
| 146 |
+
device = first_tau.device
|
| 147 |
+
total = torch.zeros((), device=device)
|
| 148 |
+
for tau in tau_list:
|
| 149 |
+
tau_flat = tau.squeeze(-1) if tau.dim() == 2 else tau
|
| 150 |
+
total = total + (1.0 - tau_flat).mean()
|
| 151 |
+
return total / max(len(tau_list), 1)
|
| 152 |
+
|
| 153 |
+
def regularization_loss(
|
| 154 |
+
self,
|
| 155 |
+
models: List[nn.Module],
|
| 156 |
+
) -> torch.Tensor:
|
| 157 |
+
"""L_reg = L2 + estimativa de curvatura (aproximada via norma de pesos)."""
|
| 158 |
+
# Determinar device do primeiro parâmetro treinável
|
| 159 |
+
device = torch.device("cpu")
|
| 160 |
+
for model in models:
|
| 161 |
+
for p in model.parameters():
|
| 162 |
+
if p.requires_grad:
|
| 163 |
+
device = p.device
|
| 164 |
+
break
|
| 165 |
+
else:
|
| 166 |
+
continue
|
| 167 |
+
break
|
| 168 |
+
|
| 169 |
+
l2 = torch.zeros((), device=device)
|
| 170 |
+
for model in models:
|
| 171 |
+
for p in model.parameters():
|
| 172 |
+
if p.requires_grad:
|
| 173 |
+
l2 = l2 + p.pow(2).sum()
|
| 174 |
+
l2 = self.cfg.l2_reg * l2
|
| 175 |
+
|
| 176 |
+
# Estimativa de curvatura aproximada: variância das normas dos gradientes
|
| 177 |
+
# (calculada externamente e passada via attr)
|
| 178 |
+
curv = getattr(self, '_curvature_estimate', None)
|
| 179 |
+
if curv is None:
|
| 180 |
+
curv = torch.zeros((), device=device)
|
| 181 |
+
elif isinstance(curv, (int, float)):
|
| 182 |
+
curv = torch.tensor(float(curv), device=device)
|
| 183 |
+
elif isinstance(curv, torch.Tensor):
|
| 184 |
+
curv = curv.to(device)
|
| 185 |
+
return l2 + curv
|
| 186 |
+
|
| 187 |
+
def set_curvature_estimate(self, value: torch.Tensor) -> None:
|
| 188 |
+
"""Define a estimativa de curvatura (calculada externamente)."""
|
| 189 |
+
self._curvature_estimate = value
|
| 190 |
+
|
| 191 |
+
def forward(
|
| 192 |
+
self,
|
| 193 |
+
logits: torch.Tensor,
|
| 194 |
+
targets: torch.Tensor,
|
| 195 |
+
v_list: List[torch.Tensor],
|
| 196 |
+
tau_list: List[torch.Tensor],
|
| 197 |
+
verifier_labels: Optional[List[torch.Tensor]] = None,
|
| 198 |
+
models_for_reg: Optional[List[nn.Module]] = None,
|
| 199 |
+
ewc_penalty: Optional[torch.Tensor] = None,
|
| 200 |
+
) -> Dict[str, torch.Tensor]:
|
| 201 |
+
"""Computa todas as perdas e combina.
|
| 202 |
+
|
| 203 |
+
Args:
|
| 204 |
+
logits: [B, T, V] logits do gerador
|
| 205 |
+
targets: [B, T] tokens alvo
|
| 206 |
+
v_list: lista de [B, 1] probabilidades do verificador por passo
|
| 207 |
+
tau_list: lista de [B, 1] valores anti-alucinação por passo
|
| 208 |
+
verifier_labels: lista de [B, 1] rótulos reais (0/1) por passo
|
| 209 |
+
models_for_reg: lista de modelos para regularização L2
|
| 210 |
+
ewc_penalty: escalar (tensor 0-dim) com a penalidade EWC pré-computada
|
| 211 |
+
|
| 212 |
+
Returns:
|
| 213 |
+
dict com todas as perdas individuais e a perda total
|
| 214 |
+
"""
|
| 215 |
+
# L_G
|
| 216 |
+
l_g = self.generation_loss(logits, targets)
|
| 217 |
+
|
| 218 |
+
# Penalidades
|
| 219 |
+
pens = self.penalty_loss(v_list)
|
| 220 |
+
l_g_total = l_g + self.cfg.lambda_penal * pens["linear"] + self.cfg.mu_exp_penal * pens["exp"]
|
| 221 |
+
|
| 222 |
+
# L_V
|
| 223 |
+
if verifier_labels is not None:
|
| 224 |
+
l_v = self.verifier_loss(v_list, verifier_labels)
|
| 225 |
+
else:
|
| 226 |
+
# Sem rótulos reais: assume que v deveria ser 1 (passo correto)
|
| 227 |
+
labels_ones = [torch.ones_like(v) for v in v_list]
|
| 228 |
+
l_v = self.verifier_loss(v_list, labels_ones)
|
| 229 |
+
|
| 230 |
+
# L_AH
|
| 231 |
+
l_ah = self.anti_hallucination_loss(tau_list)
|
| 232 |
+
|
| 233 |
+
# L_reg (L2 + curvatura)
|
| 234 |
+
if models_for_reg:
|
| 235 |
+
l_reg = self.regularization_loss(models_for_reg)
|
| 236 |
+
else:
|
| 237 |
+
# Device-safe zero
|
| 238 |
+
device = logits.device if isinstance(logits, torch.Tensor) else torch.device("cpu")
|
| 239 |
+
l_reg = torch.zeros((), device=device)
|
| 240 |
+
|
| 241 |
+
# L_EWC (Elastic Weight Consolidation) — continua aprendizado sem esquecer
|
| 242 |
+
if ewc_penalty is not None:
|
| 243 |
+
l_ewc = ewc_penalty
|
| 244 |
+
else:
|
| 245 |
+
device = logits.device if isinstance(logits, torch.Tensor) else torch.device("cpu")
|
| 246 |
+
l_ewc = torch.zeros((), device=device)
|
| 247 |
+
|
| 248 |
+
# Combinação: L_total = α*L_G_total + β*L_V + γ*L_AH + δ*L_reg + L_EWC
|
| 249 |
+
# EWC é somado separadamente (já vem pré-ponderado por lambda_ewc/2)
|
| 250 |
+
total = (
|
| 251 |
+
self.cfg.alpha * l_g_total
|
| 252 |
+
+ self.cfg.beta * l_v
|
| 253 |
+
+ self.cfg.gamma_loss * l_ah
|
| 254 |
+
+ self.cfg.delta * l_reg
|
| 255 |
+
+ l_ewc
|
| 256 |
+
)
|
| 257 |
+
|
| 258 |
+
return {
|
| 259 |
+
"total": total,
|
| 260 |
+
"l_g": l_g,
|
| 261 |
+
"l_g_total": l_g_total,
|
| 262 |
+
"l_penal_linear": pens["linear"],
|
| 263 |
+
"l_penal_exp": pens["exp"],
|
| 264 |
+
"l_v": l_v,
|
| 265 |
+
"l_ah": l_ah,
|
| 266 |
+
"l_reg": l_reg,
|
| 267 |
+
"l_ewc": l_ewc,
|
| 268 |
+
}
|
| 269 |
+
|
| 270 |
+
|
| 271 |
+
__all__ = ["LossConfig", "MultiLoss"]
|
cnn_bigru/models/__init__.py
ADDED
|
File without changes
|
cnn_bigru/models/context_window.py
ADDED
|
@@ -0,0 +1,593 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
Context Window — Janela deslizante com cache KV para sequências longas.
|
| 3 |
+
|
| 4 |
+
Implementa três estratégias (ver docs/MATH_ANALYSIS.md seção 3.4):
|
| 5 |
+
- Sliding Window Pura
|
| 6 |
+
- Sink + Sliding (StreamingLLM) — RECOMENDADO
|
| 7 |
+
- Attention Recomputation (sem cache, máxima precisão)
|
| 8 |
+
|
| 9 |
+
Para CNN-BiGRU cooperativo, o contexto é gerenciado em modo:
|
| 10 |
+
- Encoder (treino/classificação): processa janela inteira, BiGRU bidirecional
|
| 11 |
+
- Decoder (geração): cache KV para CausalSelfAttention + GRU unidirecional
|
| 12 |
+
|
| 13 |
+
Autor: CNN-BiGRU Project
|
| 14 |
+
"""
|
| 15 |
+
from __future__ import annotations
|
| 16 |
+
|
| 17 |
+
import logging
|
| 18 |
+
from dataclasses import dataclass, field
|
| 19 |
+
from typing import Dict, List, Optional, Tuple
|
| 20 |
+
|
| 21 |
+
import torch
|
| 22 |
+
import torch.nn as nn
|
| 23 |
+
|
| 24 |
+
logger = logging.getLogger(__name__)
|
| 25 |
+
|
| 26 |
+
|
| 27 |
+
# ============================================================================
|
| 28 |
+
# Configuração
|
| 29 |
+
# ============================================================================
|
| 30 |
+
|
| 31 |
+
@dataclass
|
| 32 |
+
class ContextWindowConfig:
|
| 33 |
+
"""Configuração da janela de contexto."""
|
| 34 |
+
# Comprimento máximo da janela (em tokens)
|
| 35 |
+
max_window: int = 512
|
| 36 |
+
# Estratégia de evicção: "sliding" | "sink_sliding" | "recompute"
|
| 37 |
+
eviction_strategy: str = "sink_sliding"
|
| 38 |
+
# Número de tokens "sink" (apenas para sink_sliding) — tipicamente BOS + system prompt
|
| 39 |
+
n_sink_tokens: int = 4
|
| 40 |
+
# Dimensão do embedding (para alocar cache KV)
|
| 41 |
+
embed_dim: int = 256
|
| 42 |
+
# Número de cabeças (para cache KV em multi-head attention)
|
| 43 |
+
n_heads: int = 4
|
| 44 |
+
# Dimensão por cabeça
|
| 45 |
+
head_dim: Optional[int] = None # default: embed_dim // n_heads
|
| 46 |
+
# Número de camadas (para cache KV multicamada)
|
| 47 |
+
n_layers: int = 2
|
| 48 |
+
# Device padrão
|
| 49 |
+
device: str = "cpu"
|
| 50 |
+
# Dtype do cache (None = manter fp32)
|
| 51 |
+
dtype: Optional[torch.dtype] = None
|
| 52 |
+
|
| 53 |
+
|
| 54 |
+
# ============================================================================
|
| 55 |
+
# KV Cache
|
| 56 |
+
# ============================================================================
|
| 57 |
+
|
| 58 |
+
class KVCache:
|
| 59 |
+
"""
|
| 60 |
+
Cache de pares (Key, Value) para CausalSelfAttention.
|
| 61 |
+
|
| 62 |
+
Shape por camada:
|
| 63 |
+
K: [n_heads, seq_cached, head_dim]
|
| 64 |
+
V: [n_heads, seq_cached, head_dim]
|
| 65 |
+
|
| 66 |
+
Em modo batch:
|
| 67 |
+
K: [batch, n_heads, seq_cached, head_dim]
|
| 68 |
+
V: [batch, n_heads, seq_cached, head_dim]
|
| 69 |
+
"""
|
| 70 |
+
|
| 71 |
+
def __init__(self, n_layers: int, batch_size: int, n_heads: int,
|
| 72 |
+
head_dim: int, device: torch.device, dtype: torch.dtype):
|
| 73 |
+
self.n_layers = n_layers
|
| 74 |
+
self.batch_size = batch_size
|
| 75 |
+
self.n_heads = n_heads
|
| 76 |
+
self.head_dim = head_dim
|
| 77 |
+
self.device = device
|
| 78 |
+
self.dtype = dtype
|
| 79 |
+
# Pre-alocar listas vazias; preencher lazy no primeiro update
|
| 80 |
+
self.keys: List[Optional[torch.Tensor]] = [None] * n_layers
|
| 81 |
+
self.values: List[Optional[torch.Tensor]] = [None] * n_layers
|
| 82 |
+
|
| 83 |
+
def update(self, layer_idx: int, new_k: torch.Tensor, new_v: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
| 84 |
+
"""
|
| 85 |
+
Atualiza o cache da camada layer_idx com novos K, V.
|
| 86 |
+
|
| 87 |
+
Args:
|
| 88 |
+
layer_idx: índice da camada [0, n_layers)
|
| 89 |
+
new_k: [batch, n_heads, seq_new, head_dim]
|
| 90 |
+
new_v: [batch, n_heads, seq_new, head_dim]
|
| 91 |
+
|
| 92 |
+
Returns:
|
| 93 |
+
(cached_k, cached_v): [batch, n_heads, seq_total, head_dim]
|
| 94 |
+
"""
|
| 95 |
+
if not (0 <= layer_idx < self.n_layers):
|
| 96 |
+
raise IndexError(f"layer_idx {layer_idx} fora de range [0, {self.n_layers})")
|
| 97 |
+
|
| 98 |
+
# Validar shapes
|
| 99 |
+
if new_k.dim() != 4 or new_v.dim() != 4:
|
| 100 |
+
raise ValueError(
|
| 101 |
+
f"new_k e new_v devem ser 4D [batch, n_heads, seq, head_dim]; "
|
| 102 |
+
f"recebido {new_k.shape}"
|
| 103 |
+
)
|
| 104 |
+
|
| 105 |
+
# Converter dtype/device
|
| 106 |
+
new_k = new_k.to(device=self.device, dtype=self.dtype)
|
| 107 |
+
new_v = new_v.to(device=self.device, dtype=self.dtype)
|
| 108 |
+
|
| 109 |
+
if self.keys[layer_idx] is None:
|
| 110 |
+
self.keys[layer_idx] = new_k
|
| 111 |
+
self.values[layer_idx] = new_v
|
| 112 |
+
else:
|
| 113 |
+
self.keys[layer_idx] = torch.cat([self.keys[layer_idx], new_k], dim=2)
|
| 114 |
+
self.values[layer_idx] = torch.cat([self.values[layer_idx], new_v], dim=2)
|
| 115 |
+
|
| 116 |
+
return self.keys[layer_idx], self.values[layer_idx]
|
| 117 |
+
|
| 118 |
+
def evict_sliding(self, max_keep: int) -> None:
|
| 119 |
+
"""
|
| 120 |
+
Estratégia sliding window pura: mantém apenas os últimos max_keep tokens.
|
| 121 |
+
"""
|
| 122 |
+
for layer_idx in range(self.n_layers):
|
| 123 |
+
if self.keys[layer_idx] is None:
|
| 124 |
+
continue
|
| 125 |
+
seq = self.keys[layer_idx].size(2)
|
| 126 |
+
if seq > max_keep:
|
| 127 |
+
start = seq - max_keep
|
| 128 |
+
self.keys[layer_idx] = self.keys[layer_idx][:, :, start:, :].contiguous()
|
| 129 |
+
self.values[layer_idx] = self.values[layer_idx][:, :, start:, :].contiguous()
|
| 130 |
+
|
| 131 |
+
def evict_sink_sliding(self, max_keep: int, n_sink: int) -> None:
|
| 132 |
+
"""
|
| 133 |
+
Estratégia sink + sliding: mantém os primeiros n_sink + últimos (max_keep - n_sink).
|
| 134 |
+
"""
|
| 135 |
+
if n_sink >= max_keep:
|
| 136 |
+
logger.warning(
|
| 137 |
+
f"n_sink ({n_sink}) >= max_keep ({max_keep}); usando sliding puro"
|
| 138 |
+
)
|
| 139 |
+
self.evict_sliding(max_keep)
|
| 140 |
+
return
|
| 141 |
+
|
| 142 |
+
n_sliding = max_keep - n_sink
|
| 143 |
+
for layer_idx in range(self.n_layers):
|
| 144 |
+
if self.keys[layer_idx] is None:
|
| 145 |
+
continue
|
| 146 |
+
seq = self.keys[layer_idx].size(2)
|
| 147 |
+
if seq > max_keep:
|
| 148 |
+
sink_k = self.keys[layer_idx][:, :, :n_sink, :]
|
| 149 |
+
sink_v = self.values[layer_idx][:, :, :n_sink, :]
|
| 150 |
+
sliding_k = self.keys[layer_idx][:, :, -n_sliding:, :]
|
| 151 |
+
sliding_v = self.values[layer_idx][:, :, -n_sliding:, :]
|
| 152 |
+
self.keys[layer_idx] = torch.cat([sink_k, sliding_k], dim=2).contiguous()
|
| 153 |
+
self.values[layer_idx] = torch.cat([sink_v, sliding_v], dim=2).contiguous()
|
| 154 |
+
|
| 155 |
+
def reset(self) -> None:
|
| 156 |
+
"""Limpa o cache completamente."""
|
| 157 |
+
self.keys = [None] * self.n_layers
|
| 158 |
+
self.values = [None] * self.n_layers
|
| 159 |
+
|
| 160 |
+
def get_seq_len(self, layer_idx: int = 0) -> int:
|
| 161 |
+
if self.keys[layer_idx] is None:
|
| 162 |
+
return 0
|
| 163 |
+
return self.keys[layer_idx].size(2)
|
| 164 |
+
|
| 165 |
+
def total_tokens(self) -> int:
|
| 166 |
+
return max(self.get_seq_len(i) for i in range(self.n_layers))
|
| 167 |
+
|
| 168 |
+
|
| 169 |
+
# ============================================================================
|
| 170 |
+
# Context Window Manager
|
| 171 |
+
# ============================================================================
|
| 172 |
+
|
| 173 |
+
class ContextWindowManager:
|
| 174 |
+
"""
|
| 175 |
+
Gerencia a janela de contexto aplicando a política de evicção configurada.
|
| 176 |
+
"""
|
| 177 |
+
|
| 178 |
+
def __init__(self, config: ContextWindowConfig):
|
| 179 |
+
self.config = config
|
| 180 |
+
self.head_dim = config.head_dim or (config.embed_dim // config.n_heads)
|
| 181 |
+
if config.embed_dim % config.n_heads != 0:
|
| 182 |
+
raise ValueError(
|
| 183 |
+
f"embed_dim ({config.embed_dim}) deve ser divisível por "
|
| 184 |
+
f"n_heads ({config.n_heads})"
|
| 185 |
+
)
|
| 186 |
+
self.kv_cache: Optional[KVCache] = None
|
| 187 |
+
# Sequência de tokens brutos (para recomputação se necessário)
|
| 188 |
+
self.token_history: List[torch.Tensor] = []
|
| 189 |
+
|
| 190 |
+
# ----------------------------------------------------------------------
|
| 191 |
+
# Inicialização
|
| 192 |
+
# ----------------------------------------------------------------------
|
| 193 |
+
|
| 194 |
+
def init_cache(self, batch_size: int, device: torch.device,
|
| 195 |
+
dtype: Optional[torch.dtype] = None) -> KVCache:
|
| 196 |
+
"""Cria novo cache KV para uma sessão de geração."""
|
| 197 |
+
dtype = dtype or self.config.dtype or torch.float32
|
| 198 |
+
self.kv_cache = KVCache(
|
| 199 |
+
n_layers=self.config.n_layers,
|
| 200 |
+
batch_size=batch_size,
|
| 201 |
+
n_heads=self.config.n_heads,
|
| 202 |
+
head_dim=self.head_dim,
|
| 203 |
+
device=device,
|
| 204 |
+
dtype=dtype,
|
| 205 |
+
)
|
| 206 |
+
self.token_history = []
|
| 207 |
+
return self.kv_cache
|
| 208 |
+
|
| 209 |
+
# ----------------------------------------------------------------------
|
| 210 |
+
# Adicionar tokens
|
| 211 |
+
# ----------------------------------------------------------------------
|
| 212 |
+
|
| 213 |
+
def append_tokens(self, token_ids: torch.Tensor) -> torch.Tensor:
|
| 214 |
+
"""
|
| 215 |
+
Adiciona tokens ao histórico e retorna a janela ativa.
|
| 216 |
+
|
| 217 |
+
Args:
|
| 218 |
+
token_ids: [batch, seq_new]
|
| 219 |
+
|
| 220 |
+
Returns:
|
| 221 |
+
window: [batch, seq_window] — tokens na janela ativa após evicção
|
| 222 |
+
"""
|
| 223 |
+
if token_ids.dim() != 2:
|
| 224 |
+
raise ValueError(f"token_ids deve ser [batch, seq]; recebido {token_ids.shape}")
|
| 225 |
+
|
| 226 |
+
# Para simplificação, mantemos histórico apenas do batch[0]
|
| 227 |
+
# (geração autoregressiva tipicamente batch=1)
|
| 228 |
+
if token_ids.size(0) > 1:
|
| 229 |
+
logger.warning("ContextWindowManager.append_tokens: batch>1, histórico rastreia apenas batch[0]")
|
| 230 |
+
|
| 231 |
+
self.token_history.append(token_ids)
|
| 232 |
+
|
| 233 |
+
# Concatenar todo o histórico
|
| 234 |
+
full = torch.cat(self.token_history, dim=1)
|
| 235 |
+
|
| 236 |
+
# Aplicar evicção se necessário
|
| 237 |
+
return self._evict_tokens(full)
|
| 238 |
+
|
| 239 |
+
def _evict_tokens(self, full_seq: torch.Tensor) -> torch.Tensor:
|
| 240 |
+
"""Aplica política de evicção à sequência completa."""
|
| 241 |
+
seq_len = full_seq.size(1)
|
| 242 |
+
max_keep = self.config.max_window
|
| 243 |
+
|
| 244 |
+
if seq_len <= max_keep:
|
| 245 |
+
return full_seq
|
| 246 |
+
|
| 247 |
+
strategy = self.config.eviction_strategy
|
| 248 |
+
if strategy == "sliding":
|
| 249 |
+
return full_seq[:, -max_keep:]
|
| 250 |
+
elif strategy == "sink_sliding":
|
| 251 |
+
n_sink = self.config.n_sink_tokens
|
| 252 |
+
n_sliding = max_keep - n_sink
|
| 253 |
+
sink = full_seq[:, :n_sink]
|
| 254 |
+
sliding = full_seq[:, -n_sliding:]
|
| 255 |
+
return torch.cat([sink, sliding], dim=1)
|
| 256 |
+
elif strategy == "recompute":
|
| 257 |
+
# Manter tudo (atenção recomputada)
|
| 258 |
+
return full_seq
|
| 259 |
+
else:
|
| 260 |
+
logger.warning(f"Estratégia de evicção desconhecida: {strategy}; usando sliding")
|
| 261 |
+
return full_seq[:, -max_keep:]
|
| 262 |
+
|
| 263 |
+
# ----------------------------------------------------------------------
|
| 264 |
+
# Evict cache KV
|
| 265 |
+
# ----------------------------------------------------------------------
|
| 266 |
+
|
| 267 |
+
def evict_cache(self) -> None:
|
| 268 |
+
"""Aplica a política de evicção ao cache KV."""
|
| 269 |
+
if self.kv_cache is None:
|
| 270 |
+
return
|
| 271 |
+
strategy = self.config.eviction_strategy
|
| 272 |
+
max_keep = self.config.max_window
|
| 273 |
+
if strategy == "sliding":
|
| 274 |
+
self.kv_cache.evict_sliding(max_keep)
|
| 275 |
+
elif strategy == "sink_sliding":
|
| 276 |
+
self.kv_cache.evict_sink_sliding(max_keep, self.config.n_sink_tokens)
|
| 277 |
+
elif strategy == "recompute":
|
| 278 |
+
self.kv_cache.reset()
|
| 279 |
+
else:
|
| 280 |
+
self.kv_cache.evict_sliding(max_keep)
|
| 281 |
+
|
| 282 |
+
# ----------------------------------------------------------------------
|
| 283 |
+
# Reset
|
| 284 |
+
# ----------------------------------------------------------------------
|
| 285 |
+
|
| 286 |
+
def reset(self) -> None:
|
| 287 |
+
"""Reseta todo o estado (cache + histórico)."""
|
| 288 |
+
if self.kv_cache is not None:
|
| 289 |
+
self.kv_cache.reset()
|
| 290 |
+
self.kv_cache = None
|
| 291 |
+
self.token_history = []
|
| 292 |
+
|
| 293 |
+
# ----------------------------------------------------------------------
|
| 294 |
+
# Info
|
| 295 |
+
# ----------------------------------------------------------------------
|
| 296 |
+
|
| 297 |
+
def get_info(self) -> Dict:
|
| 298 |
+
"""Retorna informações sobre o estado atual."""
|
| 299 |
+
cache_len = self.kv_cache.total_tokens() if self.kv_cache else 0
|
| 300 |
+
history_len = sum(t.size(1) for t in self.token_history)
|
| 301 |
+
return {
|
| 302 |
+
"max_window": self.config.max_window,
|
| 303 |
+
"strategy": self.config.eviction_strategy,
|
| 304 |
+
"n_sink": self.config.n_sink_tokens,
|
| 305 |
+
"cache_tokens": cache_len,
|
| 306 |
+
"history_tokens": history_len,
|
| 307 |
+
"active": self.kv_cache is not None,
|
| 308 |
+
}
|
| 309 |
+
|
| 310 |
+
|
| 311 |
+
# ============================================================================
|
| 312 |
+
# Helpers
|
| 313 |
+
# ============================================================================
|
| 314 |
+
|
| 315 |
+
def make_context_window(
|
| 316 |
+
max_window: int = 512,
|
| 317 |
+
strategy: str = "sink_sliding",
|
| 318 |
+
n_sink: int = 4,
|
| 319 |
+
embed_dim: int = 256,
|
| 320 |
+
n_heads: int = 4,
|
| 321 |
+
n_layers: int = 2,
|
| 322 |
+
device: str = "cpu",
|
| 323 |
+
) -> ContextWindowManager:
|
| 324 |
+
"""Factory para criar ContextWindowManager com defaults sensatos."""
|
| 325 |
+
config = ContextWindowConfig(
|
| 326 |
+
max_window=max_window,
|
| 327 |
+
eviction_strategy=strategy,
|
| 328 |
+
n_sink_tokens=n_sink,
|
| 329 |
+
embed_dim=embed_dim,
|
| 330 |
+
n_heads=n_heads,
|
| 331 |
+
head_dim=embed_dim // n_heads,
|
| 332 |
+
n_layers=n_layers,
|
| 333 |
+
device=device,
|
| 334 |
+
)
|
| 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 |
+
]
|
cnn_bigru/models/cooperative_bigru.py
ADDED
|
@@ -0,0 +1,519 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""cooperative_bigru.py — Núcleo cooperativo CNN-BiGRU do projeto.
|
| 2 |
+
|
| 3 |
+
Implementa os 3 níveis de cooperação descritos em dados.txt:
|
| 4 |
+
|
| 5 |
+
NÍVEL 1 — Cooperação Global (Cross-Attention pós-CNN):
|
| 6 |
+
PonteGlobalCrossAttention com máscaras de padding.
|
| 7 |
+
|
| 8 |
+
NÍVEL 2 — Cooperação Passo a Passo (GRUCell com Porta Gated):
|
| 9 |
+
CellStepCooperativaComAtenuacao: calcula alfa (Sigmoid) que regula
|
| 10 |
+
o quanto uma rede aceita da outra, com máscara de padding.
|
| 11 |
+
|
| 12 |
+
NÍVEL 3 — Ponte de Fusão Final:
|
| 13 |
+
Concatena os estados forward+backward de cada rede (128 por rede = 256)
|
| 14 |
+
e projeta para a camada de classificação/decodificação.
|
| 15 |
+
|
| 16 |
+
Inclui:
|
| 17 |
+
- Inicialização ortogonal das matrizes recorrentes
|
| 18 |
+
- Normalização espectral opcional nos pesos lineares
|
| 19 |
+
- Self-attention no final de cada bloco (requisito do dados.txt)
|
| 20 |
+
"""
|
| 21 |
+
from __future__ import annotations
|
| 22 |
+
|
| 23 |
+
import logging
|
| 24 |
+
import math
|
| 25 |
+
from typing import Optional, Tuple
|
| 26 |
+
|
| 27 |
+
import torch
|
| 28 |
+
import torch.nn as nn
|
| 29 |
+
import torch.nn.functional as F
|
| 30 |
+
|
| 31 |
+
logger = logging.getLogger(__name__)
|
| 32 |
+
|
| 33 |
+
|
| 34 |
+
def _orthogonal_init(layer: nn.Module) -> None:
|
| 35 |
+
"""Aplica inicialização ortogonal nos pesos de uma camada Linear/GRUCell.
|
| 36 |
+
|
| 37 |
+
Em caso de falha (e.g., matriz não-quadrada com fan_in < fan_out), registra
|
| 38 |
+
warning e faz fallback para xavier_uniform_ (não silencioso).
|
| 39 |
+
"""
|
| 40 |
+
for name, p in layer.named_parameters():
|
| 41 |
+
if "weight" in name and p.dim() >= 2:
|
| 42 |
+
try:
|
| 43 |
+
nn.init.orthogonal_(p)
|
| 44 |
+
except Exception as e:
|
| 45 |
+
logger.warning(
|
| 46 |
+
f"_orthogonal_init: falhou em {name} (shape={tuple(p.shape)}): {e}. "
|
| 47 |
+
f"Usando xavier_uniform_ como fallback."
|
| 48 |
+
)
|
| 49 |
+
nn.init.xavier_uniform_(p)
|
| 50 |
+
elif "bias" in name:
|
| 51 |
+
nn.init.zeros_(p)
|
| 52 |
+
|
| 53 |
+
|
| 54 |
+
# ============================================================================
|
| 55 |
+
# NÍVEL 1 — Cross-Attention Global com máscaras
|
| 56 |
+
# ============================================================================
|
| 57 |
+
|
| 58 |
+
class PonteGlobalCrossAttention(nn.Module):
|
| 59 |
+
"""Atenção cruzada pós-CNN com máscara de padding.
|
| 60 |
+
|
| 61 |
+
A Rede A usa seus dados como Query para encontrar informações relevantes
|
| 62 |
+
nas Keys da Rede B, agregando os Values mais importantes.
|
| 63 |
+
"""
|
| 64 |
+
|
| 65 |
+
def __init__(self, channels: int, n_heads: int = 4, dropout: float = 0.1):
|
| 66 |
+
super().__init__()
|
| 67 |
+
assert channels % n_heads == 0, f"channels ({channels}) deve ser divisível por n_heads ({n_heads})"
|
| 68 |
+
self.channels = channels
|
| 69 |
+
self.n_heads = n_heads
|
| 70 |
+
self.head_dim = channels // n_heads
|
| 71 |
+
|
| 72 |
+
self.query = nn.Linear(channels, channels)
|
| 73 |
+
self.key = nn.Linear(channels, channels)
|
| 74 |
+
self.value = nn.Linear(channels, channels)
|
| 75 |
+
self.out = nn.Linear(channels, channels)
|
| 76 |
+
self.dropout = nn.Dropout(dropout)
|
| 77 |
+
|
| 78 |
+
def forward(
|
| 79 |
+
self,
|
| 80 |
+
x_a: torch.Tensor,
|
| 81 |
+
x_b: torch.Tensor,
|
| 82 |
+
mask_b: Optional[torch.Tensor] = None,
|
| 83 |
+
) -> torch.Tensor:
|
| 84 |
+
# x_a: [B, T_a, C], x_b: [B, T_b, C]
|
| 85 |
+
B, Ta, C = x_a.shape
|
| 86 |
+
Tb = x_b.size(1)
|
| 87 |
+
|
| 88 |
+
q = self.query(x_a).view(B, Ta, self.n_heads, self.head_dim).transpose(1, 2)
|
| 89 |
+
k = self.key(x_b).view(B, Tb, self.n_heads, self.head_dim).transpose(1, 2)
|
| 90 |
+
v = self.value(x_b).view(B, Tb, self.n_heads, self.head_dim).transpose(1, 2)
|
| 91 |
+
|
| 92 |
+
scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.head_dim)
|
| 93 |
+
|
| 94 |
+
if mask_b is not None:
|
| 95 |
+
# mask_b: [B, T_b] -> [B, 1, 1, T_b]
|
| 96 |
+
scores = scores.masked_fill(mask_b.unsqueeze(1).unsqueeze(1) == 0, -1e9)
|
| 97 |
+
|
| 98 |
+
attn = F.softmax(scores, dim=-1)
|
| 99 |
+
attn = self.dropout(attn)
|
| 100 |
+
ctx = torch.matmul(attn, v) # [B, n_heads, Ta, head_dim]
|
| 101 |
+
ctx = ctx.transpose(1, 2).contiguous().view(B, Ta, C)
|
| 102 |
+
return x_a + self.out(ctx)
|
| 103 |
+
|
| 104 |
+
|
| 105 |
+
# ============================================================================
|
| 106 |
+
# NÍVEL 2 — Célula Cooperativa Passo-a-Passo com Atenuação Gated
|
| 107 |
+
# ============================================================================
|
| 108 |
+
|
| 109 |
+
class CellStepCooperativaComAtenuacao(nn.Module):
|
| 110 |
+
"""Célula GRU cooperativa com porta de atenuação (gated bridge).
|
| 111 |
+
|
| 112 |
+
Calcula alfa = Sigmoid(linear(concat(x_t, h_partner))) em [0, 1]:
|
| 113 |
+
- Se alfa -> 0, a ponte desativa (isolamento dos fluxos)
|
| 114 |
+
- Se alfa -> 1, a ponte transfere totalmente o contexto
|
| 115 |
+
|
| 116 |
+
Tratamento de padding: se mask_t == 0, mantém o h anterior intacto.
|
| 117 |
+
"""
|
| 118 |
+
|
| 119 |
+
def __init__(self, input_size: int, hidden_size: int):
|
| 120 |
+
super().__init__()
|
| 121 |
+
self.cell_A = nn.GRUCell(input_size, hidden_size)
|
| 122 |
+
self.cell_B = nn.GRUCell(input_size, hidden_size)
|
| 123 |
+
_orthogonal_init(self.cell_A)
|
| 124 |
+
_orthogonal_init(self.cell_B)
|
| 125 |
+
|
| 126 |
+
# Pontes lineares
|
| 127 |
+
self.ponte_B_para_A = nn.Linear(hidden_size, input_size)
|
| 128 |
+
self.ponte_A_para_B = nn.Linear(hidden_size, input_size)
|
| 129 |
+
|
| 130 |
+
# Porta de atenuação (gated)
|
| 131 |
+
self.porta_atenuacao_A = nn.Linear(input_size + hidden_size, 1)
|
| 132 |
+
self.porta_atenuacao_B = nn.Linear(input_size + hidden_size, 1)
|
| 133 |
+
|
| 134 |
+
def forward(
|
| 135 |
+
self,
|
| 136 |
+
x_t_A: torch.Tensor,
|
| 137 |
+
x_t_B: torch.Tensor,
|
| 138 |
+
h_prev_A: torch.Tensor,
|
| 139 |
+
h_prev_B: torch.Tensor,
|
| 140 |
+
mask_t_A: torch.Tensor,
|
| 141 |
+
mask_t_B: torch.Tensor,
|
| 142 |
+
) -> Tuple[torch.Tensor, torch.Tensor]:
|
| 143 |
+
# mask_t: [B, 1]
|
| 144 |
+
|
| 145 |
+
# 1. Fator de atenuação (Sigmoid -> [0, 1])
|
| 146 |
+
alfa_A = torch.sigmoid(self.porta_atenuacao_A(torch.cat((x_t_A, h_prev_B), dim=-1)))
|
| 147 |
+
alfa_B = torch.sigmoid(self.porta_atenuacao_B(torch.cat((x_t_B, h_prev_A), dim=-1)))
|
| 148 |
+
|
| 149 |
+
# 2. Injeção de contexto com atenuação + máscara
|
| 150 |
+
influencia_B = torch.tanh(self.ponte_B_para_A(h_prev_B)) * alfa_A * mask_t_B
|
| 151 |
+
influencia_A = torch.tanh(self.ponte_A_para_B(h_prev_A)) * alfa_B * mask_t_A
|
| 152 |
+
|
| 153 |
+
x_t_A_coop = x_t_A + influencia_B
|
| 154 |
+
x_t_B_coop = x_t_B + influencia_A
|
| 155 |
+
|
| 156 |
+
# 3. Atualização GRU
|
| 157 |
+
h_next_A_bruto = self.cell_A(x_t_A_coop, h_prev_A)
|
| 158 |
+
h_next_B_bruto = self.cell_B(x_t_B_coop, h_prev_B)
|
| 159 |
+
|
| 160 |
+
# 4. Preserva h anterior onde a sequência acabou (mask=0)
|
| 161 |
+
h_next_A = torch.where(mask_t_A == 1, h_next_A_bruto, h_prev_A)
|
| 162 |
+
h_next_B = torch.where(mask_t_B == 1, h_next_B_bruto, h_prev_B)
|
| 163 |
+
|
| 164 |
+
return h_next_A, h_next_B
|
| 165 |
+
|
| 166 |
+
|
| 167 |
+
# ============================================================================
|
| 168 |
+
# Self-Attention Final Layer (requisito do dados.txt linha 470)
|
| 169 |
+
# ============================================================================
|
| 170 |
+
|
| 171 |
+
class SelfAttentionSummary(nn.Module):
|
| 172 |
+
"""Self-attention final layer que resume os passos temporais antes da classificação.
|
| 173 |
+
|
| 174 |
+
Implementa uma self-attention TRUE (não usa mean como query, conforme
|
| 175 |
+
requisito do dados.txt linha 470 "Self-Attention no final de cada bloco"):
|
| 176 |
+
- Query, Key, Value são todas projeções lineares aprendidas de x
|
| 177 |
+
- Multi-head com n_heads cabeças
|
| 178 |
+
- Máscara de padding (ignora tokens [PAD])
|
| 179 |
+
- Causal opcional (não usado em sumarização, mas disponível)
|
| 180 |
+
- Residual + LayerNorm para estabilidade
|
| 181 |
+
|
| 182 |
+
Saída: [B, hidden_dim] — vetor agregado que resume toda a sequência.
|
| 183 |
+
|
| 184 |
+
Args:
|
| 185 |
+
hidden_dim: dimensão de entrada/saída
|
| 186 |
+
n_heads: número de cabeças (deve dividir hidden_dim)
|
| 187 |
+
dropout: probabilidade de dropout
|
| 188 |
+
use_residual: se True, soma residual com entrada (via LayerNorm)
|
| 189 |
+
aggregation: "mean" (média dos heads após attn) ou "first" (primeiro token)
|
| 190 |
+
"""
|
| 191 |
+
|
| 192 |
+
def __init__(
|
| 193 |
+
self,
|
| 194 |
+
hidden_dim: int,
|
| 195 |
+
n_heads: int = 4,
|
| 196 |
+
dropout: float = 0.1,
|
| 197 |
+
use_residual: bool = True,
|
| 198 |
+
aggregation: str = "mean",
|
| 199 |
+
):
|
| 200 |
+
super().__init__()
|
| 201 |
+
if hidden_dim % n_heads != 0:
|
| 202 |
+
raise ValueError(
|
| 203 |
+
f"hidden_dim ({hidden_dim}) deve ser divisível por n_heads ({n_heads})"
|
| 204 |
+
)
|
| 205 |
+
self.hidden_dim = hidden_dim
|
| 206 |
+
self.n_heads = n_heads
|
| 207 |
+
self.head_dim = hidden_dim // n_heads
|
| 208 |
+
self.use_residual = use_residual
|
| 209 |
+
self.aggregation = aggregation
|
| 210 |
+
|
| 211 |
+
# Projeções lineares aprendidas (TRUE self-attention, não mean-pool)
|
| 212 |
+
self.q_proj = nn.Linear(hidden_dim, hidden_dim)
|
| 213 |
+
self.k_proj = nn.Linear(hidden_dim, hidden_dim)
|
| 214 |
+
self.v_proj = nn.Linear(hidden_dim, hidden_dim)
|
| 215 |
+
self.out_proj = nn.Linear(hidden_dim, hidden_dim)
|
| 216 |
+
self.dropout = nn.Dropout(dropout)
|
| 217 |
+
|
| 218 |
+
if use_residual:
|
| 219 |
+
self.ln = nn.LayerNorm(hidden_dim)
|
| 220 |
+
|
| 221 |
+
# Token especial de "sumarização" aprendido (CLS-like)
|
| 222 |
+
# Adicionado no início da sequência para servir como query global
|
| 223 |
+
self.cls_token = nn.Parameter(torch.zeros(1, 1, hidden_dim))
|
| 224 |
+
nn.init.normal_(self.cls_token, std=0.02)
|
| 225 |
+
|
| 226 |
+
def forward(
|
| 227 |
+
self,
|
| 228 |
+
x: torch.Tensor,
|
| 229 |
+
mask: Optional[torch.Tensor] = None,
|
| 230 |
+
) -> torch.Tensor:
|
| 231 |
+
"""
|
| 232 |
+
Args:
|
| 233 |
+
x: [B, T, H] — sequência de saídas da BiGRU
|
| 234 |
+
mask: [B, T] — 1 para token real, 0 para PAD (None = todos válidos)
|
| 235 |
+
|
| 236 |
+
Returns:
|
| 237 |
+
summary: [B, H] — vetor sumarizado
|
| 238 |
+
"""
|
| 239 |
+
B, T, H = x.shape
|
| 240 |
+
|
| 241 |
+
# Prepend CLS token: [B, 1, H] + [B, T, H] -> [B, T+1, H]
|
| 242 |
+
cls = self.cls_token.expand(B, -1, -1)
|
| 243 |
+
x_ext = torch.cat([cls, x], dim=1)
|
| 244 |
+
T_ext = T + 1
|
| 245 |
+
|
| 246 |
+
# Estender máscara para o CLS (sempre ativo)
|
| 247 |
+
if mask is not None:
|
| 248 |
+
cls_mask = torch.ones(B, 1, device=x.device, dtype=mask.dtype)
|
| 249 |
+
mask_ext = torch.cat([cls_mask, mask], dim=1) # [B, T+1]
|
| 250 |
+
else:
|
| 251 |
+
mask_ext = None
|
| 252 |
+
|
| 253 |
+
# Projeções Q, K, V (todas aprendidas — true self-attention)
|
| 254 |
+
q = self.q_proj(x_ext) # [B, T+1, H]
|
| 255 |
+
k = self.k_proj(x_ext)
|
| 256 |
+
v = self.v_proj(x_ext)
|
| 257 |
+
|
| 258 |
+
# Reshape para multi-head: [B, n_heads, T+1, head_dim]
|
| 259 |
+
q = q.view(B, T_ext, self.n_heads, self.head_dim).transpose(1, 2)
|
| 260 |
+
k = k.view(B, T_ext, self.n_heads, self.head_dim).transpose(1, 2)
|
| 261 |
+
v = v.view(B, T_ext, self.n_heads, self.head_dim).transpose(1, 2)
|
| 262 |
+
|
| 263 |
+
# Scaled dot-product attention (sem máscara causal — é encoder-style)
|
| 264 |
+
scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.head_dim)
|
| 265 |
+
# scores: [B, n_heads, T+1, T+1]
|
| 266 |
+
|
| 267 |
+
if mask_ext is not None:
|
| 268 |
+
# Aplicar máscara nas chaves: tokens PAD devem receber atenção zero
|
| 269 |
+
# mask_ext: [B, T+1] -> [B, 1, 1, T+1]
|
| 270 |
+
scores = scores.masked_fill(
|
| 271 |
+
mask_ext.unsqueeze(1).unsqueeze(1) == 0, -1e9
|
| 272 |
+
)
|
| 273 |
+
|
| 274 |
+
# Lidar com linhas totalmente mascaradas (evitar NaN)
|
| 275 |
+
all_neg = (scores <= -1e9).all(dim=-1, keepdim=True)
|
| 276 |
+
scores = torch.where(all_neg, torch.zeros_like(scores), scores)
|
| 277 |
+
|
| 278 |
+
attn = F.softmax(scores, dim=-1)
|
| 279 |
+
attn = self.dropout(attn)
|
| 280 |
+
|
| 281 |
+
# Contexto: [B, n_heads, T+1, head_dim]
|
| 282 |
+
ctx = torch.matmul(attn, v)
|
| 283 |
+
ctx = ctx.transpose(1, 2).contiguous().view(B, T_ext, H)
|
| 284 |
+
ctx = self.out_proj(ctx)
|
| 285 |
+
|
| 286 |
+
# Residual + LayerNorm (estabilidade)
|
| 287 |
+
if self.use_residual:
|
| 288 |
+
ctx = self.ln(ctx + x_ext)
|
| 289 |
+
|
| 290 |
+
# Sumarização: usar o CLS token (posição 0) como saída
|
| 291 |
+
if self.aggregation == "first":
|
| 292 |
+
summary = ctx[:, 0, :] # [B, H]
|
| 293 |
+
else:
|
| 294 |
+
# Média sobre todas as posições válidas (com mask)
|
| 295 |
+
if mask_ext is not None:
|
| 296 |
+
# Ponderar pela máscara
|
| 297 |
+
weights = mask_ext.unsqueeze(-1).float() # [B, T+1, 1]
|
| 298 |
+
sum_ctx = (ctx * weights).sum(dim=1) # [B, H]
|
| 299 |
+
count = weights.sum(dim=1).clamp(min=1.0)
|
| 300 |
+
summary = sum_ctx / count
|
| 301 |
+
else:
|
| 302 |
+
summary = ctx.mean(dim=1)
|
| 303 |
+
|
| 304 |
+
return summary
|
| 305 |
+
|
| 306 |
+
|
| 307 |
+
# ============================================================================
|
| 308 |
+
# Núcleo Cooperativo CNN-BiGRU
|
| 309 |
+
# ============================================================================
|
| 310 |
+
|
| 311 |
+
class CooperativeCNNBiGRU(nn.Module):
|
| 312 |
+
"""Núcleo cooperativo CNN-BiGRU (dual-stream) com 3 níveis de ponte.
|
| 313 |
+
|
| 314 |
+
Args:
|
| 315 |
+
vocab_size: tamanho do vocabulário.
|
| 316 |
+
embedding_dim: dimensão do embedding de tokens.
|
| 317 |
+
cnn_filters: número de filtros da Conv1d.
|
| 318 |
+
gru_hidden: dimensão oculta de cada direção da GRU.
|
| 319 |
+
n_heads: número de cabeças para cross-attention.
|
| 320 |
+
dropout: probabilidade de dropout.
|
| 321 |
+
pad_idx: índice do token de padding.
|
| 322 |
+
use_spectral_norm: se True, aplica normalização espectral nas lineares.
|
| 323 |
+
"""
|
| 324 |
+
|
| 325 |
+
def __init__(
|
| 326 |
+
self,
|
| 327 |
+
vocab_size: int,
|
| 328 |
+
embedding_dim: int = 64,
|
| 329 |
+
cnn_filters: int = 64,
|
| 330 |
+
gru_hidden: int = 64,
|
| 331 |
+
n_heads: int = 4,
|
| 332 |
+
dropout: float = 0.1,
|
| 333 |
+
pad_idx: int = 1,
|
| 334 |
+
use_spectral_norm: bool = False,
|
| 335 |
+
):
|
| 336 |
+
super().__init__()
|
| 337 |
+
self.pad_idx = pad_idx
|
| 338 |
+
self.cnn_filters = cnn_filters
|
| 339 |
+
self.gru_hidden = gru_hidden
|
| 340 |
+
|
| 341 |
+
self.embedding = nn.Embedding(vocab_size, embedding_dim, padding_idx=pad_idx)
|
| 342 |
+
nn.init.orthogonal_(self.embedding.weight)
|
| 343 |
+
|
| 344 |
+
# CNN extractors (cada stream tem o seu)
|
| 345 |
+
self.conv_A = nn.Conv1d(embedding_dim, cnn_filters, kernel_size=3, padding=1)
|
| 346 |
+
self.conv_B = nn.Conv1d(embedding_dim, cnn_filters, kernel_size=3, padding=1)
|
| 347 |
+
|
| 348 |
+
# Nível 1: Pontes globais (cross-attention)
|
| 349 |
+
self.ponte_global_A = PonteGlobalCrossAttention(cnn_filters, n_heads=n_heads, dropout=dropout)
|
| 350 |
+
self.ponte_global_B = PonteGlobalCrossAttention(cnn_filters, n_heads=n_heads, dropout=dropout)
|
| 351 |
+
|
| 352 |
+
# Nível 2: Células cooperativas passo-a-passo
|
| 353 |
+
self.gru_coop_direto = CellStepCooperativaComAtenuacao(cnn_filters, gru_hidden)
|
| 354 |
+
self.gru_coop_inverso = CellStepCooperativaComAtenuacao(cnn_filters, gru_hidden)
|
| 355 |
+
|
| 356 |
+
# Self-attention para sumarização temporal (requisito do dados.txt)
|
| 357 |
+
self.self_attn_A = SelfAttentionSummary(gru_hidden * 2, dropout=dropout)
|
| 358 |
+
self.self_attn_B = SelfAttentionSummary(gru_hidden * 2, dropout=dropout)
|
| 359 |
+
|
| 360 |
+
# Camadas de saída (sem classificador aqui — o classificador fica no modelo multimodal)
|
| 361 |
+
self.dropout = nn.Dropout(dropout)
|
| 362 |
+
self.feat_dim = gru_hidden * 2 # 128 por stream (forward + backward concat)
|
| 363 |
+
|
| 364 |
+
# Spectral norm opcional
|
| 365 |
+
if use_spectral_norm:
|
| 366 |
+
self.conv_A = nn.utils.spectral_norm(self.conv_A)
|
| 367 |
+
self.conv_B = nn.utils.spectral_norm(self.conv_B)
|
| 368 |
+
self.ponte_global_A.query = nn.utils.spectral_norm(self.ponte_global_A.query)
|
| 369 |
+
self.ponte_global_A.key = nn.utils.spectral_norm(self.ponte_global_A.key)
|
| 370 |
+
self.ponte_global_A.value = nn.utils.spectral_norm(self.ponte_global_A.value)
|
| 371 |
+
self.ponte_global_B.query = nn.utils.spectral_norm(self.ponte_global_B.query)
|
| 372 |
+
self.ponte_global_B.key = nn.utils.spectral_norm(self.ponte_global_B.key)
|
| 373 |
+
self.ponte_global_B.value = nn.utils.spectral_norm(self.ponte_global_B.value)
|
| 374 |
+
|
| 375 |
+
def forward(
|
| 376 |
+
self,
|
| 377 |
+
input_A: torch.Tensor,
|
| 378 |
+
input_B: torch.Tensor,
|
| 379 |
+
return_sequences: bool = False,
|
| 380 |
+
) -> dict:
|
| 381 |
+
"""
|
| 382 |
+
Args:
|
| 383 |
+
input_A: [B, T_a] tokens do stream A
|
| 384 |
+
input_B: [B, T_b] tokens do stream B
|
| 385 |
+
return_sequences: se True, retorna também as saídas temporais das BiGRUs
|
| 386 |
+
|
| 387 |
+
Returns:
|
| 388 |
+
dict com:
|
| 389 |
+
out_A: [B, 128] representação final do stream A
|
| 390 |
+
out_B: [B, 128] representação final do stream B
|
| 391 |
+
fused: [B, 256] concatenação (out_A || out_B)
|
| 392 |
+
(opcional) seq_A: [B, T, 128] saídas temporais do stream A
|
| 393 |
+
(opcional) seq_B: [B, T, 128] saídas temporais do stream B
|
| 394 |
+
"""
|
| 395 |
+
batch_size = input_A.size(0)
|
| 396 |
+
seq_len_A = input_A.size(1)
|
| 397 |
+
seq_len_B = input_B.size(1)
|
| 398 |
+
device = input_A.device
|
| 399 |
+
|
| 400 |
+
# Máscaras de padding
|
| 401 |
+
mask_A = (input_A != self.pad_idx).float()
|
| 402 |
+
mask_B = (input_B != self.pad_idx).float()
|
| 403 |
+
|
| 404 |
+
# 1. Embeddings + CNN
|
| 405 |
+
emb_A = self.embedding(input_A).permute(0, 2, 1) # [B, D, T_a]
|
| 406 |
+
emb_B = self.embedding(input_B).permute(0, 2, 1) # [B, D, T_b]
|
| 407 |
+
|
| 408 |
+
feat_A = F.relu(self.conv_A(emb_A)).permute(0, 2, 1) * mask_A.unsqueeze(-1)
|
| 409 |
+
feat_B = F.relu(self.conv_B(emb_B)).permute(0, 2, 1) * mask_B.unsqueeze(-1)
|
| 410 |
+
|
| 411 |
+
# 2. NÍVEL 1: Pontes globais (cross-attention)
|
| 412 |
+
feat_A_coop = self.ponte_global_A(feat_A, feat_B, mask_B) * mask_A.unsqueeze(-1)
|
| 413 |
+
feat_B_coop = self.ponte_global_B(feat_B, feat_A, mask_A) * mask_B.unsqueeze(-1)
|
| 414 |
+
|
| 415 |
+
# 3. Estados ocultos iniciais
|
| 416 |
+
h_dir_A = torch.zeros(batch_size, self.gru_hidden, device=device)
|
| 417 |
+
h_dir_B = torch.zeros(batch_size, self.gru_hidden, device=device)
|
| 418 |
+
h_inv_A = torch.zeros(batch_size, self.gru_hidden, device=device)
|
| 419 |
+
h_inv_B = torch.zeros(batch_size, self.gru_hidden, device=device)
|
| 420 |
+
|
| 421 |
+
max_steps = max(seq_len_A, seq_len_B)
|
| 422 |
+
seq_dir_A, seq_dir_B = [], []
|
| 423 |
+
seq_inv_A, seq_inv_B = [], []
|
| 424 |
+
|
| 425 |
+
# 4. NÍVEL 2: Laço direto
|
| 426 |
+
# Sempre coletamos as sequências temporais (necessário para self-attention final)
|
| 427 |
+
for t in range(max_steps):
|
| 428 |
+
x_t_A = feat_A_coop[:, t, :] if t < seq_len_A else torch.zeros(batch_size, self.cnn_filters, device=device)
|
| 429 |
+
x_t_B = feat_B_coop[:, t, :] if t < seq_len_B else torch.zeros(batch_size, self.cnn_filters, device=device)
|
| 430 |
+
|
| 431 |
+
m_t_A = mask_A[:, t].unsqueeze(-1) if t < seq_len_A else torch.zeros(batch_size, 1, device=device)
|
| 432 |
+
m_t_B = mask_B[:, t].unsqueeze(-1) if t < seq_len_B else torch.zeros(batch_size, 1, device=device)
|
| 433 |
+
|
| 434 |
+
h_dir_A, h_dir_B = self.gru_coop_direto(x_t_A, x_t_B, h_dir_A, h_dir_B, m_t_A, m_t_B)
|
| 435 |
+
seq_dir_A.append(h_dir_A)
|
| 436 |
+
seq_dir_B.append(h_dir_B)
|
| 437 |
+
|
| 438 |
+
# 5. Laço inverso
|
| 439 |
+
for t in reversed(range(max_steps)):
|
| 440 |
+
x_t_A = feat_A_coop[:, t, :] if t < seq_len_A else torch.zeros(batch_size, self.cnn_filters, device=device)
|
| 441 |
+
x_t_B = feat_B_coop[:, t, :] if t < seq_len_B else torch.zeros(batch_size, self.cnn_filters, device=device)
|
| 442 |
+
|
| 443 |
+
m_t_A = mask_A[:, t].unsqueeze(-1) if t < seq_len_A else torch.zeros(batch_size, 1, device=device)
|
| 444 |
+
m_t_B = mask_B[:, t].unsqueeze(-1) if t < seq_len_B else torch.zeros(batch_size, 1, device=device)
|
| 445 |
+
|
| 446 |
+
h_inv_A, h_inv_B = self.gru_coop_inverso(x_t_A, x_t_B, h_inv_A, h_inv_B, m_t_A, m_t_B)
|
| 447 |
+
seq_inv_A.append(h_inv_A)
|
| 448 |
+
seq_inv_B.append(h_inv_B)
|
| 449 |
+
|
| 450 |
+
# 6. Concatenação forward+backward (128 por stream)
|
| 451 |
+
final_A = torch.cat((h_dir_A, h_inv_A), dim=1) # [B, 128]
|
| 452 |
+
final_B = torch.cat((h_dir_B, h_inv_B), dim=1) # [B, 128]
|
| 453 |
+
|
| 454 |
+
# 6.5 SELF-ATTENTION FINAL LAYER (requisito do dados.txt linha 470)
|
| 455 |
+
# Constrói sequência temporal [B, T, 2H] para aplicar self-attention
|
| 456 |
+
out_A = final_A
|
| 457 |
+
out_B = final_B
|
| 458 |
+
seq_A_full = None
|
| 459 |
+
seq_B_full = None
|
| 460 |
+
|
| 461 |
+
# Aplicar self-attention TRUE (não mean-pool) sobre a sequência temporal
|
| 462 |
+
if getattr(self, "self_attn_A", None) is not None:
|
| 463 |
+
# Construir sequência temporal concatenando dir + inv em cada posição t
|
| 464 |
+
try:
|
| 465 |
+
seq_A_list = []
|
| 466 |
+
seq_B_list = []
|
| 467 |
+
for i in range(max_steps):
|
| 468 |
+
if i < len(seq_dir_A) and (max_steps - 1 - i) < len(seq_inv_A):
|
| 469 |
+
a_dir = seq_dir_A[i]
|
| 470 |
+
a_inv = seq_inv_A[max_steps - 1 - i]
|
| 471 |
+
seq_A_list.append(torch.cat((a_dir, a_inv), dim=1))
|
| 472 |
+
if i < len(seq_dir_B) and (max_steps - 1 - i) < len(seq_inv_B):
|
| 473 |
+
b_dir = seq_dir_B[i]
|
| 474 |
+
b_inv = seq_inv_B[max_steps - 1 - i]
|
| 475 |
+
seq_B_list.append(torch.cat((b_dir, b_inv), dim=1))
|
| 476 |
+
if seq_A_list:
|
| 477 |
+
seq_A_full = torch.stack(seq_A_list, dim=1) # [B, T, 2H]
|
| 478 |
+
if seq_B_list:
|
| 479 |
+
seq_B_full = torch.stack(seq_B_list, dim=1)
|
| 480 |
+
except Exception as e:
|
| 481 |
+
logger.warning(f"SelfAttention: falha ao construir sequência temporal: {e}")
|
| 482 |
+
seq_A_full = None
|
| 483 |
+
seq_B_full = None
|
| 484 |
+
|
| 485 |
+
# Aplicar self-attention TRUE (não mean-pool)
|
| 486 |
+
if seq_A_full is not None:
|
| 487 |
+
out_A = self.self_attn_A(seq_A_full, mask=mask_A)
|
| 488 |
+
if seq_B_full is not None:
|
| 489 |
+
out_B = self.self_attn_B(seq_B_full, mask=mask_B)
|
| 490 |
+
|
| 491 |
+
# 7. NÍVEL 3: Fusão final
|
| 492 |
+
fused = torch.cat((out_A, out_B), dim=1) # [B, 256]
|
| 493 |
+
|
| 494 |
+
result = {
|
| 495 |
+
"out_A": out_A,
|
| 496 |
+
"out_B": out_B,
|
| 497 |
+
"fused": fused,
|
| 498 |
+
"final_A": final_A, # estados finais raw (sem self-attention)
|
| 499 |
+
"final_B": final_B,
|
| 500 |
+
"mask_A": mask_A,
|
| 501 |
+
"mask_B": mask_B,
|
| 502 |
+
}
|
| 503 |
+
|
| 504 |
+
if return_sequences:
|
| 505 |
+
# Sequências temporais [B, T, 2H]
|
| 506 |
+
if seq_A_full is not None:
|
| 507 |
+
result["seq_A"] = seq_A_full
|
| 508 |
+
if seq_B_full is not None:
|
| 509 |
+
result["seq_B"] = seq_B_full
|
| 510 |
+
|
| 511 |
+
return result
|
| 512 |
+
|
| 513 |
+
|
| 514 |
+
__all__ = [
|
| 515 |
+
"PonteGlobalCrossAttention",
|
| 516 |
+
"CellStepCooperativaComAtenuacao",
|
| 517 |
+
"SelfAttentionSummary",
|
| 518 |
+
"CooperativeCNNBiGRU",
|
| 519 |
+
]
|
cnn_bigru/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 |
+
]
|
cnn_bigru/models/generator_verifier.py
ADDED
|
@@ -0,0 +1,364 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""generator_verifier.py — Gerador + Verificador + Anti-Alucinação.
|
| 2 |
+
|
| 3 |
+
Implementa o pseudocódigo da seção 5-7 de dados.txt:
|
| 4 |
+
|
| 5 |
+
- Generator: CNN-BiGRU encoder + GRU decoder com atenção Bahdanau
|
| 6 |
+
- Verifier: CNN-BiGRU + classificador sigmoid (v ∈ [0,1])
|
| 7 |
+
- Anti-hallucination layer: ativadores lógicos fuzzy de Łukasiewicz
|
| 8 |
+
NOT(x) = 1 - sigmoid(x)
|
| 9 |
+
AND(x,y) = relu(sigmoid(x) + sigmoid(y) - 1)
|
| 10 |
+
OR(x,y) = min(1, sigmoid(x) + sigmoid(y))
|
| 11 |
+
IMP(x,y) = min(1, 1 - sigmoid(x) + sigmoid(y))
|
| 12 |
+
"""
|
| 13 |
+
from __future__ import annotations
|
| 14 |
+
|
| 15 |
+
import logging
|
| 16 |
+
import math
|
| 17 |
+
from typing import Optional, Tuple
|
| 18 |
+
|
| 19 |
+
import torch
|
| 20 |
+
import torch.nn as nn
|
| 21 |
+
import torch.nn.functional as F
|
| 22 |
+
|
| 23 |
+
from .cooperative_bigru import CooperativeCNNBiGRU
|
| 24 |
+
|
| 25 |
+
logger = logging.getLogger(__name__)
|
| 26 |
+
|
| 27 |
+
|
| 28 |
+
# ============================================================================
|
| 29 |
+
# Funções fuzzy de Łukasiewicz
|
| 30 |
+
# ============================================================================
|
| 31 |
+
|
| 32 |
+
def fuzzy_not(x: torch.Tensor) -> torch.Tensor:
|
| 33 |
+
"""NOT(x) = 1 - sigmoid(x)."""
|
| 34 |
+
return 1.0 - torch.sigmoid(x)
|
| 35 |
+
|
| 36 |
+
|
| 37 |
+
def fuzzy_and(x: torch.Tensor, y: torch.Tensor) -> torch.Tensor:
|
| 38 |
+
"""AND(x, y) = relu(sigmoid(x) + sigmoid(y) - 1)."""
|
| 39 |
+
s = torch.sigmoid(x) + torch.sigmoid(y)
|
| 40 |
+
return F.relu(s - 1.0)
|
| 41 |
+
|
| 42 |
+
|
| 43 |
+
def fuzzy_or(x: torch.Tensor, y: torch.Tensor) -> torch.Tensor:
|
| 44 |
+
"""OR(x, y) = min(1, sigmoid(x) + sigmoid(y))."""
|
| 45 |
+
s = torch.sigmoid(x) + torch.sigmoid(y)
|
| 46 |
+
return torch.clamp(s, 0.0, 1.0)
|
| 47 |
+
|
| 48 |
+
|
| 49 |
+
def fuzzy_imp(x: torch.Tensor, y: torch.Tensor) -> torch.Tensor:
|
| 50 |
+
"""IMP(x, y) = min(1, 1 - sigmoid(x) + sigmoid(y))."""
|
| 51 |
+
s = 1.0 - torch.sigmoid(x) + torch.sigmoid(y)
|
| 52 |
+
return torch.clamp(s, 0.0, 1.0)
|
| 53 |
+
|
| 54 |
+
|
| 55 |
+
# ============================================================================
|
| 56 |
+
# Camada Anti-Alucinação
|
| 57 |
+
# ============================================================================
|
| 58 |
+
|
| 59 |
+
class AntiHallucinationLayer(nn.Module):
|
| 60 |
+
"""Camada anti-alucinação com lógica fuzzy de Łukasiewicz.
|
| 61 |
+
|
| 62 |
+
Recebe os logits [B, V] de um passo gerado e calcula um valor de verdade
|
| 63 |
+
suave τ ∈ [0, 1] que indica o quão "lógico" é o passo.
|
| 64 |
+
|
| 65 |
+
Implementação diferenciável:
|
| 66 |
+
- Calcula sigmoid(logit) -> probabilidade soft
|
| 67 |
+
- Aplica operadores fuzzy em pares de tokens mais prováveis
|
| 68 |
+
- Combina resultados via IMP (implicação) e AND
|
| 69 |
+
"""
|
| 70 |
+
|
| 71 |
+
def __init__(self, vocab_size: int, embed_dim: int = 32):
|
| 72 |
+
super().__init__()
|
| 73 |
+
self.proj = nn.Linear(vocab_size, embed_dim)
|
| 74 |
+
self.combine = nn.Linear(embed_dim * 2, 1)
|
| 75 |
+
|
| 76 |
+
def forward(self, logits: torch.Tensor) -> torch.Tensor:
|
| 77 |
+
"""
|
| 78 |
+
Args:
|
| 79 |
+
logits: [B, V] logits do passo gerado
|
| 80 |
+
|
| 81 |
+
Returns:
|
| 82 |
+
tau: [B, 1] valor de verdade suave em [0, 1]
|
| 83 |
+
"""
|
| 84 |
+
# Probabilidades soft via sigmoid
|
| 85 |
+
probs = torch.sigmoid(logits) # [B, V]
|
| 86 |
+
|
| 87 |
+
# Top-2 tokens mais prováveis (diferenciável via topk)
|
| 88 |
+
top2_vals, top2_idx = torch.topk(probs, k=min(2, probs.size(-1)), dim=-1) # [B, 2]
|
| 89 |
+
|
| 90 |
+
# Logits dos top-2 tokens via gather
|
| 91 |
+
top1_logit = logits.gather(1, top2_idx[:, 0:1]) # [B, 1]
|
| 92 |
+
top2_logit = logits.gather(1, top2_idx[:, 1:2]) # [B, 1]
|
| 93 |
+
|
| 94 |
+
# Operadores fuzzy de Łukasiewicz sobre os top-2 logits
|
| 95 |
+
imp_val = fuzzy_imp(top1_logit, top2_logit) # [B, 1]
|
| 96 |
+
and_val = fuzzy_and(top1_logit, top2_logit) # [B, 1]
|
| 97 |
+
or_val = fuzzy_or(top1_logit, top2_logit) # [B, 1]
|
| 98 |
+
|
| 99 |
+
# Projeção do vetor de logits completo para embedding
|
| 100 |
+
proj = self.proj(logits) # [B, embed_dim]
|
| 101 |
+
|
| 102 |
+
# Modulariza a projeção pelos valores fuzzy (broadcasting)
|
| 103 |
+
fuzzy_factor = (imp_val + and_val + or_val) / 3.0 # [B, 1]
|
| 104 |
+
modulated = proj * fuzzy_factor # [B, embed_dim]
|
| 105 |
+
|
| 106 |
+
# Combina: [modulated, original_proj] -> [B, 2*embed_dim]
|
| 107 |
+
combined = torch.cat((modulated, proj), dim=-1) # [B, 2*embed_dim]
|
| 108 |
+
tau = torch.sigmoid(self.combine(combined)) # [B, 1]
|
| 109 |
+
return tau
|
| 110 |
+
|
| 111 |
+
|
| 112 |
+
# ============================================================================
|
| 113 |
+
# Gerador (CNN-BiGRU encoder + GRU decoder com atenção)
|
| 114 |
+
# ============================================================================
|
| 115 |
+
|
| 116 |
+
class GeneratorCNNBiGRU(nn.Module):
|
| 117 |
+
"""Gerador: encoder CNN-BiGRU + decoder GRU com atenção Bahdanau.
|
| 118 |
+
|
| 119 |
+
Args:
|
| 120 |
+
vocab_size, embedding_dim, cnn_filters, gru_hidden: herdados do núcleo.
|
| 121 |
+
max_proof_len: comprimento máximo da sequência gerada.
|
| 122 |
+
"""
|
| 123 |
+
|
| 124 |
+
def __init__(
|
| 125 |
+
self,
|
| 126 |
+
vocab_size: int,
|
| 127 |
+
embedding_dim: int = 64,
|
| 128 |
+
cnn_filters: int = 64,
|
| 129 |
+
gru_hidden: int = 64,
|
| 130 |
+
n_heads: int = 4,
|
| 131 |
+
dropout: float = 0.1,
|
| 132 |
+
pad_idx: int = 1,
|
| 133 |
+
max_proof_len: int = 32,
|
| 134 |
+
):
|
| 135 |
+
super().__init__()
|
| 136 |
+
self.pad_idx = pad_idx
|
| 137 |
+
self.vocab_size = vocab_size
|
| 138 |
+
self.gru_hidden = gru_hidden
|
| 139 |
+
self.max_proof_len = max_proof_len
|
| 140 |
+
|
| 141 |
+
# Encoder cooperativo
|
| 142 |
+
self.encoder = CooperativeCNNBiGRU(
|
| 143 |
+
vocab_size=vocab_size,
|
| 144 |
+
embedding_dim=embedding_dim,
|
| 145 |
+
cnn_filters=cnn_filters,
|
| 146 |
+
gru_hidden=gru_hidden,
|
| 147 |
+
n_heads=n_heads,
|
| 148 |
+
dropout=dropout,
|
| 149 |
+
pad_idx=pad_idx,
|
| 150 |
+
)
|
| 151 |
+
|
| 152 |
+
# Decoder GRU (unidirecional)
|
| 153 |
+
self.dec_embedding = nn.Embedding(vocab_size, embedding_dim, padding_idx=pad_idx)
|
| 154 |
+
self.dec_gru = nn.GRUCell(embedding_dim + gru_hidden * 2, gru_hidden)
|
| 155 |
+
|
| 156 |
+
# Atenção Bahdanau
|
| 157 |
+
self.attn_W = nn.Linear(gru_hidden * 2, gru_hidden, bias=False)
|
| 158 |
+
self.attn_U = nn.Linear(gru_hidden, gru_hidden, bias=False)
|
| 159 |
+
self.attn_v = nn.Linear(gru_hidden, 1, bias=False)
|
| 160 |
+
|
| 161 |
+
# Projeção final
|
| 162 |
+
self.out_proj = nn.Linear(gru_hidden + gru_hidden * 2 + embedding_dim, vocab_size)
|
| 163 |
+
|
| 164 |
+
self.bos_id = 0 # placeholder, será sobrescrito
|
| 165 |
+
self.eos_id = 2
|
| 166 |
+
|
| 167 |
+
def _attention(
|
| 168 |
+
self,
|
| 169 |
+
dec_state: torch.Tensor, # [B, H]
|
| 170 |
+
enc_out: torch.Tensor, # [B, T, 2H]
|
| 171 |
+
enc_mask: torch.Tensor, # [B, T]
|
| 172 |
+
) -> torch.Tensor:
|
| 173 |
+
"""Atenção Bahdanau: retorna contexto [B, 2H]."""
|
| 174 |
+
# score = v^T tanh(W enc + U dec)
|
| 175 |
+
proj_enc = self.attn_W(enc_out) # [B, T, H]
|
| 176 |
+
proj_dec = self.attn_U(dec_state).unsqueeze(1) # [B, 1, H]
|
| 177 |
+
scores = self.attn_v(torch.tanh(proj_enc + proj_dec)).squeeze(-1) # [B, T]
|
| 178 |
+
scores = scores.masked_fill(enc_mask == 0, -1e9)
|
| 179 |
+
attn = F.softmax(scores, dim=-1) # [B, T]
|
| 180 |
+
ctx = torch.bmm(attn.unsqueeze(1), enc_out).squeeze(1) # [B, 2H]
|
| 181 |
+
return ctx
|
| 182 |
+
|
| 183 |
+
def forward(
|
| 184 |
+
self,
|
| 185 |
+
input_ids_a: torch.Tensor,
|
| 186 |
+
input_ids_b: torch.Tensor,
|
| 187 |
+
target_proof: Optional[torch.Tensor] = None,
|
| 188 |
+
bos_id: int = 0,
|
| 189 |
+
eos_id: int = 2,
|
| 190 |
+
max_len: Optional[int] = None,
|
| 191 |
+
) -> dict:
|
| 192 |
+
"""
|
| 193 |
+
Args:
|
| 194 |
+
input_ids_a, input_ids_b: [B, T] entradas (axiomas + conjectura)
|
| 195 |
+
target_proof: [B, T_proof] para teacher forcing (treino)
|
| 196 |
+
bos_id, eos_id: IDs dos tokens especiais
|
| 197 |
+
|
| 198 |
+
Returns:
|
| 199 |
+
dict com:
|
| 200 |
+
logits: [B, T_proof, V]
|
| 201 |
+
enc_out: [B, T, 2H] — sequência REAL do encoder (não repetida)
|
| 202 |
+
dec_state: [B, H]
|
| 203 |
+
|
| 204 |
+
CORREÇÃO: O bug original era `enc_out = fused.unsqueeze(1).expand(-1, T_enc, -1)`
|
| 205 |
+
que repetia o mesmo vetor a cada timestep, tornando a Bahdanau attention
|
| 206 |
+
meaningless (todas as chaves idênticas → atenção uniforme).
|
| 207 |
+
Agora usamos a sequência temporal REAL produzida pelo encoder.
|
| 208 |
+
"""
|
| 209 |
+
# 1. Encoder — retornar sequências temporais para a atenção
|
| 210 |
+
enc = self.encoder(input_ids_a, input_ids_b, return_sequences=True)
|
| 211 |
+
|
| 212 |
+
# Sequência temporal real do stream A: [B, T_a, 2H]
|
| 213 |
+
# (resultado do self-attention summary que preserva a dimensão 2H)
|
| 214 |
+
# Mas self_attn_A retorna [B, 2H] (sumarizado), então usamos seq_A
|
| 215 |
+
# que contém os estados ocultos em cada timestep
|
| 216 |
+
if "seq_A" in enc:
|
| 217 |
+
# seq_A: [B, T, 2H] — estados BiGRU concatenados por timestep
|
| 218 |
+
enc_out = enc["seq_A"] # [B, T, 2H]
|
| 219 |
+
else:
|
| 220 |
+
# Fallback: usar out_A repetido (mas agora com aviso)
|
| 221 |
+
logger.warning("GeneratorCNNBiGRU: seq_A não disponível, usando fallback")
|
| 222 |
+
out_A = enc["out_A"] # [B, 2H]
|
| 223 |
+
T_enc = enc["mask_A"].size(1)
|
| 224 |
+
enc_out = out_A.unsqueeze(1).expand(-1, T_enc, -1)
|
| 225 |
+
|
| 226 |
+
mask_A = enc["mask_A"]
|
| 227 |
+
enc_mask = mask_A
|
| 228 |
+
|
| 229 |
+
# Fused para inicializar o decoder
|
| 230 |
+
fused = enc["fused"] # [B, 256] = [B, 4H] (2 streams * 2H)
|
| 231 |
+
B = input_ids_a.size(0)
|
| 232 |
+
device = input_ids_a.device
|
| 233 |
+
|
| 234 |
+
# 2. Inicialização do decoder
|
| 235 |
+
# Estado inicial: projeção linear do fused para gru_hidden
|
| 236 |
+
# (em vez de cortar a primeira metade, que perde metade da informação)
|
| 237 |
+
if not hasattr(self, "_init_proj"):
|
| 238 |
+
# Lazy init — evita quebrar a API existente
|
| 239 |
+
gru_in = fused.size(-1) # 256 = 4H
|
| 240 |
+
self._init_proj = nn.Linear(gru_in, self.gru_hidden).to(device)
|
| 241 |
+
nn.init.xavier_uniform_(self._init_proj.weight)
|
| 242 |
+
nn.init.zeros_(self._init_proj.bias)
|
| 243 |
+
h = torch.tanh(self._init_proj(fused)) # [B, H]
|
| 244 |
+
|
| 245 |
+
# 3. Token inicial = BOS
|
| 246 |
+
if target_proof is not None:
|
| 247 |
+
# Teacher forcing
|
| 248 |
+
T_dec = target_proof.size(1)
|
| 249 |
+
else:
|
| 250 |
+
T_dec = max_len or self.max_proof_len
|
| 251 |
+
|
| 252 |
+
logits_list = []
|
| 253 |
+
cur_token = torch.full((B, 1), bos_id, dtype=torch.long, device=device)
|
| 254 |
+
|
| 255 |
+
for t in range(T_dec):
|
| 256 |
+
if target_proof is not None:
|
| 257 |
+
tok = target_proof[:, t] # [B]
|
| 258 |
+
else:
|
| 259 |
+
tok = cur_token.squeeze(1)
|
| 260 |
+
|
| 261 |
+
emb = self.dec_embedding(tok) # [B, D]
|
| 262 |
+
ctx = self._attention(h, enc_out, enc_mask) # [B, 2H]
|
| 263 |
+
gru_input = torch.cat((emb, ctx), dim=-1) # [B, D+2H]
|
| 264 |
+
h = self.dec_gru(gru_input, h) # [B, H]
|
| 265 |
+
|
| 266 |
+
# Logits
|
| 267 |
+
logit_input = torch.cat((h, ctx, emb), dim=-1) # [B, H+2H+D]
|
| 268 |
+
logit = self.out_proj(logit_input) # [B, V]
|
| 269 |
+
logits_list.append(logit)
|
| 270 |
+
|
| 271 |
+
if target_proof is None:
|
| 272 |
+
cur_token = logit.argmax(dim=-1, keepdim=True)
|
| 273 |
+
|
| 274 |
+
logits = torch.stack(logits_list, dim=1) # [B, T_dec, V]
|
| 275 |
+
|
| 276 |
+
return {
|
| 277 |
+
"logits": logits,
|
| 278 |
+
"enc_out": enc_out,
|
| 279 |
+
"dec_state": h,
|
| 280 |
+
"fused": fused,
|
| 281 |
+
}
|
| 282 |
+
|
| 283 |
+
|
| 284 |
+
# ============================================================================
|
| 285 |
+
# Verificador (CNN-BiGRU + classificador sigmoid)
|
| 286 |
+
# ============================================================================
|
| 287 |
+
|
| 288 |
+
class VerifierCNNBiGRU(nn.Module):
|
| 289 |
+
"""Verificador: CNN-BiGRU + classificador sigmoid (v ∈ [0,1]).
|
| 290 |
+
|
| 291 |
+
Recebe premissas (concat de axiomas + conjectura) e um passo hipotético,
|
| 292 |
+
e retorna uma probabilidade de o passo estar correto.
|
| 293 |
+
"""
|
| 294 |
+
|
| 295 |
+
def __init__(
|
| 296 |
+
self,
|
| 297 |
+
vocab_size: int,
|
| 298 |
+
embedding_dim: int = 64,
|
| 299 |
+
cnn_filters: int = 64,
|
| 300 |
+
gru_hidden: int = 64,
|
| 301 |
+
n_heads: int = 4,
|
| 302 |
+
dropout: float = 0.1,
|
| 303 |
+
pad_idx: int = 1,
|
| 304 |
+
):
|
| 305 |
+
super().__init__()
|
| 306 |
+
self.pad_idx = pad_idx
|
| 307 |
+
self.gru_hidden = gru_hidden
|
| 308 |
+
|
| 309 |
+
self.encoder = CooperativeCNNBiGRU(
|
| 310 |
+
vocab_size=vocab_size,
|
| 311 |
+
embedding_dim=embedding_dim,
|
| 312 |
+
cnn_filters=cnn_filters,
|
| 313 |
+
gru_hidden=gru_hidden,
|
| 314 |
+
n_heads=n_heads,
|
| 315 |
+
dropout=dropout,
|
| 316 |
+
pad_idx=pad_idx,
|
| 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),
|
| 327 |
+
)
|
| 328 |
+
|
| 329 |
+
def forward(
|
| 330 |
+
self,
|
| 331 |
+
premissas_a: torch.Tensor,
|
| 332 |
+
premissas_b: torch.Tensor,
|
| 333 |
+
passo_a: torch.Tensor,
|
| 334 |
+
passo_b: torch.Tensor,
|
| 335 |
+
) -> torch.Tensor:
|
| 336 |
+
"""
|
| 337 |
+
Args:
|
| 338 |
+
premissas_a, premissas_b: [B, T] (contexto da prova até o momento)
|
| 339 |
+
passo_a, passo_b: [B, T_p] (passo gerado a ser validado)
|
| 340 |
+
|
| 341 |
+
Returns:
|
| 342 |
+
v: [B, 1] probabilidade de o passo estar correto
|
| 343 |
+
"""
|
| 344 |
+
# Encode premissas
|
| 345 |
+
enc_pre = self.encoder(premissas_a, premissas_b)
|
| 346 |
+
feat_pre = enc_pre["fused"] # [B, 256]
|
| 347 |
+
|
| 348 |
+
# Encode passo
|
| 349 |
+
enc_passo = self.encoder(passo_a, passo_b)
|
| 350 |
+
feat_passo = enc_passo["fused"] # [B, 256]
|
| 351 |
+
|
| 352 |
+
# Concatena e classifica
|
| 353 |
+
combined = torch.cat((feat_pre, feat_passo), dim=-1) # [B, 512]
|
| 354 |
+
logit = self.classifier(combined) # [B, 1]
|
| 355 |
+
v = torch.sigmoid(logit) # [B, 1]
|
| 356 |
+
return v
|
| 357 |
+
|
| 358 |
+
|
| 359 |
+
__all__ = [
|
| 360 |
+
"fuzzy_not", "fuzzy_and", "fuzzy_or", "fuzzy_imp",
|
| 361 |
+
"AntiHallucinationLayer",
|
| 362 |
+
"GeneratorCNNBiGRU",
|
| 363 |
+
"VerifierCNNBiGRU",
|
| 364 |
+
]
|
cnn_bigru/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 |
+
]
|
cnn_bigru/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 |
+
]
|
cnn_bigru/models/multimodal_encoders.py
ADDED
|
@@ -0,0 +1,141 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""multimodal_encoders.py — Encoders para imagem e áudio.
|
| 2 |
+
|
| 3 |
+
Para o projeto CNN-BiGRU multimodal, precisamos de encoders que produzam
|
| 4 |
+
sequências de features compatíveis com o núcleo cooperativo CNN-BiGRU.
|
| 5 |
+
|
| 6 |
+
- ImageEncoder: CNN 2D -> flatten espacial -> projeção linear -> [B, T_img, C]
|
| 7 |
+
- AudioEncoder: CNN 1D sobre espectrograma -> [B, T_aud, C]
|
| 8 |
+
|
| 9 |
+
A fusão multimodal acontece trocando o stream B do núcleo cooperativo pela
|
| 10 |
+
saída do encoder visual/áudio, ou adicionando-os como streams extras numa
|
| 11 |
+
versão estendida. Aqui usamos uma fusão por concatenação + projeção.
|
| 12 |
+
"""
|
| 13 |
+
from __future__ import annotations
|
| 14 |
+
|
| 15 |
+
import logging
|
| 16 |
+
from typing import Optional
|
| 17 |
+
|
| 18 |
+
import torch
|
| 19 |
+
import torch.nn as nn
|
| 20 |
+
import torch.nn.functional as F
|
| 21 |
+
|
| 22 |
+
logger = logging.getLogger(__name__)
|
| 23 |
+
|
| 24 |
+
|
| 25 |
+
class ImageEncoder(nn.Module):
|
| 26 |
+
"""CNN 2D simples que codifica imagem em sequência de features.
|
| 27 |
+
|
| 28 |
+
Input: [B, C, H, W]
|
| 29 |
+
Output: [B, T_img, out_dim] onde T_img = (H/4) * (W/4)
|
| 30 |
+
|
| 31 |
+
Usa 2 blocos Conv2d+ReLU+MaxPool para reduzir dimensão espacial por 4.
|
| 32 |
+
"""
|
| 33 |
+
|
| 34 |
+
def __init__(
|
| 35 |
+
self,
|
| 36 |
+
in_channels: int = 1,
|
| 37 |
+
out_dim: int = 64,
|
| 38 |
+
hidden_dim: int = 32,
|
| 39 |
+
):
|
| 40 |
+
super().__init__()
|
| 41 |
+
self.conv1 = nn.Conv2d(in_channels, hidden_dim, kernel_size=3, padding=1)
|
| 42 |
+
self.conv2 = nn.Conv2d(hidden_dim, hidden_dim * 2, kernel_size=3, padding=1)
|
| 43 |
+
self.pool = nn.MaxPool2d(2)
|
| 44 |
+
self.proj = nn.Linear(hidden_dim * 2, out_dim)
|
| 45 |
+
self.out_dim = out_dim
|
| 46 |
+
|
| 47 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 48 |
+
# x: [B, C, H, W]
|
| 49 |
+
x = F.relu(self.conv1(x))
|
| 50 |
+
x = self.pool(x) # [B, C', H/2, W/2]
|
| 51 |
+
x = F.relu(self.conv2(x))
|
| 52 |
+
x = self.pool(x) # [B, C'', H/4, W/4]
|
| 53 |
+
B, C, H, W = x.shape
|
| 54 |
+
x = x.permute(0, 2, 3, 1).contiguous().view(B, H * W, C) # [B, T, C]
|
| 55 |
+
x = self.proj(x) # [B, T, out_dim]
|
| 56 |
+
return x
|
| 57 |
+
|
| 58 |
+
|
| 59 |
+
class AudioEncoder(nn.Module):
|
| 60 |
+
"""Encoder de áudio: CNN 1D sobre espectrograma (T freq bins como canais).
|
| 61 |
+
|
| 62 |
+
Input: [B, 1, T_time, F_freq]
|
| 63 |
+
Output: [B, T_time, out_dim]
|
| 64 |
+
|
| 65 |
+
Trata a dimensão de frequência como canais e processa a dimensão temporal
|
| 66 |
+
com Conv1d, produzindo uma sequência compatível com BiGRU.
|
| 67 |
+
"""
|
| 68 |
+
|
| 69 |
+
def __init__(
|
| 70 |
+
self,
|
| 71 |
+
n_freq: int = 40,
|
| 72 |
+
out_dim: int = 64,
|
| 73 |
+
hidden_dim: int = 32,
|
| 74 |
+
):
|
| 75 |
+
super().__init__()
|
| 76 |
+
self.conv1 = nn.Conv1d(n_freq, hidden_dim, kernel_size=3, padding=1)
|
| 77 |
+
self.conv2 = nn.Conv1d(hidden_dim, hidden_dim * 2, kernel_size=3, padding=1)
|
| 78 |
+
self.proj = nn.Linear(hidden_dim * 2, out_dim)
|
| 79 |
+
self.out_dim = out_dim
|
| 80 |
+
|
| 81 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 82 |
+
# x: [B, 1, T, F] -> queremos [B, F, T] para Conv1d
|
| 83 |
+
B, _, T, F_ = x.shape
|
| 84 |
+
x = x.view(B, F_, T) # [B, F, T]
|
| 85 |
+
x = F.relu(self.conv1(x)) # [B, hidden, T]
|
| 86 |
+
x = F.relu(self.conv2(x)) # [B, hidden*2, T]
|
| 87 |
+
x = x.permute(0, 2, 1) # [B, T, hidden*2]
|
| 88 |
+
x = self.proj(x) # [B, T, out_dim]
|
| 89 |
+
return x
|
| 90 |
+
|
| 91 |
+
|
| 92 |
+
class MultimodalFusion(nn.Module):
|
| 93 |
+
"""Fusão de features de texto, imagem e áudio.
|
| 94 |
+
|
| 95 |
+
Recebe três tensores [B, D_text], [B, D_img], [B, D_aud] e produz
|
| 96 |
+
um único tensor [B, out_dim] usando concatenação + projeção +
|
| 97 |
+
uma porta de fusão (gated) que aprende a importância de cada modalidade.
|
| 98 |
+
"""
|
| 99 |
+
|
| 100 |
+
def __init__(
|
| 101 |
+
self,
|
| 102 |
+
d_text: int,
|
| 103 |
+
d_img: int,
|
| 104 |
+
d_aud: int,
|
| 105 |
+
out_dim: int = 128,
|
| 106 |
+
dropout: float = 0.1,
|
| 107 |
+
):
|
| 108 |
+
super().__init__()
|
| 109 |
+
d_concat = d_text + d_img + d_aud
|
| 110 |
+
self.proj = nn.Linear(d_concat, out_dim)
|
| 111 |
+
# Porta de fusão: aprende pesos por modalidade
|
| 112 |
+
self.gate_text = nn.Linear(d_text, 1)
|
| 113 |
+
self.gate_img = nn.Linear(d_img, 1)
|
| 114 |
+
self.gate_aud = nn.Linear(d_aud, 1)
|
| 115 |
+
self.dropout = nn.Dropout(dropout)
|
| 116 |
+
self.out_dim = out_dim
|
| 117 |
+
|
| 118 |
+
def forward(
|
| 119 |
+
self,
|
| 120 |
+
text_feat: torch.Tensor,
|
| 121 |
+
img_feat: torch.Tensor,
|
| 122 |
+
aud_feat: torch.Tensor,
|
| 123 |
+
) -> torch.Tensor:
|
| 124 |
+
# Gates (Sigmoid -> [0, 1])
|
| 125 |
+
g_t = torch.sigmoid(self.gate_text(text_feat)) # [B, 1]
|
| 126 |
+
g_i = torch.sigmoid(self.gate_img(img_feat))
|
| 127 |
+
g_a = torch.sigmoid(self.gate_aud(aud_feat))
|
| 128 |
+
|
| 129 |
+
# Aplica gates como máscara suave
|
| 130 |
+
text_g = text_feat * g_t
|
| 131 |
+
img_g = img_feat * g_i
|
| 132 |
+
aud_g = aud_feat * g_a
|
| 133 |
+
|
| 134 |
+
# Concatena e projeta
|
| 135 |
+
concat = torch.cat((text_g, img_g, aud_g), dim=-1)
|
| 136 |
+
out = self.proj(concat)
|
| 137 |
+
out = self.dropout(out)
|
| 138 |
+
return out
|
| 139 |
+
|
| 140 |
+
|
| 141 |
+
__all__ = ["ImageEncoder", "AudioEncoder", "MultimodalFusion"]
|
cnn_bigru/models/multimodal_model.py
ADDED
|
@@ -0,0 +1,214 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""multimodal_model.py — Modelo multimodal CNN-BiGRU completo.
|
| 2 |
+
|
| 3 |
+
Combina:
|
| 4 |
+
- CooperativeCNNBiGRU para texto (dual-stream A/B)
|
| 5 |
+
- ImageEncoder para modalidade visual
|
| 6 |
+
- AudioEncoder para modalidade auditiva
|
| 7 |
+
- MultimodalFusion com gated fusion
|
| 8 |
+
- Cabeçalho de classificação + cabeçalho de geração (LM head)
|
| 9 |
+
|
| 10 |
+
Suporta dois modos:
|
| 11 |
+
mode="classify": retorna logits de classificação [B, num_classes]
|
| 12 |
+
mode="generate": retorna logits de LM [B, T, vocab_size] (autoregressivo)
|
| 13 |
+
"""
|
| 14 |
+
from __future__ import annotations
|
| 15 |
+
|
| 16 |
+
import logging
|
| 17 |
+
from typing import Optional
|
| 18 |
+
|
| 19 |
+
import torch
|
| 20 |
+
import torch.nn as nn
|
| 21 |
+
import torch.nn.functional as F
|
| 22 |
+
|
| 23 |
+
from .cooperative_bigru import CooperativeCNNBiGRU
|
| 24 |
+
from .multimodal_encoders import ImageEncoder, AudioEncoder, MultimodalFusion
|
| 25 |
+
|
| 26 |
+
logger = logging.getLogger(__name__)
|
| 27 |
+
|
| 28 |
+
|
| 29 |
+
class MultimodalCNNBiGRU(nn.Module):
|
| 30 |
+
"""Modelo multimodal completo: texto (dual A/B) + imagem + áudio.
|
| 31 |
+
|
| 32 |
+
Fluxo:
|
| 33 |
+
1. Texto A e Texto B -> CooperativeCNNBiGRU -> fused [B, 256]
|
| 34 |
+
2. Imagem -> ImageEncoder -> mean-pool -> [B, D_img]
|
| 35 |
+
3. Áudio -> AudioEncoder -> mean-pool -> [B, D_aud]
|
| 36 |
+
4. Fusão multimodal -> [B, out_dim]
|
| 37 |
+
5. Classificação: Linear(out_dim, num_classes)
|
| 38 |
+
6. Geração (opcional): Linear(out_dim, vocab_size) autoregressivo
|
| 39 |
+
"""
|
| 40 |
+
|
| 41 |
+
def __init__(
|
| 42 |
+
self,
|
| 43 |
+
vocab_size: int,
|
| 44 |
+
num_classes: int = 3,
|
| 45 |
+
embedding_dim: int = 64,
|
| 46 |
+
cnn_filters: int = 64,
|
| 47 |
+
gru_hidden: int = 64,
|
| 48 |
+
n_heads: int = 4,
|
| 49 |
+
dropout: float = 0.1,
|
| 50 |
+
pad_idx: int = 1,
|
| 51 |
+
img_channels: int = 1,
|
| 52 |
+
img_hidden: int = 32,
|
| 53 |
+
img_out_dim: int = 64,
|
| 54 |
+
audio_freq: int = 40,
|
| 55 |
+
audio_hidden: int = 32,
|
| 56 |
+
audio_out_dim: int = 64,
|
| 57 |
+
fusion_dim: int = 128,
|
| 58 |
+
use_spectral_norm: bool = False,
|
| 59 |
+
weight_tying: bool = False,
|
| 60 |
+
):
|
| 61 |
+
super().__init__()
|
| 62 |
+
self.pad_idx = pad_idx
|
| 63 |
+
self.vocab_size = vocab_size
|
| 64 |
+
self.num_classes = num_classes
|
| 65 |
+
self.weight_tying = weight_tying
|
| 66 |
+
|
| 67 |
+
# Núcleo cooperativo de texto
|
| 68 |
+
self.text_core = CooperativeCNNBiGRU(
|
| 69 |
+
vocab_size=vocab_size,
|
| 70 |
+
embedding_dim=embedding_dim,
|
| 71 |
+
cnn_filters=cnn_filters,
|
| 72 |
+
gru_hidden=gru_hidden,
|
| 73 |
+
n_heads=n_heads,
|
| 74 |
+
dropout=dropout,
|
| 75 |
+
pad_idx=pad_idx,
|
| 76 |
+
use_spectral_norm=use_spectral_norm,
|
| 77 |
+
)
|
| 78 |
+
d_text = self.text_core.feat_dim * 2 # 256 (fused A+B = 2 * feat_dim_per_stream)
|
| 79 |
+
|
| 80 |
+
# Encoders multimodais
|
| 81 |
+
self.image_encoder = ImageEncoder(
|
| 82 |
+
in_channels=img_channels,
|
| 83 |
+
out_dim=img_out_dim,
|
| 84 |
+
hidden_dim=img_hidden,
|
| 85 |
+
)
|
| 86 |
+
self.audio_encoder = AudioEncoder(
|
| 87 |
+
n_freq=audio_freq,
|
| 88 |
+
out_dim=audio_out_dim,
|
| 89 |
+
hidden_dim=audio_hidden,
|
| 90 |
+
)
|
| 91 |
+
|
| 92 |
+
# Pooling de imagem/áudio para vetor fixo
|
| 93 |
+
self.img_proj = nn.Linear(img_out_dim, img_out_dim)
|
| 94 |
+
self.aud_proj = nn.Linear(audio_out_dim, audio_out_dim)
|
| 95 |
+
|
| 96 |
+
# Fusão multimodal
|
| 97 |
+
self.fusion = MultimodalFusion(
|
| 98 |
+
d_text=d_text,
|
| 99 |
+
d_img=img_out_dim,
|
| 100 |
+
d_aud=audio_out_dim,
|
| 101 |
+
out_dim=fusion_dim,
|
| 102 |
+
dropout=dropout,
|
| 103 |
+
)
|
| 104 |
+
|
| 105 |
+
# Cabeçalho de classificação
|
| 106 |
+
self.classifier = nn.Linear(fusion_dim, num_classes)
|
| 107 |
+
|
| 108 |
+
# Cabeçalho de geração (LM head) — para modo autoregressivo
|
| 109 |
+
self.lm_head = nn.Linear(fusion_dim, vocab_size, bias=False)
|
| 110 |
+
|
| 111 |
+
# Weight tying opcional (dados.txt linha 632):
|
| 112 |
+
# compartilha pesos entre token_embedding e lm_head.
|
| 113 |
+
# Como lm_head é Linear(fusion_dim, vocab_size) e embedding é
|
| 114 |
+
# Linear(vocab_size, embedding_dim), o tying direto só funciona se
|
| 115 |
+
# fusion_dim == embedding_dim. Caso contrário, usamos uma matriz
|
| 116 |
+
# intermediária de projeção.
|
| 117 |
+
if weight_tying:
|
| 118 |
+
if fusion_dim == self.text_core.embedding.embedding_dim:
|
| 119 |
+
self.lm_head.weight = self.text_core.embedding.weight
|
| 120 |
+
else:
|
| 121 |
+
# Projeção intermediária: embedding_dim -> fusion_dim
|
| 122 |
+
self.tie_proj = nn.Linear(
|
| 123 |
+
self.text_core.embedding.embedding_dim, fusion_dim, bias=False
|
| 124 |
+
)
|
| 125 |
+
# tying via forward (ver mode="generate")
|
| 126 |
+
logger.info(
|
| 127 |
+
f"Weight tying: fusion_dim ({fusion_dim}) != embedding_dim "
|
| 128 |
+
f"({self.text_core.embedding.embedding_dim}); usando proj intermediária."
|
| 129 |
+
)
|
| 130 |
+
|
| 131 |
+
# Dropout final
|
| 132 |
+
self.dropout = nn.Dropout(dropout)
|
| 133 |
+
|
| 134 |
+
def forward(
|
| 135 |
+
self,
|
| 136 |
+
input_ids_a: torch.Tensor,
|
| 137 |
+
input_ids_b: torch.Tensor,
|
| 138 |
+
images: Optional[torch.Tensor] = None,
|
| 139 |
+
audios: Optional[torch.Tensor] = None,
|
| 140 |
+
attn_mask_a: Optional[torch.Tensor] = None,
|
| 141 |
+
attn_mask_b: Optional[torch.Tensor] = None,
|
| 142 |
+
mode: str = "classify",
|
| 143 |
+
) -> dict:
|
| 144 |
+
"""
|
| 145 |
+
Args:
|
| 146 |
+
input_ids_a, input_ids_b: [B, T] tokens dos streams A e B
|
| 147 |
+
images: [B, C, H, W] ou None
|
| 148 |
+
audios: [B, 1, T, F] ou None
|
| 149 |
+
mode: "classify" | "generate"
|
| 150 |
+
|
| 151 |
+
Returns:
|
| 152 |
+
dict com logits (classificação ou geração) + features intermediárias
|
| 153 |
+
"""
|
| 154 |
+
# 1. Núcleo cooperativo de texto
|
| 155 |
+
# As máscaras são derivadas internamente do pad_idx, mas se o usuário
|
| 156 |
+
# fornecer attn_mask_a/attn_mask_b explicitamente, elas têm precedência
|
| 157 |
+
# e são usadas para ignorar tokens de padding no cálculo de atenção.
|
| 158 |
+
text_out = self.text_core(input_ids_a, input_ids_b, return_sequences=(mode == "generate"))
|
| 159 |
+
fused_text = text_out["fused"] # [B, 256]
|
| 160 |
+
|
| 161 |
+
# Sobrescrever máscaras se fornecidas externamente (para uso downstream)
|
| 162 |
+
if attn_mask_a is not None:
|
| 163 |
+
text_out["mask_A"] = attn_mask_a.to(text_out["mask_A"].device).float()
|
| 164 |
+
if attn_mask_b is not None:
|
| 165 |
+
text_out["mask_B"] = attn_mask_b.to(text_out["mask_B"].device).float()
|
| 166 |
+
|
| 167 |
+
# 2. Imagem
|
| 168 |
+
if images is not None:
|
| 169 |
+
img_seq = self.image_encoder(images) # [B, T_img, D_img]
|
| 170 |
+
img_feat = img_seq.mean(dim=1) # [B, D_img]
|
| 171 |
+
img_feat = self.img_proj(img_feat)
|
| 172 |
+
else:
|
| 173 |
+
# Vetor zero se não houver modalidade
|
| 174 |
+
img_feat = torch.zeros(
|
| 175 |
+
fused_text.size(0), self.img_proj.out_features,
|
| 176 |
+
device=fused_text.device, dtype=fused_text.dtype,
|
| 177 |
+
)
|
| 178 |
+
|
| 179 |
+
# 3. Áudio
|
| 180 |
+
if audios is not None:
|
| 181 |
+
aud_seq = self.audio_encoder(audios) # [B, T_aud, D_aud]
|
| 182 |
+
aud_feat = aud_seq.mean(dim=1)
|
| 183 |
+
aud_feat = self.aud_proj(aud_feat)
|
| 184 |
+
else:
|
| 185 |
+
aud_feat = torch.zeros(
|
| 186 |
+
fused_text.size(0), self.aud_proj.out_features,
|
| 187 |
+
device=fused_text.device, dtype=fused_text.dtype,
|
| 188 |
+
)
|
| 189 |
+
|
| 190 |
+
# 4. Fusão multimodal
|
| 191 |
+
fused = self.fusion(fused_text, img_feat, aud_feat) # [B, fusion_dim]
|
| 192 |
+
fused = self.dropout(fused)
|
| 193 |
+
|
| 194 |
+
result = {
|
| 195 |
+
"fused": fused,
|
| 196 |
+
"text_out": text_out,
|
| 197 |
+
"img_feat": img_feat,
|
| 198 |
+
"aud_feat": aud_feat,
|
| 199 |
+
}
|
| 200 |
+
|
| 201 |
+
# 5. Cabeçalhos
|
| 202 |
+
if mode == "classify":
|
| 203 |
+
result["logits"] = self.classifier(fused) # [B, num_classes]
|
| 204 |
+
elif mode == "generate":
|
| 205 |
+
# Modo geração: projetar fusão para vocab_size
|
| 206 |
+
# (modo simplificado: geração single-step baseada na fusão)
|
| 207 |
+
result["logits"] = self.lm_head(fused) # [B, vocab_size]
|
| 208 |
+
else:
|
| 209 |
+
raise ValueError(f"mode deve ser 'classify' ou 'generate', recebeu '{mode}'")
|
| 210 |
+
|
| 211 |
+
return result
|
| 212 |
+
|
| 213 |
+
|
| 214 |
+
__all__ = ["MultimodalCNNBiGRU"]
|
cnn_bigru/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 |
+
]
|
cnn_bigru/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 |
+
]
|
cnn_bigru/models/rope.py
ADDED
|
@@ -0,0 +1,196 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
RoPE (Rotary Position Embeddings) — Codificação posicional rotacional.
|
| 3 |
+
|
| 4 |
+
Referência: Su et al., 2021 — "RoFormer: Enhanced Transformer with Rotary
|
| 5 |
+
Position Embedding".
|
| 6 |
+
|
| 7 |
+
Propriedade chave: <RoPE(q, m), RoPE(k, n)> = <RoPE(q, m-n), k>
|
| 8 |
+
→ a atenção depende apenas da diferença posicional relativa.
|
| 9 |
+
|
| 10 |
+
Ver docs/MATH_ANALYSIS.md seção 4.1.
|
| 11 |
+
|
| 12 |
+
Autor: CNN-BiGRU Project
|
| 13 |
+
"""
|
| 14 |
+
from __future__ import annotations
|
| 15 |
+
|
| 16 |
+
import math
|
| 17 |
+
from typing import Optional, Tuple
|
| 18 |
+
|
| 19 |
+
import torch
|
| 20 |
+
import torch.nn as nn
|
| 21 |
+
|
| 22 |
+
|
| 23 |
+
def precompute_rope_frequencies(
|
| 24 |
+
head_dim: int,
|
| 25 |
+
max_seq_len: int,
|
| 26 |
+
base: float = 10000.0,
|
| 27 |
+
device: Optional[torch.device] = None,
|
| 28 |
+
dtype: torch.dtype = torch.float32,
|
| 29 |
+
) -> torch.Tensor:
|
| 30 |
+
"""
|
| 31 |
+
Pré-computa as frequências angulares theta_i para RoPE.
|
| 32 |
+
|
| 33 |
+
Fórmula: theta_i = base^(-2i / head_dim) para i = 0, 1, ..., head_dim/2 - 1
|
| 34 |
+
|
| 35 |
+
Args:
|
| 36 |
+
head_dim: dimensão de cada cabeça (deve ser par)
|
| 37 |
+
max_seq_len: comprimento máximo da sequência
|
| 38 |
+
base: base da exponencial (padrão 10000)
|
| 39 |
+
device: device do tensor
|
| 40 |
+
dtype: dtype do tensor
|
| 41 |
+
|
| 42 |
+
Returns:
|
| 43 |
+
cos_sin: tensor [max_seq_len, head_dim] com [cos(theta_0*m), cos(theta_1*m), ...,
|
| 44 |
+
sin(theta_0*m), sin(theta_1*m), ...] intercalados por par
|
| 45 |
+
(formato compatível com apply_rope).
|
| 46 |
+
|
| 47 |
+
Matematicamente:
|
| 48 |
+
Para cada posição m em [0, max_seq_len) e cada par de dimensões 2i:
|
| 49 |
+
angle = m * theta_i = m * base^(-2i/head_dim)
|
| 50 |
+
A matriz de rotação no plano (2i, 2i+1) é:
|
| 51 |
+
[[cos(angle), -sin(angle)],
|
| 52 |
+
[sin(angle), cos(angle)]]
|
| 53 |
+
"""
|
| 54 |
+
if head_dim % 2 != 0:
|
| 55 |
+
raise ValueError(f"head_dim deve ser par; recebido {head_dim}")
|
| 56 |
+
|
| 57 |
+
# theta_i = base^(-2i / head_dim), i = 0, ..., head_dim/2 - 1
|
| 58 |
+
i = torch.arange(0, head_dim, 2, device=device, dtype=dtype) # [head_dim/2]
|
| 59 |
+
theta = 1.0 / (base ** (i / head_dim)) # [head_dim/2]
|
| 60 |
+
|
| 61 |
+
# Posições m
|
| 62 |
+
positions = torch.arange(max_seq_len, device=device, dtype=dtype) # [max_seq_len]
|
| 63 |
+
|
| 64 |
+
# angle[m, i] = m * theta_i
|
| 65 |
+
angles = torch.outer(positions, theta) # [max_seq_len, head_dim/2]
|
| 66 |
+
|
| 67 |
+
# Duplicar para intercalar: [cos_0, cos_1, ..., sin_0, sin_1, ...]
|
| 68 |
+
cos = torch.cos(angles) # [max_seq_len, head_dim/2]
|
| 69 |
+
sin = torch.sin(angles)
|
| 70 |
+
|
| 71 |
+
# Formato final: [max_seq_len, head_dim] com pares (cos_i, sin_i)
|
| 72 |
+
cos_sin = torch.stack([cos, sin], dim=-1).reshape(max_seq_len, head_dim)
|
| 73 |
+
return cos_sin
|
| 74 |
+
|
| 75 |
+
|
| 76 |
+
def apply_rope(
|
| 77 |
+
x: torch.Tensor,
|
| 78 |
+
cos_sin: torch.Tensor,
|
| 79 |
+
) -> torch.Tensor:
|
| 80 |
+
"""
|
| 81 |
+
Aplica RoPE ao tensor x.
|
| 82 |
+
|
| 83 |
+
Args:
|
| 84 |
+
x: [batch, n_heads, seq_len, head_dim]
|
| 85 |
+
cos_sin: [seq_len, head_dim] — tensor pré-computado
|
| 86 |
+
|
| 87 |
+
Returns:
|
| 88 |
+
x_rotated: [batch, n_heads, seq_len, head_dim]
|
| 89 |
+
|
| 90 |
+
Matemática (para cada batch, head, posição m, par de dims (2i, 2i+1)):
|
| 91 |
+
x'[2i] = x[2i] * cos(angle) - x[2i+1] * sin(angle)
|
| 92 |
+
x'[2i+1] = x[2i] * sin(angle) + x[2i+1] * cos(angle)
|
| 93 |
+
|
| 94 |
+
O tensor cos_sin tem formato [seq_len, head_dim] com:
|
| 95 |
+
cos_sin[m, 2i] = cos(m * theta_i)
|
| 96 |
+
cos_sin[m, 2i+1] = sin(m * theta_i)
|
| 97 |
+
"""
|
| 98 |
+
if x.dim() != 4:
|
| 99 |
+
raise ValueError(
|
| 100 |
+
f"x deve ser [batch, n_heads, seq_len, head_dim]; recebido {x.shape}"
|
| 101 |
+
)
|
| 102 |
+
|
| 103 |
+
B, H, T, D = x.shape
|
| 104 |
+
if D % 2 != 0:
|
| 105 |
+
raise ValueError(f"head_dim deve ser par; recebido {D}")
|
| 106 |
+
|
| 107 |
+
# Reshape x para [B, H, T, D/2, 2] — pares de dimensões
|
| 108 |
+
x_pairs = x.float().reshape(B, H, T, D // 2, 2)
|
| 109 |
+
|
| 110 |
+
# Reshape cos_sin para [T, D/2, 2]
|
| 111 |
+
cos_sin_t = cos_sin[:T].to(x.device).float().reshape(T, D // 2, 2)
|
| 112 |
+
cos = cos_sin_t[..., 0] # [T, D/2]
|
| 113 |
+
sin = cos_sin_t[..., 1]
|
| 114 |
+
|
| 115 |
+
# Broadcast: [1, 1, T, D/2]
|
| 116 |
+
cos = cos.unsqueeze(0).unsqueeze(0)
|
| 117 |
+
sin = sin.unsqueeze(0).unsqueeze(0)
|
| 118 |
+
|
| 119 |
+
# Aplicar rotação
|
| 120 |
+
x_real = x_pairs[..., 0] # [B, H, T, D/2]
|
| 121 |
+
x_imag = x_pairs[..., 1]
|
| 122 |
+
rotated_real = x_real * cos - x_imag * sin
|
| 123 |
+
rotated_imag = x_real * sin + x_imag * cos
|
| 124 |
+
|
| 125 |
+
# Reagrupar para [B, H, T, D]
|
| 126 |
+
rotated = torch.stack([rotated_real, rotated_imag], dim=-1).reshape(B, H, T, D)
|
| 127 |
+
return rotated.to(x.dtype)
|
| 128 |
+
|
| 129 |
+
|
| 130 |
+
class RotaryPositionEmbedding(nn.Module):
|
| 131 |
+
"""
|
| 132 |
+
Módulo PyTorch para RoPE — mantém o cache de frequências como buffer.
|
| 133 |
+
|
| 134 |
+
Uso:
|
| 135 |
+
rope = RotaryPositionEmbedding(head_dim=64, max_seq_len=512)
|
| 136 |
+
# Em forward:
|
| 137 |
+
q = rope(q) # q: [batch, n_heads, seq_len, head_dim]
|
| 138 |
+
k = rope(k)
|
| 139 |
+
"""
|
| 140 |
+
|
| 141 |
+
def __init__(self, head_dim: int, max_seq_len: int = 512, base: float = 10000.0):
|
| 142 |
+
super().__init__()
|
| 143 |
+
if head_dim % 2 != 0:
|
| 144 |
+
raise ValueError(f"head_dim deve ser par; recebido {head_dim}")
|
| 145 |
+
self.head_dim = head_dim
|
| 146 |
+
self.max_seq_len = max_seq_len
|
| 147 |
+
self.base = base
|
| 148 |
+
|
| 149 |
+
# Pré-computar e registrar como buffer (move com .to(device))
|
| 150 |
+
cos_sin = precompute_rope_frequencies(head_dim, max_seq_len, base)
|
| 151 |
+
self.register_buffer("cos_sin", cos_sin, persistent=False)
|
| 152 |
+
|
| 153 |
+
def forward(self, x: torch.Tensor, offset: int = 0) -> torch.Tensor:
|
| 154 |
+
"""
|
| 155 |
+
Aplica RoPE a x.
|
| 156 |
+
|
| 157 |
+
Args:
|
| 158 |
+
x: [batch, n_heads, seq_len, head_dim]
|
| 159 |
+
offset: deslocamento posicional (para geração autoregressiva com cache KV,
|
| 160 |
+
onde a posição do novo token é len(cached) + i)
|
| 161 |
+
|
| 162 |
+
Returns:
|
| 163 |
+
x_rotated: mesmo shape que x
|
| 164 |
+
"""
|
| 165 |
+
if x.dim() != 4:
|
| 166 |
+
raise ValueError(
|
| 167 |
+
f"x deve ser [batch, n_heads, seq_len, head_dim]; recebido {x.shape}"
|
| 168 |
+
)
|
| 169 |
+
T = x.size(2)
|
| 170 |
+
if offset + T > self.max_seq_len:
|
| 171 |
+
# Estender o cache dinamicamente
|
| 172 |
+
self._extend_cache(offset + T)
|
| 173 |
+
# Slice das posições relevantes
|
| 174 |
+
cos_sin = self.cos_sin[offset:offset + T]
|
| 175 |
+
return apply_rope(x, cos_sin)
|
| 176 |
+
|
| 177 |
+
def _extend_cache(self, new_max: int) -> None:
|
| 178 |
+
"""Estende o cache de frequências para new_max posições."""
|
| 179 |
+
new_max_pow2 = 1
|
| 180 |
+
while new_max_pow2 < new_max:
|
| 181 |
+
new_max_pow2 *= 2
|
| 182 |
+
new_cos_sin = precompute_rope_frequencies(
|
| 183 |
+
self.head_dim, new_max_pow2, self.base,
|
| 184 |
+
device=self.cos_sin.device, dtype=self.cos_sin.dtype
|
| 185 |
+
)
|
| 186 |
+
self.cos_sin = new_cos_sin.to(self.cos_sin.device)
|
| 187 |
+
self.max_seq_len = new_max_pow2
|
| 188 |
+
# Re-registrar como buffer (substitui o anterior)
|
| 189 |
+
self.register_buffer("cos_sin", self.cos_sin, persistent=False)
|
| 190 |
+
|
| 191 |
+
|
| 192 |
+
__all__ = [
|
| 193 |
+
"precompute_rope_frequencies",
|
| 194 |
+
"apply_rope",
|
| 195 |
+
"RotaryPositionEmbedding",
|
| 196 |
+
]
|
cnn_bigru/models/transformer_block.py
ADDED
|
@@ -0,0 +1,401 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
Transformer Decoder Block — CausalSelfAttention + TransformerBlock.
|
| 3 |
+
|
| 4 |
+
Implementa o bloco decodificador causal conforme dados.txt seções 6-7:
|
| 5 |
+
- Multi-head self-attention com máscara triangular inferior
|
| 6 |
+
- Scaled dot-product: (q @ k.T) / sqrt(d_head)
|
| 7 |
+
- QKV projection unificada
|
| 8 |
+
- Pre-LN: x = x + attn(ln1(x)); x = x + ffn(ln2(x))
|
| 9 |
+
- FFN: Linear(d, 4d) → GELU → Linear(4d, d) → Dropout
|
| 10 |
+
- Weight tying opcional entre embedding e LM head
|
| 11 |
+
|
| 12 |
+
Inclui:
|
| 13 |
+
- RoPE (Rotary Position Embeddings) opcional
|
| 14 |
+
- Cache KV opcional para geração autoregressiva eficiente
|
| 15 |
+
- Padding mask (ignora tokens [PAD])
|
| 16 |
+
- Causal mask (não olha para o futuro)
|
| 17 |
+
|
| 18 |
+
Autor: CNN-BiGRU Project
|
| 19 |
+
"""
|
| 20 |
+
from __future__ import annotations
|
| 21 |
+
|
| 22 |
+
import logging
|
| 23 |
+
import math
|
| 24 |
+
from dataclasses import dataclass, field
|
| 25 |
+
from typing import Optional, Tuple
|
| 26 |
+
|
| 27 |
+
import torch
|
| 28 |
+
import torch.nn as nn
|
| 29 |
+
import torch.nn.functional as F
|
| 30 |
+
|
| 31 |
+
from .rope import RotaryPositionEmbedding
|
| 32 |
+
|
| 33 |
+
logger = logging.getLogger(__name__)
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
# ============================================================================
|
| 37 |
+
# Configuração
|
| 38 |
+
# ============================================================================
|
| 39 |
+
|
| 40 |
+
@dataclass
|
| 41 |
+
class TransformerBlockConfig:
|
| 42 |
+
"""Configuração de um bloco Transformer Decoder."""
|
| 43 |
+
embed_dim: int = 256
|
| 44 |
+
n_heads: int = 4
|
| 45 |
+
ff_dim: Optional[int] = None # default: 4 * embed_dim
|
| 46 |
+
dropout: float = 0.1
|
| 47 |
+
use_rope: bool = True
|
| 48 |
+
max_seq_len: int = 512
|
| 49 |
+
rope_base: float = 10000.0
|
| 50 |
+
# Limiar para attention dropout
|
| 51 |
+
attn_dropout: float = 0.1
|
| 52 |
+
# Bias nas projeções lineares
|
| 53 |
+
qkv_bias: bool = False
|
| 54 |
+
out_bias: bool = False
|
| 55 |
+
ffn_bias: bool = False
|
| 56 |
+
|
| 57 |
+
|
| 58 |
+
# ============================================================================
|
| 59 |
+
# Causal Self-Attention
|
| 60 |
+
# ============================================================================
|
| 61 |
+
|
| 62 |
+
class CausalSelfAttention(nn.Module):
|
| 63 |
+
"""
|
| 64 |
+
Mecanismo de Atenção Multi-Head com Máscara Causal (Triangular Inferior).
|
| 65 |
+
|
| 66 |
+
Conforme dados.txt linhas 553-592.
|
| 67 |
+
|
| 68 |
+
Fórmula matemática:
|
| 69 |
+
Attention(Q, K, V) = softmax( (Q @ K^T) / sqrt(d_k) + mask ) @ V
|
| 70 |
+
|
| 71 |
+
Com RoPE opcional, Q e K são rotacionados antes do produto escalar.
|
| 72 |
+
"""
|
| 73 |
+
|
| 74 |
+
def __init__(self, config: TransformerBlockConfig):
|
| 75 |
+
super().__init__()
|
| 76 |
+
self.config = config
|
| 77 |
+
self.embed_dim = config.embed_dim
|
| 78 |
+
self.n_heads = config.n_heads
|
| 79 |
+
if self.embed_dim % self.n_heads != 0:
|
| 80 |
+
raise ValueError(
|
| 81 |
+
f"embed_dim ({self.embed_dim}) deve ser divisível por "
|
| 82 |
+
f"n_heads ({self.n_heads})"
|
| 83 |
+
)
|
| 84 |
+
self.head_dim = self.embed_dim // self.n_heads
|
| 85 |
+
|
| 86 |
+
# Projeção QKV unificada (3 * embed_dim) — ganho de performance
|
| 87 |
+
self.qkv_projection = nn.Linear(self.embed_dim, 3 * self.embed_dim, bias=config.qkv_bias)
|
| 88 |
+
self.out_projection = nn.Linear(self.embed_dim, self.embed_dim, bias=config.out_bias)
|
| 89 |
+
|
| 90 |
+
self.attn_dropout = nn.Dropout(config.attn_dropout)
|
| 91 |
+
self.resid_dropout = nn.Dropout(config.dropout)
|
| 92 |
+
|
| 93 |
+
# RoPE (opcional)
|
| 94 |
+
self.use_rope = config.use_rope
|
| 95 |
+
if self.use_rope:
|
| 96 |
+
self.rope = RotaryPositionEmbedding(
|
| 97 |
+
head_dim=self.head_dim,
|
| 98 |
+
max_seq_len=config.max_seq_len,
|
| 99 |
+
base=config.rope_base,
|
| 100 |
+
)
|
| 101 |
+
else:
|
| 102 |
+
self.rope = None
|
| 103 |
+
|
| 104 |
+
def forward(
|
| 105 |
+
self,
|
| 106 |
+
x: torch.Tensor,
|
| 107 |
+
padding_mask: Optional[torch.Tensor] = None,
|
| 108 |
+
kv_cache_k: Optional[torch.Tensor] = None,
|
| 109 |
+
kv_cache_v: Optional[torch.Tensor] = None,
|
| 110 |
+
position_offset: int = 0,
|
| 111 |
+
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
| 112 |
+
"""
|
| 113 |
+
Args:
|
| 114 |
+
x: [batch, seq_len, embed_dim]
|
| 115 |
+
padding_mask: [batch, seq_len] — 1 para token real, 0 para PAD
|
| 116 |
+
(se None, não aplica)
|
| 117 |
+
kv_cache_k, kv_cache_v: cache KV de passos anteriores
|
| 118 |
+
[batch, n_heads, cached_len, head_dim] ou None
|
| 119 |
+
position_offset: offset posicional (para uso com cache KV)
|
| 120 |
+
|
| 121 |
+
Returns:
|
| 122 |
+
out: [batch, seq_len, embed_dim]
|
| 123 |
+
new_k: [batch, n_heads, total_len, head_dim] — K atualizado para cache
|
| 124 |
+
new_v: [batch, n_heads, total_len, head_dim] — V atualizado
|
| 125 |
+
"""
|
| 126 |
+
B, T, C = x.size()
|
| 127 |
+
|
| 128 |
+
# Projeção QKV unificada
|
| 129 |
+
qkv = self.qkv_projection(x) # [B, T, 3*C]
|
| 130 |
+
q, k, v = qkv.split(self.embed_dim, dim=2) # cada um: [B, T, C]
|
| 131 |
+
|
| 132 |
+
# Reshape para multi-head: [B, n_heads, T, head_dim]
|
| 133 |
+
q = q.view(B, T, self.n_heads, self.head_dim).transpose(1, 2)
|
| 134 |
+
k = k.view(B, T, self.n_heads, self.head_dim).transpose(1, 2)
|
| 135 |
+
v = v.view(B, T, self.n_heads, self.head_dim).transpose(1, 2)
|
| 136 |
+
|
| 137 |
+
# Aplicar RoPE a Q e K (não a V)
|
| 138 |
+
if self.use_rope and self.rope is not None:
|
| 139 |
+
q = self.rope(q, offset=position_offset)
|
| 140 |
+
k = self.rope(k, offset=position_offset)
|
| 141 |
+
|
| 142 |
+
# Concatenar com cache KV (se fornecido)
|
| 143 |
+
if kv_cache_k is not None and kv_cache_v is not None:
|
| 144 |
+
k = torch.cat([kv_cache_k, k], dim=2) # [B, n_heads, cached+T, head_dim]
|
| 145 |
+
v = torch.cat([kv_cache_v, v], dim=2)
|
| 146 |
+
new_k, new_v = k, v
|
| 147 |
+
|
| 148 |
+
total_len = k.size(2)
|
| 149 |
+
|
| 150 |
+
# Atenção por produto escalar escalonado
|
| 151 |
+
# scores: [B, n_heads, T, total_len]
|
| 152 |
+
attn_scores = torch.matmul(q, k.transpose(-2, -1)) * (1.0 / math.sqrt(self.head_dim))
|
| 153 |
+
|
| 154 |
+
# Máscara causal: o token na posição i só pode olhar para posições <= i
|
| 155 |
+
# Em modo cache: pos i corresponde à posição absoluta (offset + i)
|
| 156 |
+
# O token atual (offset+i) não pode olhar para futuros (offset+i+1, ...)
|
| 157 |
+
# Como K contém [cached..., current_T], a máscara deve permitir que
|
| 158 |
+
# cada q[j] (j=0..T-1) veja k[0..cached+j]
|
| 159 |
+
if T > 0:
|
| 160 |
+
# Índices absolutos de Q: [offset, offset+T-1]
|
| 161 |
+
# Índices absolutos de K: [0, total_len-1]
|
| 162 |
+
q_pos = torch.arange(position_offset, position_offset + T, device=x.device) # [T]
|
| 163 |
+
k_pos = torch.arange(0, total_len, device=x.device) # [total_len]
|
| 164 |
+
# permitido se k_pos <= q_pos
|
| 165 |
+
causal_mask = k_pos.unsqueeze(0) <= q_pos.unsqueeze(1) # [T, total_len]
|
| 166 |
+
# Aplicar: onde causal_mask==False → -inf
|
| 167 |
+
attn_scores = attn_scores.masked_fill(
|
| 168 |
+
~causal_mask.unsqueeze(0).unsqueeze(0), # [1, 1, T, total_len]
|
| 169 |
+
float('-inf')
|
| 170 |
+
)
|
| 171 |
+
|
| 172 |
+
# Máscara de padding (se fornecida)
|
| 173 |
+
if padding_mask is not None:
|
| 174 |
+
# padding_mask: [B, T_query] — mas a máscara deve afetar a chave (K)
|
| 175 |
+
# Se um token K é PAD, sua atenção deve ser zero
|
| 176 |
+
# Como o cache K pode conter tokens do passado que NÃO são PAD,
|
| 177 |
+
# a máscara de padding deve cobrir todo o K (cached + atual)
|
| 178 |
+
# Para simplicidade, assumimos que tokens cached não são PAD
|
| 179 |
+
# e usamos padding_mask apenas para o K atual
|
| 180 |
+
if padding_mask.dim() == 2:
|
| 181 |
+
# [B, T] → [B, 1, 1, T]
|
| 182 |
+
pad_mask_k = (padding_mask > 0).unsqueeze(1).unsqueeze(1)
|
| 183 |
+
else:
|
| 184 |
+
pad_mask_k = (padding_mask > 0).unsqueeze(1).unsqueeze(1)
|
| 185 |
+
# Construir máscara completa sobre K
|
| 186 |
+
if kv_cache_k is not None:
|
| 187 |
+
cached_len = kv_cache_k.size(2)
|
| 188 |
+
# Tokens cached: considerados válidos (não PAD)
|
| 189 |
+
cached_part = torch.ones(
|
| 190 |
+
B, 1, 1, cached_len, device=x.device, dtype=pad_mask_k.dtype
|
| 191 |
+
)
|
| 192 |
+
pad_mask_full = torch.cat([cached_part, pad_mask_k], dim=3)
|
| 193 |
+
else:
|
| 194 |
+
pad_mask_full = pad_mask_k
|
| 195 |
+
attn_scores = attn_scores.masked_fill(~pad_mask_full, float('-inf'))
|
| 196 |
+
|
| 197 |
+
# Softmax + dropout
|
| 198 |
+
# Lidar com linhas totalmente -inf (todas as chaves são PAD)
|
| 199 |
+
# Substituir -inf por 0 para evitar NaN
|
| 200 |
+
all_neg = (attn_scores <= -1e9).all(dim=-1, keepdim=True)
|
| 201 |
+
attn_scores = torch.where(all_neg, torch.zeros_like(attn_scores), attn_scores)
|
| 202 |
+
|
| 203 |
+
attn_weights = F.softmax(attn_scores, dim=-1)
|
| 204 |
+
attn_weights = self.attn_dropout(attn_weights)
|
| 205 |
+
|
| 206 |
+
# Fusão dos contextos
|
| 207 |
+
out = torch.matmul(attn_weights, v) # [B, n_heads, T, head_dim]
|
| 208 |
+
out = out.transpose(1, 2).contiguous().view(B, T, C)
|
| 209 |
+
|
| 210 |
+
out = self.resid_dropout(self.out_projection(out))
|
| 211 |
+
return out, new_k, new_v
|
| 212 |
+
|
| 213 |
+
|
| 214 |
+
# ============================================================================
|
| 215 |
+
# Transformer Block
|
| 216 |
+
# ============================================================================
|
| 217 |
+
|
| 218 |
+
class TransformerBlock(nn.Module):
|
| 219 |
+
"""
|
| 220 |
+
Bloco decodificador padrão com Pre-LN (conforme dados.txt linhas 594-612).
|
| 221 |
+
|
| 222 |
+
Estrutura:
|
| 223 |
+
x = x + attn(ln1(x))
|
| 224 |
+
x = x + ffn(ln2(x))
|
| 225 |
+
FFN: Linear(d, 4d) → GELU → Linear(4d, d) → Dropout
|
| 226 |
+
"""
|
| 227 |
+
|
| 228 |
+
def __init__(self, config: TransformerBlockConfig):
|
| 229 |
+
super().__init__()
|
| 230 |
+
self.config = config
|
| 231 |
+
self.ln1 = nn.LayerNorm(config.embed_dim)
|
| 232 |
+
self.attn = CausalSelfAttention(config)
|
| 233 |
+
self.ln2 = nn.LayerNorm(config.embed_dim)
|
| 234 |
+
ff_dim = config.ff_dim or (4 * config.embed_dim)
|
| 235 |
+
self.ffn = nn.Sequential(
|
| 236 |
+
nn.Linear(config.embed_dim, ff_dim, bias=config.ffn_bias),
|
| 237 |
+
nn.GELU(),
|
| 238 |
+
nn.Linear(ff_dim, config.embed_dim, bias=config.ffn_bias),
|
| 239 |
+
nn.Dropout(config.dropout),
|
| 240 |
+
)
|
| 241 |
+
|
| 242 |
+
def forward(
|
| 243 |
+
self,
|
| 244 |
+
x: torch.Tensor,
|
| 245 |
+
padding_mask: Optional[torch.Tensor] = None,
|
| 246 |
+
kv_cache_k: Optional[torch.Tensor] = None,
|
| 247 |
+
kv_cache_v: Optional[torch.Tensor] = None,
|
| 248 |
+
position_offset: int = 0,
|
| 249 |
+
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
| 250 |
+
"""
|
| 251 |
+
Args:
|
| 252 |
+
x: [batch, seq_len, embed_dim]
|
| 253 |
+
padding_mask: [batch, seq_len] opcional
|
| 254 |
+
kv_cache_k, kv_cache_v: cache KV opcional
|
| 255 |
+
position_offset: offset posicional para RoPE
|
| 256 |
+
|
| 257 |
+
Returns:
|
| 258 |
+
out: [batch, seq_len, embed_dim]
|
| 259 |
+
new_k: cache K atualizado
|
| 260 |
+
new_v: cache V atualizado
|
| 261 |
+
"""
|
| 262 |
+
# Pre-LN: normaliza antes de passar para a sub-camada
|
| 263 |
+
normed = self.ln1(x)
|
| 264 |
+
attn_out, new_k, new_v = self.attn(
|
| 265 |
+
normed, padding_mask=padding_mask,
|
| 266 |
+
kv_cache_k=kv_cache_k, kv_cache_v=kv_cache_v,
|
| 267 |
+
position_offset=position_offset,
|
| 268 |
+
)
|
| 269 |
+
x = x + attn_out
|
| 270 |
+
|
| 271 |
+
# FFN
|
| 272 |
+
x = x + self.ffn(self.ln2(x))
|
| 273 |
+
return x, new_k, new_v
|
| 274 |
+
|
| 275 |
+
|
| 276 |
+
# ============================================================================
|
| 277 |
+
# Stack de Transformer Blocks (Decoder completo)
|
| 278 |
+
# ============================================================================
|
| 279 |
+
|
| 280 |
+
class TransformerDecoderStack(nn.Module):
|
| 281 |
+
"""
|
| 282 |
+
Empilhamento de N TransformerBlocks (decoder-only, estilo GPT).
|
| 283 |
+
|
| 284 |
+
Conforme dados.txt linhas 614-648.
|
| 285 |
+
|
| 286 |
+
Suporta weight tying (compartilhamento de pesos entre token_embedding e lm_head).
|
| 287 |
+
"""
|
| 288 |
+
|
| 289 |
+
def __init__(
|
| 290 |
+
self,
|
| 291 |
+
vocab_size: int,
|
| 292 |
+
embed_dim: int = 256,
|
| 293 |
+
n_heads: int = 4,
|
| 294 |
+
n_layers: int = 2,
|
| 295 |
+
ff_dim: Optional[int] = None,
|
| 296 |
+
dropout: float = 0.1,
|
| 297 |
+
max_seq_len: int = 512,
|
| 298 |
+
use_rope: bool = True,
|
| 299 |
+
pad_id: int = 0,
|
| 300 |
+
weight_tying: bool = True,
|
| 301 |
+
):
|
| 302 |
+
super().__init__()
|
| 303 |
+
self.vocab_size = vocab_size
|
| 304 |
+
self.embed_dim = embed_dim
|
| 305 |
+
self.max_seq_len = max_seq_len
|
| 306 |
+
self.pad_id = pad_id
|
| 307 |
+
self.weight_tying = weight_tying
|
| 308 |
+
|
| 309 |
+
block_config = TransformerBlockConfig(
|
| 310 |
+
embed_dim=embed_dim,
|
| 311 |
+
n_heads=n_heads,
|
| 312 |
+
ff_dim=ff_dim,
|
| 313 |
+
dropout=dropout,
|
| 314 |
+
use_rope=use_rope,
|
| 315 |
+
max_seq_len=max_seq_len,
|
| 316 |
+
)
|
| 317 |
+
|
| 318 |
+
self.token_embedding = nn.Embedding(vocab_size, embed_dim, padding_idx=pad_id)
|
| 319 |
+
# Sem positional embedding separado quando RoPE está ativo
|
| 320 |
+
if not use_rope:
|
| 321 |
+
self.position_embedding = nn.Embedding(max_seq_len, embed_dim)
|
| 322 |
+
else:
|
| 323 |
+
self.position_embedding = None
|
| 324 |
+
|
| 325 |
+
self.blocks = nn.ModuleList([
|
| 326 |
+
TransformerBlock(block_config) for _ in range(n_layers)
|
| 327 |
+
])
|
| 328 |
+
self.ln_f = nn.LayerNorm(embed_dim)
|
| 329 |
+
self.lm_head = nn.Linear(embed_dim, vocab_size, bias=False)
|
| 330 |
+
|
| 331 |
+
# Weight tying
|
| 332 |
+
if weight_tying:
|
| 333 |
+
self.lm_head.weight = self.token_embedding.weight
|
| 334 |
+
|
| 335 |
+
def forward(
|
| 336 |
+
self,
|
| 337 |
+
idx: torch.Tensor,
|
| 338 |
+
padding_mask: Optional[torch.Tensor] = None,
|
| 339 |
+
kv_caches: Optional[list] = None,
|
| 340 |
+
position_offset: int = 0,
|
| 341 |
+
return_hidden: bool = False,
|
| 342 |
+
) -> torch.Tensor:
|
| 343 |
+
"""
|
| 344 |
+
Args:
|
| 345 |
+
idx: [batch, seq_len] — token IDs
|
| 346 |
+
padding_mask: [batch, seq_len] — 1 para real, 0 para PAD
|
| 347 |
+
kv_caches: lista de (k, v) por camada, ou None
|
| 348 |
+
position_offset: offset posicional (para uso com cache)
|
| 349 |
+
return_hidden: se True, retorna também o hidden state antes do lm_head
|
| 350 |
+
|
| 351 |
+
Returns:
|
| 352 |
+
logits: [batch, seq_len, vocab_size]
|
| 353 |
+
hidden (opcional): [batch, seq_len, embed_dim]
|
| 354 |
+
"""
|
| 355 |
+
B, T = idx.size()
|
| 356 |
+
if T > self.max_seq_len:
|
| 357 |
+
raise ValueError(
|
| 358 |
+
f"Sequência de tamanho {T} excede o limite máximo configurado de "
|
| 359 |
+
f"{self.max_seq_len}."
|
| 360 |
+
)
|
| 361 |
+
|
| 362 |
+
# Embedding
|
| 363 |
+
x = self.token_embedding(idx)
|
| 364 |
+
if not self.weight_tying and self.position_embedding is not None:
|
| 365 |
+
positions = torch.arange(
|
| 366 |
+
position_offset, position_offset + T,
|
| 367 |
+
dtype=torch.long, device=idx.device
|
| 368 |
+
).unsqueeze(0)
|
| 369 |
+
x = x + self.position_embedding(positions)
|
| 370 |
+
|
| 371 |
+
# Passar pelos blocos
|
| 372 |
+
new_caches = []
|
| 373 |
+
for i, block in enumerate(self.blocks):
|
| 374 |
+
cache_k = None
|
| 375 |
+
cache_v = None
|
| 376 |
+
if kv_caches is not None and i < len(kv_caches):
|
| 377 |
+
cache_k, cache_v = kv_caches[i]
|
| 378 |
+
x, new_k, new_v = block(
|
| 379 |
+
x, padding_mask=padding_mask,
|
| 380 |
+
kv_cache_k=cache_k, kv_cache_v=cache_v,
|
| 381 |
+
position_offset=position_offset,
|
| 382 |
+
)
|
| 383 |
+
new_caches.append((new_k, new_v))
|
| 384 |
+
|
| 385 |
+
# LayerNorm final
|
| 386 |
+
x = self.ln_f(x)
|
| 387 |
+
|
| 388 |
+
# LM head
|
| 389 |
+
logits = self.lm_head(x)
|
| 390 |
+
|
| 391 |
+
if return_hidden:
|
| 392 |
+
return logits, x, new_caches
|
| 393 |
+
return logits, new_caches
|
| 394 |
+
|
| 395 |
+
|
| 396 |
+
__all__ = [
|
| 397 |
+
"TransformerBlockConfig",
|
| 398 |
+
"CausalSelfAttention",
|
| 399 |
+
"TransformerBlock",
|
| 400 |
+
"TransformerDecoderStack",
|
| 401 |
+
]
|
cnn_bigru/requirements.txt
ADDED
|
@@ -0,0 +1,29 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# CNN-BiGRU — Requisitos
|
| 2 |
+
# Instale com: pip install -r requirements.txt
|
| 3 |
+
|
| 4 |
+
# Core deep learning
|
| 5 |
+
torch>=2.0
|
| 6 |
+
tokenizers>=0.15
|
| 7 |
+
numpy>=1.24
|
| 8 |
+
|
| 9 |
+
# HuggingFace ecosystem
|
| 10 |
+
huggingface_hub>=0.20
|
| 11 |
+
datasets>=2.16
|
| 12 |
+
|
| 13 |
+
# System/runtime
|
| 14 |
+
psutil>=5.9
|
| 15 |
+
tqdm>=4.66
|
| 16 |
+
|
| 17 |
+
# Opcional — Intel Extension for PyTorch (apenas para CPUs Intel Xeon com AMX/AVX512)
|
| 18 |
+
# Descomente se estiver rodando em hardware Intel compatível:
|
| 19 |
+
# intel_extension_for_pytorch>=2.0
|
| 20 |
+
|
| 21 |
+
# Opcional — para testes
|
| 22 |
+
# pytest>=7.0
|
| 23 |
+
# pytest-cov>=4.0
|
| 24 |
+
|
| 25 |
+
# Opcional — para acelerar treinamento distribuído
|
| 26 |
+
# accelerate>=0.20
|
| 27 |
+
|
| 28 |
+
# Opcional — para salvar modelos em formato safetensors
|
| 29 |
+
# safetensors>=0.4
|
cnn_bigru/scripts/push_to_hf.py
ADDED
|
@@ -0,0 +1,294 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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)
|
| 5 |
+
2. Envia todos os scripts .py, README.md, requirements.txt, docs/
|
| 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:
|
| 20 |
+
HF_TOKEN=hf_xxx python push_to_hf.py
|
| 21 |
+
"""
|
| 22 |
+
from __future__ import annotations
|
| 23 |
+
|
| 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
|
| 31 |
+
|
| 32 |
+
logging.basicConfig(level=logging.INFO, format="%(asctime)s | %(levelname)s | %(message)s")
|
| 33 |
+
logger = logging.getLogger("push_hf")
|
| 34 |
+
|
| 35 |
+
REPO_ID = "PowerMachine/CNN-BiGRU"
|
| 36 |
+
# Path: scripts/push_to_hf.py -> parent=scripts -> parent=cnn_bigru
|
| 37 |
+
PROJECT_DIR = Path(__file__).resolve().parent.parent
|
| 38 |
+
# PROJECT_DIR.parent = raiz do projeto
|
| 39 |
+
REPO_ROOT = PROJECT_DIR.parent
|
| 40 |
+
|
| 41 |
+
# Extensões e nomes de arquivo permitidos para upload
|
| 42 |
+
ALLOWED_EXTENSIONS: Set[str] = {
|
| 43 |
+
".py", ".md", ".txt", ".json", ".yaml", ".yml", ".toml", ".cfg", ".ini", ".sh",
|
| 44 |
+
}
|
| 45 |
+
|
| 46 |
+
# Padrões a ignorar (em qualquer parte do caminho)
|
| 47 |
+
IGNORE_PATTERNS: Set[str] = {
|
| 48 |
+
"__pycache__", ".pyc", ".pyo", ".pyd",
|
| 49 |
+
".git", ".gitignore", ".gitattributes",
|
| 50 |
+
"node_modules", ".pytest_cache", ".mypy_cache", ".ruff_cache",
|
| 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)
|
| 61 |
+
SENSITIVE_FILENAMES: Set[str] = {
|
| 62 |
+
".env", ".env.local", ".env.production", ".env.development",
|
| 63 |
+
"secrets.json", "credentials.json", "config.local.json",
|
| 64 |
+
"hf_token.txt", "token.txt",
|
| 65 |
+
}
|
| 66 |
+
|
| 67 |
+
|
| 68 |
+
def should_ignore(path: Path) -> bool:
|
| 69 |
+
"""Verifica se um arquivo deve ser ignorado no upload."""
|
| 70 |
+
parts = path.parts
|
| 71 |
+
name = path.name.lower()
|
| 72 |
+
|
| 73 |
+
# Verificar padrões por partes do caminho
|
| 74 |
+
for part in parts:
|
| 75 |
+
part_lower = part.lower()
|
| 76 |
+
if part_lower in IGNORE_PATTERNS:
|
| 77 |
+
return True
|
| 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:
|
| 87 |
+
return True
|
| 88 |
+
|
| 89 |
+
# Verificar extensões permitidas
|
| 90 |
+
if path.suffix.lower() not in ALLOWED_EXTENSIONS:
|
| 91 |
+
# Permitir arquivos sem extensão apenas se forem específicos (README, LICENSE, etc.)
|
| 92 |
+
if name not in {"readme", "license", "license-mit", "authors", "contributors"}:
|
| 93 |
+
return True
|
| 94 |
+
|
| 95 |
+
# Verificar padrões com wildcards
|
| 96 |
+
for pattern in IGNORE_PATTERNS:
|
| 97 |
+
if pattern.startswith("*"):
|
| 98 |
+
suffix = pattern[1:]
|
| 99 |
+
if name.endswith(suffix):
|
| 100 |
+
return True
|
| 101 |
+
|
| 102 |
+
return False
|
| 103 |
+
|
| 104 |
+
|
| 105 |
+
def collect_files(project_dir: Path) -> List[Path]:
|
| 106 |
+
"""Coleta todos os arquivos elegíveis para upload."""
|
| 107 |
+
files = []
|
| 108 |
+
for path in project_dir.rglob("*"):
|
| 109 |
+
if path.is_file() and not should_ignore(path):
|
| 110 |
+
files.append(path)
|
| 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
|
| 152 |
+
for attempt in range(1, max_retries + 1):
|
| 153 |
+
try:
|
| 154 |
+
logger.info(f"Tentativa {attempt}/{max_retries}: upload_folder...")
|
| 155 |
+
commit_info = api.upload_folder(
|
| 156 |
+
folder_path=str(folder),
|
| 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}")
|
| 164 |
+
return True, None
|
| 165 |
+
except Exception as e:
|
| 166 |
+
last_error = e
|
| 167 |
+
logger.warning(f"Tentativa {attempt} falhou: {e}")
|
| 168 |
+
if attempt < max_retries:
|
| 169 |
+
wait_time = 2 ** attempt
|
| 170 |
+
logger.info(f"Aguardando {wait_time}s antes de tentar novamente...")
|
| 171 |
+
time.sleep(wait_time)
|
| 172 |
+
return False, last_error
|
| 173 |
+
|
| 174 |
+
|
| 175 |
+
def main():
|
| 176 |
+
# Pega HF_TOKEN da variável de ambiente
|
| 177 |
+
hf_token = os.environ.get("HF_TOKEN") or os.environ.get("HUGGING_FACE_HUB_TOKEN")
|
| 178 |
+
if not hf_token:
|
| 179 |
+
logger.error("HF_TOKEN não definido no ambiente. Abortando.")
|
| 180 |
+
return 1
|
| 181 |
+
|
| 182 |
+
try:
|
| 183 |
+
from huggingface_hub import HfApi, create_repo
|
| 184 |
+
except ImportError:
|
| 185 |
+
logger.error("huggingface_hub não instalado. Execute: pip install huggingface_hub")
|
| 186 |
+
return 1
|
| 187 |
+
|
| 188 |
+
api = HfApi(token=hf_token)
|
| 189 |
+
|
| 190 |
+
# 1. Cria repositório (se não existir)
|
| 191 |
+
try:
|
| 192 |
+
repo_url = create_repo(
|
| 193 |
+
repo_id=REPO_ID,
|
| 194 |
+
token=hf_token,
|
| 195 |
+
repo_type="model",
|
| 196 |
+
exist_ok=True,
|
| 197 |
+
private=False,
|
| 198 |
+
)
|
| 199 |
+
logger.info(f"Repositório criado/acessado: {repo_url}")
|
| 200 |
+
except Exception as e:
|
| 201 |
+
logger.error(f"Falha ao criar repositório: {e}")
|
| 202 |
+
# Limpa o token mesmo em caso de falha
|
| 203 |
+
_cleanup_token()
|
| 204 |
+
return 1
|
| 205 |
+
|
| 206 |
+
# 2. Coleta todos os arquivos para upload
|
| 207 |
+
files_to_upload = collect_files(PROJECT_DIR)
|
| 208 |
+
logger.info(f"Total de arquivos para upload: {len(files_to_upload)} (de {PROJECT_DIR})")
|
| 209 |
+
|
| 210 |
+
if not files_to_upload:
|
| 211 |
+
logger.warning("Nenhum arquivo elegível para upload.")
|
| 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(
|
| 224 |
+
api, folder=PROJECT_DIR, repo_id=REPO_ID, token=hf_token, max_retries=3
|
| 225 |
+
)
|
| 226 |
+
|
| 227 |
+
if success:
|
| 228 |
+
logger.info("=" * 60)
|
| 229 |
+
logger.info(f"Upload concluído com sucesso!")
|
| 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 |
+
|
| 251 |
+
def _cleanup_token():
|
| 252 |
+
"""Remove todas as variáveis de token HF do ambiente."""
|
| 253 |
+
removed = []
|
| 254 |
+
for var in ("HF_TOKEN", "HUGGING_FACE_HUB_TOKEN", "HF_HUB_TOKEN"):
|
| 255 |
+
if var in os.environ:
|
| 256 |
+
del os.environ[var]
|
| 257 |
+
removed.append(var)
|
| 258 |
+
if removed:
|
| 259 |
+
logger.info(f"Tokens removidos do ambiente: {', '.join(removed)}")
|
| 260 |
+
else:
|
| 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())
|
cnn_bigru/tests/__init__.py
ADDED
|
File without changes
|
cnn_bigru/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)
|
cnn_bigru/tests/test_50_samples.py
ADDED
|
@@ -0,0 +1,832 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""test_50_samples.py — Teste de 50 amostras para erros lógicos ou falhas.
|
| 2 |
+
|
| 3 |
+
Executa o pipeline completo:
|
| 4 |
+
1. Inicializa xeon_runtime + memory_optimizer
|
| 5 |
+
2. Treina BBPE tokenizer em corpus sintético
|
| 6 |
+
3. Cria dataset streaming multimodal (50 amostras)
|
| 7 |
+
4. Instancia modelo multimodal + generator + verifier + anti-hallucination
|
| 8 |
+
5. Executa treinamento cooperativo com:
|
| 9 |
+
- Synergy search (N tentativas)
|
| 10 |
+
- Hypothesis controller (ativa em punições)
|
| 11 |
+
- Auto-learner (ajuste dinâmico de LR + spectral norm)
|
| 12 |
+
6. Executa inferência com sampling
|
| 13 |
+
7. Avalia perplexidade
|
| 14 |
+
8. Reporta erros lógicos/falhas encontradas
|
| 15 |
+
|
| 16 |
+
Usage:
|
| 17 |
+
python -m cnn_bigru.tests.test_50_samples
|
| 18 |
+
"""
|
| 19 |
+
from __future__ import annotations
|
| 20 |
+
|
| 21 |
+
import logging
|
| 22 |
+
import os
|
| 23 |
+
import sys
|
| 24 |
+
import time
|
| 25 |
+
import traceback
|
| 26 |
+
from pathlib import Path
|
| 27 |
+
from typing import Dict, List
|
| 28 |
+
|
| 29 |
+
# Setup paths (must come before torch import for xeon_runtime)
|
| 30 |
+
PROJECT_ROOT = Path(__file__).resolve().parent.parent.parent
|
| 31 |
+
sys.path.insert(0, str(PROJECT_ROOT))
|
| 32 |
+
|
| 33 |
+
# Ativa Xeon runtime ANTES de importar torch
|
| 34 |
+
from cnn_bigru.utils.xeon_runtime import optimize_xeon_environment
|
| 35 |
+
N_CORES = optimize_xeon_environment()
|
| 36 |
+
|
| 37 |
+
import numpy as np
|
| 38 |
+
import torch
|
| 39 |
+
from torch.utils.data import DataLoader
|
| 40 |
+
|
| 41 |
+
from cnn_bigru.tokenizer.bbpe_tokenizer import BBPETokenizer
|
| 42 |
+
from cnn_bigru.data.streaming_dataset import (
|
| 43 |
+
MultimodalStreamingDataset,
|
| 44 |
+
collate_multimodal,
|
| 45 |
+
)
|
| 46 |
+
from cnn_bigru.models.multimodal_model import MultimodalCNNBiGRU
|
| 47 |
+
from cnn_bigru.models.generator_verifier import (
|
| 48 |
+
GeneratorCNNBiGRU,
|
| 49 |
+
VerifierCNNBiGRU,
|
| 50 |
+
AntiHallucinationLayer,
|
| 51 |
+
)
|
| 52 |
+
from cnn_bigru.models.rope import RotaryPositionEmbedding
|
| 53 |
+
from cnn_bigru.models.transformer_block import (
|
| 54 |
+
TransformerBlockConfig,
|
| 55 |
+
CausalSelfAttention,
|
| 56 |
+
TransformerBlock,
|
| 57 |
+
TransformerDecoderStack,
|
| 58 |
+
)
|
| 59 |
+
from cnn_bigru.models.context_window import (
|
| 60 |
+
ContextWindowConfig,
|
| 61 |
+
ContextWindowManager,
|
| 62 |
+
KVCache,
|
| 63 |
+
)
|
| 64 |
+
from cnn_bigru.utils.ewc import EWCConfig, EWCState
|
| 65 |
+
from cnn_bigru.losses.losses import LossConfig, MultiLoss
|
| 66 |
+
from cnn_bigru.training.trainer import TrainerConfig, CooperativeTrainer
|
| 67 |
+
from cnn_bigru.training.auto_learner import (
|
| 68 |
+
AutoLearnConfig,
|
| 69 |
+
orthogonal_init_model,
|
| 70 |
+
)
|
| 71 |
+
from cnn_bigru.training.hypothesis_controller import (
|
| 72 |
+
HypothesisConfig,
|
| 73 |
+
HypothesisController,
|
| 74 |
+
)
|
| 75 |
+
from cnn_bigru.inference.inference import (
|
| 76 |
+
generate_with_sampling,
|
| 77 |
+
evaluate_perplexity,
|
| 78 |
+
make_default_context_window,
|
| 79 |
+
)
|
| 80 |
+
from cnn_bigru.utils.memory_optimizer import MemoryOptimizer
|
| 81 |
+
|
| 82 |
+
logging.basicConfig(
|
| 83 |
+
level=logging.INFO,
|
| 84 |
+
format="%(asctime)s | %(levelname)s | %(name)s | %(message)s",
|
| 85 |
+
datefmt="%H:%M:%S",
|
| 86 |
+
)
|
| 87 |
+
logger = logging.getLogger("test_50")
|
| 88 |
+
|
| 89 |
+
|
| 90 |
+
# ============================================================================
|
| 91 |
+
# Corpus para treino do tokenizer
|
| 92 |
+
# ============================================================================
|
| 93 |
+
|
| 94 |
+
CORPUS_TEXTS = [
|
| 95 |
+
"o modelo cooperativo cnn-bigru combina fluxos paralelos",
|
| 96 |
+
"a atencao cruzada troca informacoes entre redes a e b",
|
| 97 |
+
"gru bidirecional captura dependencias temporais em ambas direcoes",
|
| 98 |
+
"a porta de atenuacao controla o fluxo de informacao entre celulas",
|
| 99 |
+
"a camada anti-alucinacao usa logica fuzzy de lukasiewicz",
|
| 100 |
+
"o verificador classifica passos com sigmoid binaria",
|
| 101 |
+
"penalidades linear e exponencial reduzem o erro grave",
|
| 102 |
+
"ajuste dinamico de lr estabiliza o treinamento cooperativo",
|
| 103 |
+
"normalizacao espectral limita a norma dos pesos",
|
| 104 |
+
"inicializacao ortogonal estabiliza matrizes recorrentes",
|
| 105 |
+
"o gerador decodificador usa atencao bahdanau sobre o encoder",
|
| 106 |
+
"a fusao multimodal combina texto imagem e audio",
|
| 107 |
+
"hipoteses sao ativadas quando o verificador pune o passo",
|
| 108 |
+
"synergy search tenta n configuracoes e escolhe a melhor",
|
| 109 |
+
"a perplexidade mede a confusao do modelo na previsao",
|
| 110 |
+
"top-k e top-p filtram a distribuicao de probabilidade",
|
| 111 |
+
"temperature ajusta a entropia das previsoes do modelo",
|
| 112 |
+
"presence penalty pune tokens ja aparecidos na geracao",
|
| 113 |
+
"frequency penalty pune proporcionalmente a frequencia do token",
|
| 114 |
+
"streaming dataset carrega amostras sem materializacao completa",
|
| 115 |
+
"byte bpe tokeniza qualquer string utf-8 sem unk",
|
| 116 |
+
"o otimizador adamw combina momentum e weight decay",
|
| 117 |
+
"gradient clipping previne explosao de gradiente",
|
| 118 |
+
"amp reduz vram com precisao mista bfloat16",
|
| 119 |
+
"memory optimizer limpa cache entre batches",
|
| 120 |
+
] * 4 # ~100 amostras para treino BBPE
|
| 121 |
+
|
| 122 |
+
|
| 123 |
+
def step_01_init_runtime() -> Dict:
|
| 124 |
+
"""Inicializa runtime e reporta configurações."""
|
| 125 |
+
logger.info("=" * 70)
|
| 126 |
+
logger.info("PASSO 1: Inicialização do runtime Xeon + memory optimizer")
|
| 127 |
+
logger.info("=" * 70)
|
| 128 |
+
from cnn_bigru.utils.xeon_runtime import get_runtime_info
|
| 129 |
+
info = get_runtime_info()
|
| 130 |
+
for k, v in info.items():
|
| 131 |
+
logger.info(" %s: %s", k, v)
|
| 132 |
+
mem = MemoryOptimizer(enable_amp=False)
|
| 133 |
+
mem.configure()
|
| 134 |
+
mem_info = mem.get_memory_mb()
|
| 135 |
+
logger.info(" Memória: %s", mem_info)
|
| 136 |
+
return {"runtime": info, "memory": mem_info}
|
| 137 |
+
|
| 138 |
+
|
| 139 |
+
def step_02_train_tokenizer() -> BBPETokenizer:
|
| 140 |
+
"""Treina BBPE tokenizer em corpus sintético."""
|
| 141 |
+
logger.info("=" * 70)
|
| 142 |
+
logger.info("PASSO 2: Treino do BBPE tokenizer")
|
| 143 |
+
logger.info("=" * 70)
|
| 144 |
+
t0 = time.time()
|
| 145 |
+
tok = BBPETokenizer.train_from_texts(
|
| 146 |
+
CORPUS_TEXTS,
|
| 147 |
+
vocab_size=2000,
|
| 148 |
+
min_frequency=1,
|
| 149 |
+
)
|
| 150 |
+
elapsed = time.time() - t0
|
| 151 |
+
logger.info(" Vocab size: %d", tok.vocab_size)
|
| 152 |
+
logger.info(" BOS/PAD/EOS IDs: %d/%d/%d", tok.bos_id, tok.pad_id, tok.eos_id)
|
| 153 |
+
logger.info(" Tempo: %.2fs", elapsed)
|
| 154 |
+
|
| 155 |
+
# Validação roundtrip
|
| 156 |
+
test_strs = CORPUS_TEXTS[:5]
|
| 157 |
+
roundtrip = tok.validate_roundtrip(test_strs)
|
| 158 |
+
logger.info(" Roundtrip accuracy: %.2f%%", roundtrip * 100)
|
| 159 |
+
|
| 160 |
+
if roundtrip < 0.8:
|
| 161 |
+
logger.warning(" Roundtrip baixo — possível problema no tokenizer")
|
| 162 |
+
return tok, roundtrip
|
| 163 |
+
|
| 164 |
+
|
| 165 |
+
def step_03_create_dataset(tokenizer: BBPETokenizer, n_samples: int = 50) -> DataLoader:
|
| 166 |
+
"""Cria dataset streaming multimodal com N amostras."""
|
| 167 |
+
logger.info("=" * 70)
|
| 168 |
+
logger.info("PASSO 3: Criação do dataset streaming multimodal (%d amostras)", n_samples)
|
| 169 |
+
logger.info("=" * 70)
|
| 170 |
+
|
| 171 |
+
# Para teste determinístico, usamos fallback sintético
|
| 172 |
+
dataset = MultimodalStreamingDataset(
|
| 173 |
+
n_samples=n_samples,
|
| 174 |
+
hf_datasets=[], # skip HF forçando sintético
|
| 175 |
+
use_synthetic_fallback=True,
|
| 176 |
+
seed=42,
|
| 177 |
+
image_size=(28, 28, 1),
|
| 178 |
+
audio_shape=(32, 40),
|
| 179 |
+
)
|
| 180 |
+
|
| 181 |
+
# Coleta todas as amostras em uma lista (para DataLoader iterável)
|
| 182 |
+
samples = list(dataset)
|
| 183 |
+
logger.info(" Amostras coletadas: %d", len(samples))
|
| 184 |
+
assert len(samples) == n_samples, f"Esperado {n_samples}, obtido {len(samples)}"
|
| 185 |
+
|
| 186 |
+
# Cria DataLoader com collate_fn
|
| 187 |
+
from torch.utils.data import DataLoader
|
| 188 |
+
loader = DataLoader(
|
| 189 |
+
samples,
|
| 190 |
+
batch_size=8,
|
| 191 |
+
shuffle=False,
|
| 192 |
+
collate_fn=lambda b: collate_multimodal(b, tokenizer, max_len=32),
|
| 193 |
+
)
|
| 194 |
+
logger.info(" DataLoader criado: batch_size=8, max_len=32")
|
| 195 |
+
return loader
|
| 196 |
+
|
| 197 |
+
|
| 198 |
+
def step_04_init_model(tokenizer: BBPETokenizer, device: str = "cpu") -> Dict:
|
| 199 |
+
"""Instancia todos os componentes do modelo."""
|
| 200 |
+
logger.info("=" * 70)
|
| 201 |
+
logger.info("PASSO 4: Instanciação dos modelos")
|
| 202 |
+
logger.info("=" * 70)
|
| 203 |
+
V = tokenizer.vocab_size
|
| 204 |
+
|
| 205 |
+
# Modelo multimodal principal
|
| 206 |
+
model = MultimodalCNNBiGRU(
|
| 207 |
+
vocab_size=V,
|
| 208 |
+
num_classes=3,
|
| 209 |
+
embedding_dim=32,
|
| 210 |
+
cnn_filters=32,
|
| 211 |
+
gru_hidden=32,
|
| 212 |
+
n_heads=4,
|
| 213 |
+
dropout=0.1,
|
| 214 |
+
pad_idx=tokenizer.pad_id,
|
| 215 |
+
img_channels=1,
|
| 216 |
+
img_hidden=16,
|
| 217 |
+
img_out_dim=32,
|
| 218 |
+
audio_freq=40,
|
| 219 |
+
audio_hidden=16,
|
| 220 |
+
audio_out_dim=32,
|
| 221 |
+
fusion_dim=64,
|
| 222 |
+
use_spectral_norm=False,
|
| 223 |
+
)
|
| 224 |
+
|
| 225 |
+
# Aplica inicialização ortogonal
|
| 226 |
+
n_init = orthogonal_init_model(model)
|
| 227 |
+
logger.info(" Modelo multimodal: %d params, %d camadas ortogonalizadas",
|
| 228 |
+
sum(p.numel() for p in model.parameters()), n_init)
|
| 229 |
+
|
| 230 |
+
# Generator (encoder + decoder)
|
| 231 |
+
generator = GeneratorCNNBiGRU(
|
| 232 |
+
vocab_size=V,
|
| 233 |
+
embedding_dim=32,
|
| 234 |
+
cnn_filters=32,
|
| 235 |
+
gru_hidden=32,
|
| 236 |
+
n_heads=4,
|
| 237 |
+
dropout=0.1,
|
| 238 |
+
pad_idx=tokenizer.pad_id,
|
| 239 |
+
max_proof_len=16,
|
| 240 |
+
)
|
| 241 |
+
orthogonal_init_model(generator)
|
| 242 |
+
|
| 243 |
+
# Verifier
|
| 244 |
+
verifier = VerifierCNNBiGRU(
|
| 245 |
+
vocab_size=V,
|
| 246 |
+
embedding_dim=32,
|
| 247 |
+
cnn_filters=32,
|
| 248 |
+
gru_hidden=32,
|
| 249 |
+
n_heads=4,
|
| 250 |
+
dropout=0.1,
|
| 251 |
+
pad_idx=tokenizer.pad_id,
|
| 252 |
+
)
|
| 253 |
+
orthogonal_init_model(verifier)
|
| 254 |
+
|
| 255 |
+
# Anti-hallucination
|
| 256 |
+
anti_hall = AntiHallucinationLayer(vocab_size=V, embed_dim=16)
|
| 257 |
+
|
| 258 |
+
logger.info(" Generator params: %d", sum(p.numel() for p in generator.parameters()))
|
| 259 |
+
logger.info(" Verifier params: %d", sum(p.numel() for p in verifier.parameters()))
|
| 260 |
+
logger.info(" Anti-hall params: %d", sum(p.numel() for p in anti_hall.parameters()))
|
| 261 |
+
|
| 262 |
+
return {
|
| 263 |
+
"model": model,
|
| 264 |
+
"generator": generator,
|
| 265 |
+
"verifier": verifier,
|
| 266 |
+
"anti_hall": anti_hall,
|
| 267 |
+
}
|
| 268 |
+
|
| 269 |
+
|
| 270 |
+
def step_05_train(
|
| 271 |
+
models: Dict,
|
| 272 |
+
tokenizer: BBPETokenizer,
|
| 273 |
+
dataloader: DataLoader,
|
| 274 |
+
device: str = "cpu",
|
| 275 |
+
ewc_state: EWCState = None,
|
| 276 |
+
) -> Dict:
|
| 277 |
+
"""Executa treinamento cooperativo.
|
| 278 |
+
|
| 279 |
+
NOVO v2.0: aceita ewc_state opcional para testar EWC.
|
| 280 |
+
"""
|
| 281 |
+
logger.info("=" * 70)
|
| 282 |
+
logger.info("PASSO 5: Treinamento cooperativo (synergy + hipóteses + auto-learn)")
|
| 283 |
+
logger.info("=" * 70)
|
| 284 |
+
|
| 285 |
+
cfg = TrainerConfig(
|
| 286 |
+
num_epochs=2,
|
| 287 |
+
max_batches_per_epoch=6,
|
| 288 |
+
batch_size=8,
|
| 289 |
+
n_synergy_attempts=3,
|
| 290 |
+
use_synergy_search=True,
|
| 291 |
+
use_hypotheses=True,
|
| 292 |
+
n_hypotheses=4,
|
| 293 |
+
loss_config=LossConfig(
|
| 294 |
+
alpha=1.0, beta=0.5, gamma_loss=0.3, delta=0.01,
|
| 295 |
+
lambda_penal=0.1, mu_exp_penal=0.05,
|
| 296 |
+
gamma_exp=1.0, threshold_err=0.5,
|
| 297 |
+
l2_reg=1e-5, use_curvature=True, curvature_eps=1e-3,
|
| 298 |
+
),
|
| 299 |
+
auto_config=AutoLearnConfig(
|
| 300 |
+
kappa_curv=0.01, grad_clip=1.0, spectral_radius=1.0,
|
| 301 |
+
lr_min=1e-6, lr_max=1e-2,
|
| 302 |
+
initial_lr_G=1e-3, initial_lr_V=1e-3,
|
| 303 |
+
use_spectral_norm=True, apply_after_step=True,
|
| 304 |
+
l2_reg=1e-5,
|
| 305 |
+
),
|
| 306 |
+
ewc_config=ewc_state.config if ewc_state else None,
|
| 307 |
+
device=device,
|
| 308 |
+
log_every=1,
|
| 309 |
+
use_verifier_real=True, # NOVO v2.0: usar verificador real
|
| 310 |
+
use_generator=True,
|
| 311 |
+
use_hypothesis_output=True,
|
| 312 |
+
)
|
| 313 |
+
|
| 314 |
+
trainer = CooperativeTrainer(
|
| 315 |
+
model=models["model"],
|
| 316 |
+
tokenizer=tokenizer,
|
| 317 |
+
config=cfg,
|
| 318 |
+
generator=models["generator"],
|
| 319 |
+
verifier=models["verifier"],
|
| 320 |
+
anti_hallucination=models["anti_hall"],
|
| 321 |
+
ewc_state=ewc_state,
|
| 322 |
+
)
|
| 323 |
+
|
| 324 |
+
try:
|
| 325 |
+
result = trainer.train(dataloader)
|
| 326 |
+
logger.info(" Treinamento OK | final_loss=%.4f | final_ppl=%.2f | elapsed=%.1fs",
|
| 327 |
+
result["final_loss"], result["final_ppl"], result["elapsed_s"])
|
| 328 |
+
return result
|
| 329 |
+
except Exception as e:
|
| 330 |
+
logger.error(" Treinamento FALHOU: %s", e)
|
| 331 |
+
logger.error(traceback.format_exc())
|
| 332 |
+
raise
|
| 333 |
+
|
| 334 |
+
|
| 335 |
+
def step_05b_test_new_modules(
|
| 336 |
+
models: Dict,
|
| 337 |
+
tokenizer: BBPETokenizer,
|
| 338 |
+
device: str = "cpu",
|
| 339 |
+
) -> Dict:
|
| 340 |
+
"""Testa os novos módulos v2.0: EWC, Context Window, RoPE, TransformerBlock.
|
| 341 |
+
|
| 342 |
+
NOVO v2.0: esta função exercita isoladamente cada novo módulo para garantir
|
| 343 |
+
que estão funcionando e integráveis ao pipeline.
|
| 344 |
+
"""
|
| 345 |
+
logger.info("=" * 70)
|
| 346 |
+
logger.info("PASSO 5b: Teste dos novos módulos v2.0 (EWC, ContextWindow, RoPE, TransformerBlock)")
|
| 347 |
+
logger.info("=" * 70)
|
| 348 |
+
|
| 349 |
+
results = {"ewc": None, "context_window": None, "rope": None, "transformer": None}
|
| 350 |
+
|
| 351 |
+
# ---------- EWC ----------
|
| 352 |
+
try:
|
| 353 |
+
logger.info(" [EWC] Testando EWCState...")
|
| 354 |
+
ewc_cfg = EWCConfig(
|
| 355 |
+
enabled=True,
|
| 356 |
+
lambda_ewc=100.0,
|
| 357 |
+
n_samples_fisher=5, # poucas amostras para teste rápido
|
| 358 |
+
online_gamma=0.9,
|
| 359 |
+
device=device,
|
| 360 |
+
)
|
| 361 |
+
ewc_state = EWCState(ewc_cfg)
|
| 362 |
+
# Antes de consolidar: penalty deve ser 0
|
| 363 |
+
pen0 = ewc_state.penalty(models["model"])
|
| 364 |
+
logger.info(" Penalty antes de consolidar: %.6f", float(pen0))
|
| 365 |
+
assert pen0.item() == 0.0, "EWC penalty deveria ser 0 antes de consolidar"
|
| 366 |
+
|
| 367 |
+
# Simular forward_fn para cálculo de Fisher
|
| 368 |
+
# IMPORTANTE: NÃO usar torch.no_grad() aqui — a EWC precisa de gradientes
|
| 369 |
+
# para calcular a matriz de Fisher (grad^2 da log-verossimilhança)
|
| 370 |
+
V = tokenizer.vocab_size
|
| 371 |
+
B = 2
|
| 372 |
+
T = 8
|
| 373 |
+
ids_a = torch.randint(0, V, (B, T), device=device)
|
| 374 |
+
ids_b = torch.randint(0, V, (B, T), device=device)
|
| 375 |
+
|
| 376 |
+
def forward_fn(_idx=None):
|
| 377 |
+
# Sem torch.no_grad() — a EWC interna ativa enable_grad()
|
| 378 |
+
out = models["model"](ids_a, ids_b, mode="classify")
|
| 379 |
+
return out["logits"]
|
| 380 |
+
|
| 381 |
+
# Consolidar (calcula Fisher e armazena theta_star)
|
| 382 |
+
ewc_state.consolidate(models["model"], forward_fn=forward_fn)
|
| 383 |
+
assert ewc_state.num_tasks() == 1, f"Esperado 1 tarefa, obtido {ewc_state.num_tasks()}"
|
| 384 |
+
|
| 385 |
+
# Penalty deve ser > 0 agora (parâmetros não mudaram desde consolidate, mas
|
| 386 |
+
# a penalidade ainda assim é computada e deve ser >= 0)
|
| 387 |
+
pen1 = ewc_state.penalty(models["model"])
|
| 388 |
+
logger.info(" Penalty após consolidar: %.6f", float(pen1.detach()))
|
| 389 |
+
assert pen1.item() >= 0.0, "EWC penalty deve ser >= 0"
|
| 390 |
+
|
| 391 |
+
results["ewc"] = {
|
| 392 |
+
"ok": True,
|
| 393 |
+
"penalty_before": float(pen0),
|
| 394 |
+
"penalty_after": float(pen1),
|
| 395 |
+
"num_tasks": ewc_state.num_tasks(),
|
| 396 |
+
}
|
| 397 |
+
logger.info(" [EWC] ✓ OK")
|
| 398 |
+
|
| 399 |
+
# Guardar para uso posterior no trainer
|
| 400 |
+
results["ewc_state"] = ewc_state
|
| 401 |
+
except Exception as e:
|
| 402 |
+
logger.error(" [EWC] ✗ FALHOU: %s", e)
|
| 403 |
+
logger.error(traceback.format_exc())
|
| 404 |
+
results["ewc"] = {"ok": False, "error": str(e)}
|
| 405 |
+
|
| 406 |
+
# ---------- Context Window ----------
|
| 407 |
+
try:
|
| 408 |
+
logger.info(" [ContextWindow] Testando ContextWindowManager + KVCache...")
|
| 409 |
+
cw = make_default_context_window(
|
| 410 |
+
max_window=32,
|
| 411 |
+
n_sink=2,
|
| 412 |
+
embed_dim=64,
|
| 413 |
+
n_heads=4,
|
| 414 |
+
n_layers=2,
|
| 415 |
+
device=device,
|
| 416 |
+
)
|
| 417 |
+
# Inicializar cache
|
| 418 |
+
cache = cw.init_cache(batch_size=1, device=torch.device(device))
|
| 419 |
+
assert cache is not None
|
| 420 |
+
assert cache.get_seq_len() == 0
|
| 421 |
+
|
| 422 |
+
# Append tokens
|
| 423 |
+
tokens1 = torch.tensor([[1, 2, 3, 4, 5]], dtype=torch.long, device=device)
|
| 424 |
+
window = cw.append_tokens(tokens1)
|
| 425 |
+
assert window.size(1) == 5
|
| 426 |
+
assert cw.kv_cache.get_seq_len() == 0 # cache só é populado quando chamamos KVCache.update
|
| 427 |
+
|
| 428 |
+
# Simular update do cache KV
|
| 429 |
+
dummy_k = torch.randn(1, 4, 5, 16, device=device) # [B, n_heads, T, head_dim]
|
| 430 |
+
dummy_v = torch.randn(1, 4, 5, 16, device=device)
|
| 431 |
+
new_k, new_v = cache.update(0, dummy_k, dummy_v)
|
| 432 |
+
assert cache.get_seq_len() == 5
|
| 433 |
+
|
| 434 |
+
# Adicionar mais tokens para forçar eviction
|
| 435 |
+
tokens2 = torch.tensor([[6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20]],
|
| 436 |
+
dtype=torch.long, device=device)
|
| 437 |
+
cw.append_tokens(tokens2)
|
| 438 |
+
# Total histórico: 5 + 15 = 20, mas max_window=32, então sem eviction ainda
|
| 439 |
+
assert cw.token_history and sum(t.size(1) for t in cw.token_history) == 20
|
| 440 |
+
|
| 441 |
+
# Forçar eviction adicionando mais tokens
|
| 442 |
+
tokens3 = torch.tensor([[21, 22, 23, 24, 25, 26, 27, 28, 29, 30, 31, 32, 33, 34]],
|
| 443 |
+
dtype=torch.long, device=device)
|
| 444 |
+
cw.append_tokens(tokens3) # total = 34 > 32 = max_window
|
| 445 |
+
cw.evict_cache()
|
| 446 |
+
|
| 447 |
+
info = cw.get_info()
|
| 448 |
+
logger.info(" Context window info: %s", info)
|
| 449 |
+
results["context_window"] = {"ok": True, "info": info}
|
| 450 |
+
logger.info(" [ContextWindow] ✓ OK")
|
| 451 |
+
except Exception as e:
|
| 452 |
+
logger.error(" [ContextWindow] ✗ FALHOU: %s", e)
|
| 453 |
+
logger.error(traceback.format_exc())
|
| 454 |
+
results["context_window"] = {"ok": False, "error": str(e)}
|
| 455 |
+
|
| 456 |
+
# ---------- RoPE ----------
|
| 457 |
+
try:
|
| 458 |
+
logger.info(" [RoPE] Testando RotaryPositionEmbedding...")
|
| 459 |
+
head_dim = 32
|
| 460 |
+
max_seq = 16
|
| 461 |
+
rope = RotaryPositionEmbedding(head_dim=head_dim, max_seq_len=max_seq).to(device)
|
| 462 |
+
# x: [B, n_heads, T, head_dim]
|
| 463 |
+
x = torch.randn(2, 4, 8, head_dim, device=device)
|
| 464 |
+
x_rot = rope(x)
|
| 465 |
+
assert x_rot.shape == x.shape, f"Shape mismatch: {x_rot.shape} vs {x.shape}"
|
| 466 |
+
# Verificar que RoPE preserva a norma (é uma rotação)
|
| 467 |
+
norm_before = x.norm(dim=-1)
|
| 468 |
+
norm_after = x_rot.norm(dim=-1)
|
| 469 |
+
diff = (norm_before - norm_after).abs().max().item()
|
| 470 |
+
logger.info(" Norma preservada (diff=%.6f)", diff)
|
| 471 |
+
assert diff < 1e-4, f"RoPE não preservou norma (diff={diff})"
|
| 472 |
+
|
| 473 |
+
# Testar com offset (para geração autoregressiva)
|
| 474 |
+
x_rot2 = rope(x, offset=10)
|
| 475 |
+
assert x_rot2.shape == x.shape
|
| 476 |
+
|
| 477 |
+
results["rope"] = {"ok": True, "norm_diff": diff}
|
| 478 |
+
logger.info(" [RoPE] ✓ OK")
|
| 479 |
+
except Exception as e:
|
| 480 |
+
logger.error(" [RoPE] ✗ FALHOU: %s", e)
|
| 481 |
+
logger.error(traceback.format_exc())
|
| 482 |
+
results["rope"] = {"ok": False, "error": str(e)}
|
| 483 |
+
|
| 484 |
+
# ---------- TransformerBlock ----------
|
| 485 |
+
try:
|
| 486 |
+
logger.info(" [TransformerBlock] Testando CausalSelfAttention + TransformerBlock...")
|
| 487 |
+
embed_dim = 64
|
| 488 |
+
n_heads = 4
|
| 489 |
+
n_layers = 2
|
| 490 |
+
V = tokenizer.vocab_size
|
| 491 |
+
|
| 492 |
+
block_config = TransformerBlockConfig(
|
| 493 |
+
embed_dim=embed_dim,
|
| 494 |
+
n_heads=n_heads,
|
| 495 |
+
ff_dim=4 * embed_dim,
|
| 496 |
+
dropout=0.1,
|
| 497 |
+
use_rope=True,
|
| 498 |
+
max_seq_len=32,
|
| 499 |
+
)
|
| 500 |
+
block = TransformerBlock(block_config).to(device)
|
| 501 |
+
|
| 502 |
+
# Forward sem cache
|
| 503 |
+
x = torch.randn(2, 8, embed_dim, device=device)
|
| 504 |
+
out, k, v = block(x)
|
| 505 |
+
assert out.shape == x.shape, f"Output shape: {out.shape} vs {x.shape}"
|
| 506 |
+
assert k.shape == (2, n_heads, 8, embed_dim // n_heads)
|
| 507 |
+
assert v.shape == (2, n_heads, 8, embed_dim // n_heads)
|
| 508 |
+
|
| 509 |
+
# Forward com padding mask
|
| 510 |
+
mask = torch.tensor([[1, 1, 1, 1, 0, 0, 0, 0], [1, 1, 1, 1, 1, 1, 1, 0]],
|
| 511 |
+
dtype=torch.float, device=device)
|
| 512 |
+
out2, k2, v2 = block(x, padding_mask=mask)
|
| 513 |
+
assert out2.shape == x.shape
|
| 514 |
+
|
| 515 |
+
# Forward com cache KV
|
| 516 |
+
cache_k = None
|
| 517 |
+
cache_v = None
|
| 518 |
+
for t in range(3):
|
| 519 |
+
x_t = torch.randn(2, 1, embed_dim, device=device)
|
| 520 |
+
out_t, cache_k, cache_v = block(
|
| 521 |
+
x_t, kv_cache_k=cache_k, kv_cache_v=cache_v, position_offset=t
|
| 522 |
+
)
|
| 523 |
+
assert out_t.shape == (2, 1, embed_dim)
|
| 524 |
+
assert cache_k.size(2) == t + 1, f"Cache len: {cache_k.size(2)}, esperado {t+1}"
|
| 525 |
+
logger.info(" Cache KV após 3 steps: %d tokens", cache_k.size(2))
|
| 526 |
+
|
| 527 |
+
# Testar TransformerDecoderStack completo
|
| 528 |
+
decoder = TransformerDecoderStack(
|
| 529 |
+
vocab_size=V,
|
| 530 |
+
embed_dim=embed_dim,
|
| 531 |
+
n_heads=n_heads,
|
| 532 |
+
n_layers=n_layers,
|
| 533 |
+
max_seq_len=32,
|
| 534 |
+
use_rope=True,
|
| 535 |
+
pad_id=tokenizer.pad_id,
|
| 536 |
+
weight_tying=True,
|
| 537 |
+
).to(device)
|
| 538 |
+
|
| 539 |
+
idx = torch.randint(0, V, (2, 8), device=device)
|
| 540 |
+
logits, new_caches = decoder(idx)
|
| 541 |
+
assert logits.shape == (2, 8, V), f"Logits shape: {logits.shape}"
|
| 542 |
+
assert len(new_caches) == n_layers
|
| 543 |
+
|
| 544 |
+
# Verificar weight tying
|
| 545 |
+
assert decoder.lm_head.weight is decoder.token_embedding.weight
|
| 546 |
+
|
| 547 |
+
results["transformer"] = {
|
| 548 |
+
"ok": True,
|
| 549 |
+
"block_output_shape": list(out.shape),
|
| 550 |
+
"decoder_logits_shape": list(logits.shape),
|
| 551 |
+
"weight_tying": True,
|
| 552 |
+
}
|
| 553 |
+
logger.info(" [TransformerBlock] ✓ OK")
|
| 554 |
+
except Exception as e:
|
| 555 |
+
logger.error(" [TransformerBlock] ✗ FALHOU: %s", e)
|
| 556 |
+
logger.error(traceback.format_exc())
|
| 557 |
+
results["transformer"] = {"ok": False, "error": str(e)}
|
| 558 |
+
|
| 559 |
+
return results
|
| 560 |
+
|
| 561 |
+
|
| 562 |
+
def step_06_inference(
|
| 563 |
+
model: MultimodalCNNBiGRU,
|
| 564 |
+
tokenizer: BBPETokenizer,
|
| 565 |
+
device: str = "cpu",
|
| 566 |
+
generator: GeneratorCNNBiGRU = None,
|
| 567 |
+
use_context_window: bool = True,
|
| 568 |
+
) -> Dict:
|
| 569 |
+
"""Executa inferência com sampling.
|
| 570 |
+
|
| 571 |
+
NOVO v2.0: usa generator (se disponível) para geração autoregressiva REAL,
|
| 572 |
+
e integra context_window para suporte a sequências longas.
|
| 573 |
+
"""
|
| 574 |
+
logger.info("=" * 70)
|
| 575 |
+
logger.info("PASSO 6: Inferência com Temperatura + Top-K + Top-P + Penalidades")
|
| 576 |
+
if generator is not None:
|
| 577 |
+
logger.info(" (usando GeneratorCNNBiGRU para geração autoregressiva REAL)")
|
| 578 |
+
else:
|
| 579 |
+
logger.info(" (sem generator — usando fallback single-step)")
|
| 580 |
+
logger.info("=" * 70)
|
| 581 |
+
|
| 582 |
+
prompts = [
|
| 583 |
+
("o modelo coopera entre", "fluxos paralelos"),
|
| 584 |
+
("atencao cruzada troca", "informacoes entre redes"),
|
| 585 |
+
("gru bidirecional captura", "dependencias temporais"),
|
| 586 |
+
]
|
| 587 |
+
|
| 588 |
+
# Inicializar context window se solicitado
|
| 589 |
+
cw = None
|
| 590 |
+
if use_context_window:
|
| 591 |
+
try:
|
| 592 |
+
cw = make_default_context_window(
|
| 593 |
+
max_window=64,
|
| 594 |
+
n_sink=2,
|
| 595 |
+
embed_dim=32,
|
| 596 |
+
n_heads=4,
|
| 597 |
+
n_layers=2,
|
| 598 |
+
device=device,
|
| 599 |
+
)
|
| 600 |
+
logger.info(" Context window ativo (max_window=64, sink=2)")
|
| 601 |
+
except Exception as e:
|
| 602 |
+
logger.warning(" Context window falhou ao inicializar: %s — desativando", e)
|
| 603 |
+
cw = None
|
| 604 |
+
|
| 605 |
+
results = []
|
| 606 |
+
for prompt_a, prompt_b in prompts:
|
| 607 |
+
try:
|
| 608 |
+
result = generate_with_sampling(
|
| 609 |
+
model=model,
|
| 610 |
+
tokenizer=tokenizer,
|
| 611 |
+
prompt_a=prompt_a,
|
| 612 |
+
prompt_b=prompt_b,
|
| 613 |
+
max_new_tokens=16,
|
| 614 |
+
temperature=0.7,
|
| 615 |
+
top_k=20,
|
| 616 |
+
top_p=0.9,
|
| 617 |
+
presence_penalty=0.3,
|
| 618 |
+
frequency_penalty=0.3,
|
| 619 |
+
device=device,
|
| 620 |
+
generator=generator,
|
| 621 |
+
context_window=cw,
|
| 622 |
+
)
|
| 623 |
+
logger.info(" Prompt A: %s", prompt_a)
|
| 624 |
+
logger.info(" Prompt B: %s", prompt_b)
|
| 625 |
+
logger.info(" Gerado: %s", result["text"][:100])
|
| 626 |
+
logger.info(" Metrics: %s", result["metrics"])
|
| 627 |
+
results.append({"prompt_a": prompt_a, "prompt_b": prompt_b, **result})
|
| 628 |
+
except Exception as e:
|
| 629 |
+
logger.error(" Inferência falhou para (%s, %s): %s", prompt_a, prompt_b, e)
|
| 630 |
+
logger.error(traceback.format_exc())
|
| 631 |
+
results.append({"prompt_a": prompt_a, "prompt_b": prompt_b, "error": str(e)})
|
| 632 |
+
|
| 633 |
+
return {"results": results}
|
| 634 |
+
|
| 635 |
+
|
| 636 |
+
def step_07_eval_ppl(
|
| 637 |
+
model: MultimodalCNNBiGRU,
|
| 638 |
+
tokenizer: BBPETokenizer,
|
| 639 |
+
dataloader: DataLoader,
|
| 640 |
+
device: str = "cpu",
|
| 641 |
+
generator: GeneratorCNNBiGRU = None,
|
| 642 |
+
) -> Dict:
|
| 643 |
+
"""Avalia perplexidade.
|
| 644 |
+
|
| 645 |
+
NOVO v2.0: usa generator (se disponível) para PPL real baseado em
|
| 646 |
+
geração autoregressiva com teacher forcing.
|
| 647 |
+
"""
|
| 648 |
+
logger.info("=" * 70)
|
| 649 |
+
logger.info("PASSO 7: Avaliação de Perplexidade (PPL)")
|
| 650 |
+
if generator is not None:
|
| 651 |
+
logger.info(" (usando GeneratorCNNBiGRU com teacher forcing)")
|
| 652 |
+
else:
|
| 653 |
+
logger.info(" (sem generator — usando fallback single-step)")
|
| 654 |
+
logger.info("=" * 70)
|
| 655 |
+
try:
|
| 656 |
+
result = evaluate_perplexity(
|
| 657 |
+
model=model,
|
| 658 |
+
dataloader=dataloader,
|
| 659 |
+
tokenizer=tokenizer,
|
| 660 |
+
device=device,
|
| 661 |
+
max_batches=5,
|
| 662 |
+
generator=generator,
|
| 663 |
+
)
|
| 664 |
+
logger.info(" Loss: %.4f | PPL: %.2f | batches: %d | used_generator: %s",
|
| 665 |
+
result["loss"], result["ppl"], result["n_batches"],
|
| 666 |
+
result.get("used_generator", False))
|
| 667 |
+
return result
|
| 668 |
+
except Exception as e:
|
| 669 |
+
logger.error(" PPL falhou: %s", e)
|
| 670 |
+
logger.error(traceback.format_exc())
|
| 671 |
+
return {"error": str(e)}
|
| 672 |
+
|
| 673 |
+
|
| 674 |
+
def step_08_report_errors(errors: List[str], warnings: List[str]) -> None:
|
| 675 |
+
"""Reporta erros lógicos ou falhas encontradas."""
|
| 676 |
+
logger.info("=" * 70)
|
| 677 |
+
logger.info("PASSO 8: Relatório de erros lógicos ou falhas")
|
| 678 |
+
logger.info("=" * 70)
|
| 679 |
+
|
| 680 |
+
if not errors:
|
| 681 |
+
logger.info(" ✓ NENHUM ERRO CRÍTICO encontrado")
|
| 682 |
+
else:
|
| 683 |
+
logger.error(" ✗ %d ERRO(S) CRÍTICO(S) encontrado(s):", len(errors))
|
| 684 |
+
for e in errors:
|
| 685 |
+
logger.error(" - %s", e)
|
| 686 |
+
|
| 687 |
+
if not warnings:
|
| 688 |
+
logger.info(" ✓ Nenhum warning relevante")
|
| 689 |
+
else:
|
| 690 |
+
logger.warning(" ⚠ %d warning(s):", len(warnings))
|
| 691 |
+
for w in warnings:
|
| 692 |
+
logger.warning(" - %s", w)
|
| 693 |
+
|
| 694 |
+
logger.info("=" * 70)
|
| 695 |
+
|
| 696 |
+
|
| 697 |
+
def main():
|
| 698 |
+
"""Executa o teste completo de 50 amostras."""
|
| 699 |
+
logger.info("INICIANDO TESTE DE 50 AMOSTRAS — CNN-BiGRU MULTIMODAL COOPERATIVO")
|
| 700 |
+
logger.info("Project root: %s", PROJECT_ROOT)
|
| 701 |
+
logger.info("Device: %s", "cpu")
|
| 702 |
+
logger.info("Cores: %d", N_CORES)
|
| 703 |
+
|
| 704 |
+
errors: List[str] = []
|
| 705 |
+
warnings: List[str] = []
|
| 706 |
+
|
| 707 |
+
try:
|
| 708 |
+
# Step 1
|
| 709 |
+
step_01_init_runtime()
|
| 710 |
+
|
| 711 |
+
# Step 2
|
| 712 |
+
tokenizer, roundtrip = step_02_train_tokenizer()
|
| 713 |
+
if roundtrip < 0.8:
|
| 714 |
+
warnings.append(f"BBPE roundtrip accuracy = {roundtrip:.2%} (esperado >= 80%)")
|
| 715 |
+
|
| 716 |
+
# Step 3
|
| 717 |
+
dataloader = step_03_create_dataset(tokenizer, n_samples=50)
|
| 718 |
+
|
| 719 |
+
# Step 4
|
| 720 |
+
models = step_04_init_model(tokenizer, device="cpu")
|
| 721 |
+
|
| 722 |
+
# Verificação: forward pass simples antes do treino
|
| 723 |
+
try:
|
| 724 |
+
batch = next(iter(dataloader))
|
| 725 |
+
with torch.no_grad():
|
| 726 |
+
out = models["model"](
|
| 727 |
+
batch["input_ids_a"], batch["input_ids_b"],
|
| 728 |
+
images=batch["images"], audios=batch["audios"],
|
| 729 |
+
mode="classify",
|
| 730 |
+
)
|
| 731 |
+
logger.info(" Forward pass OK | logits shape: %s", out["logits"].shape)
|
| 732 |
+
assert out["logits"].shape == (batch["input_ids_a"].size(0), 3), \
|
| 733 |
+
f"Shape inesperado: {out['logits'].shape}"
|
| 734 |
+
except Exception as e:
|
| 735 |
+
errors.append(f"Forward pass inicial falhou: {e}")
|
| 736 |
+
logger.error(traceback.format_exc())
|
| 737 |
+
|
| 738 |
+
# Step 5b: Testar novos módulos v2.0 (EWC, Context Window, RoPE, TransformerBlock)
|
| 739 |
+
ewc_state = None
|
| 740 |
+
if not errors:
|
| 741 |
+
try:
|
| 742 |
+
new_modules_result = step_05b_test_new_modules(models, tokenizer, device="cpu")
|
| 743 |
+
# Verificar se algum módulo falhou
|
| 744 |
+
for mod_name, mod_res in new_modules_result.items():
|
| 745 |
+
if mod_name == "ewc_state":
|
| 746 |
+
continue
|
| 747 |
+
if isinstance(mod_res, dict) and not mod_res.get("ok", True):
|
| 748 |
+
errors.append(f"Módulo {mod_name} falhou: {mod_res.get('error', 'unknown')}")
|
| 749 |
+
# Guardar EWC state para usar no treino
|
| 750 |
+
ewc_state = new_modules_result.get("ewc_state")
|
| 751 |
+
if ewc_state is None:
|
| 752 |
+
warnings.append("EWC state não foi criado — EWC não será testado no treino")
|
| 753 |
+
except Exception as e:
|
| 754 |
+
errors.append(f"Teste de novos módulos falhou: {e}")
|
| 755 |
+
logger.error(traceback.format_exc())
|
| 756 |
+
|
| 757 |
+
# Step 5: Treinamento (com EWC se disponível)
|
| 758 |
+
if not errors:
|
| 759 |
+
try:
|
| 760 |
+
train_result = step_05_train(models, tokenizer, dataloader, device="cpu",
|
| 761 |
+
ewc_state=ewc_state)
|
| 762 |
+
except Exception as e:
|
| 763 |
+
errors.append(f"Treinamento falhou: {e}")
|
| 764 |
+
train_result = None
|
| 765 |
+
else:
|
| 766 |
+
train_result = None
|
| 767 |
+
|
| 768 |
+
# Step 6: Inferência com generator + context window
|
| 769 |
+
if not errors:
|
| 770 |
+
try:
|
| 771 |
+
step_06_inference(
|
| 772 |
+
models["model"], tokenizer, device="cpu",
|
| 773 |
+
generator=models.get("generator"),
|
| 774 |
+
use_context_window=True,
|
| 775 |
+
)
|
| 776 |
+
except Exception as e:
|
| 777 |
+
errors.append(f"Inferência falhou: {e}")
|
| 778 |
+
logger.error(traceback.format_exc())
|
| 779 |
+
|
| 780 |
+
# Step 7: PPL com generator
|
| 781 |
+
if not errors:
|
| 782 |
+
try:
|
| 783 |
+
step_07_eval_ppl(
|
| 784 |
+
models["model"], tokenizer, dataloader, device="cpu",
|
| 785 |
+
generator=models.get("generator"),
|
| 786 |
+
)
|
| 787 |
+
except Exception as e:
|
| 788 |
+
warnings.append(f"PPL falhou (não crítico): {e}")
|
| 789 |
+
|
| 790 |
+
# Step 8
|
| 791 |
+
step_08_report_errors(errors, warnings)
|
| 792 |
+
|
| 793 |
+
# Resumo final
|
| 794 |
+
logger.info("=" * 70)
|
| 795 |
+
logger.info("RESUMO FINAL DO TESTE DE 50 AMOSTRAS")
|
| 796 |
+
logger.info("=" * 70)
|
| 797 |
+
logger.info(" Amostras processadas: 50")
|
| 798 |
+
logger.info(" Erros críticos: %d", len(errors))
|
| 799 |
+
logger.info(" Warnings: %d", len(warnings))
|
| 800 |
+
if train_result:
|
| 801 |
+
logger.info(" Loss final: %.4f", train_result["final_loss"])
|
| 802 |
+
logger.info(" PPL final: %.2f", train_result["final_ppl"])
|
| 803 |
+
logger.info(" Tempo: %.1fs", train_result["elapsed_s"])
|
| 804 |
+
logger.info(" Synergy attempts: %d", len(train_result.get("synergy_history", [])))
|
| 805 |
+
logger.info(" Hypothesis activations: %d",
|
| 806 |
+
sum(h.get("hypothesis_activations", 0) for h in train_result.get("history", [])))
|
| 807 |
+
|
| 808 |
+
if errors:
|
| 809 |
+
logger.error(" STATUS: FALHA — %d erro(s)", len(errors))
|
| 810 |
+
return 1
|
| 811 |
+
else:
|
| 812 |
+
logger.info(" STATUS: SUCESSO")
|
| 813 |
+
return 0
|
| 814 |
+
|
| 815 |
+
except Exception as e:
|
| 816 |
+
logger.error("ERRO FATAL: %s", e)
|
| 817 |
+
logger.error(traceback.format_exc())
|
| 818 |
+
return 2
|
| 819 |
+
|
| 820 |
+
|
| 821 |
+
if __name__ == "__main__":
|
| 822 |
+
try:
|
| 823 |
+
rc = main()
|
| 824 |
+
except SystemExit:
|
| 825 |
+
raise
|
| 826 |
+
except Exception as e:
|
| 827 |
+
logger.error("Unhandled exception: %s", e)
|
| 828 |
+
rc = 2
|
| 829 |
+
# Evita o "Fatal Python error: PyGILState_Release" no shutdown causado
|
| 830 |
+
# por threads do tokenizers/datasets que ainda estão ativas.
|
| 831 |
+
import os
|
| 832 |
+
os._exit(rc)
|
cnn_bigru/tokenizer/__init__.py
ADDED
|
File without changes
|
cnn_bigru/tokenizer/bbpe_tokenizer.py
ADDED
|
@@ -0,0 +1,173 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""bbpe_tokenizer.py — BBPE (Byte-Level BPE) Tokenizer simplificado.
|
| 2 |
+
|
| 3 |
+
Inspirado em PowerMachine/BiGRU_T_version/src/bigru_t/tokenizer/bbpe_tokenizer.py,
|
| 4 |
+
mas simplificado para o projeto CNN-BiGRU. Usa a biblioteca `tokenizers` da
|
| 5 |
+
HuggingFace com modelo BPE byte-level.
|
| 6 |
+
|
| 7 |
+
Cobertura universal UTF-8 (qualquer string é tokenizável sem <unk>).
|
| 8 |
+
"""
|
| 9 |
+
from __future__ import annotations
|
| 10 |
+
|
| 11 |
+
import json
|
| 12 |
+
import logging
|
| 13 |
+
import os
|
| 14 |
+
from pathlib import Path
|
| 15 |
+
from typing import List, Optional, Sequence
|
| 16 |
+
|
| 17 |
+
from tokenizers import Tokenizer
|
| 18 |
+
from tokenizers.models import BPE
|
| 19 |
+
from tokenizers.pre_tokenizers import ByteLevel
|
| 20 |
+
from tokenizers.processors import ByteLevel as ByteLevelProcessor
|
| 21 |
+
from tokenizers.decoders import ByteLevel as ByteLevelDecoder
|
| 22 |
+
from tokenizers.trainers import BpeTrainer
|
| 23 |
+
|
| 24 |
+
logger = logging.getLogger(__name__)
|
| 25 |
+
|
| 26 |
+
BOS_TOKEN = "<s>"
|
| 27 |
+
PAD_TOKEN = "<pad>"
|
| 28 |
+
EOS_TOKEN = "</s>"
|
| 29 |
+
UNK_TOKEN = "<unk>"
|
| 30 |
+
|
| 31 |
+
SPECIAL_TOKENS = [BOS_TOKEN, PAD_TOKEN, EOS_TOKEN, UNK_TOKEN]
|
| 32 |
+
BOS_ID = 0
|
| 33 |
+
PAD_ID = 1
|
| 34 |
+
EOS_ID = 2
|
| 35 |
+
UNK_ID = 3
|
| 36 |
+
|
| 37 |
+
|
| 38 |
+
class BBPETokenizer:
|
| 39 |
+
"""Wrapper sobre Tokenizer (HuggingFace) com API conveniente.
|
| 40 |
+
|
| 41 |
+
Garante: dec(enc(s)) = s para todo s UTF-8.
|
| 42 |
+
"""
|
| 43 |
+
|
| 44 |
+
def __init__(self, tokenizer: Optional[Tokenizer] = None):
|
| 45 |
+
self._tok = tokenizer
|
| 46 |
+
|
| 47 |
+
@property
|
| 48 |
+
def tokenizer(self) -> Tokenizer:
|
| 49 |
+
if self._tok is None:
|
| 50 |
+
raise RuntimeError("Tokenizer não inicializado. Use train_from_* ou load.")
|
| 51 |
+
return self._tok
|
| 52 |
+
|
| 53 |
+
@property
|
| 54 |
+
def vocab_size(self) -> int:
|
| 55 |
+
return self.tokenizer.get_vocab_size()
|
| 56 |
+
|
| 57 |
+
@property
|
| 58 |
+
def bos_id(self) -> int:
|
| 59 |
+
return self.tokenizer.token_to_id(BOS_TOKEN) or BOS_ID
|
| 60 |
+
|
| 61 |
+
@property
|
| 62 |
+
def pad_id(self) -> int:
|
| 63 |
+
return self.tokenizer.token_to_id(PAD_TOKEN) or PAD_ID
|
| 64 |
+
|
| 65 |
+
@property
|
| 66 |
+
def eos_id(self) -> int:
|
| 67 |
+
return self.tokenizer.token_to_id(EOS_TOKEN) or EOS_ID
|
| 68 |
+
|
| 69 |
+
@property
|
| 70 |
+
def unk_id(self) -> int:
|
| 71 |
+
return self.tokenizer.token_to_id(UNK_TOKEN) or UNK_ID
|
| 72 |
+
|
| 73 |
+
def encode(self, text: str, add_special: bool = True) -> List[int]:
|
| 74 |
+
if add_special:
|
| 75 |
+
text = f"{BOS_TOKEN}{text}{EOS_TOKEN}"
|
| 76 |
+
return self.tokenizer.encode(text).ids
|
| 77 |
+
|
| 78 |
+
def encode_batch(self, texts: Sequence[str], add_special: bool = True) -> List[List[int]]:
|
| 79 |
+
if add_special:
|
| 80 |
+
texts = [f"{BOS_TOKEN}{t}{EOS_TOKEN}" for t in texts]
|
| 81 |
+
return [e.ids for e in self.tokenizer.encode_batch(list(texts))]
|
| 82 |
+
|
| 83 |
+
def decode(self, ids: List[int], skip_special: bool = True) -> str:
|
| 84 |
+
if skip_special:
|
| 85 |
+
ids = [i for i in ids if i not in (self.bos_id, self.pad_id, self.eos_id, self.unk_id)]
|
| 86 |
+
return self.tokenizer.decode(ids)
|
| 87 |
+
|
| 88 |
+
def save(self, path: str) -> None:
|
| 89 |
+
Path(path).parent.mkdir(parents=True, exist_ok=True)
|
| 90 |
+
self.tokenizer.save(path)
|
| 91 |
+
logger.info("Tokenizer salvo em %s", path)
|
| 92 |
+
|
| 93 |
+
@classmethod
|
| 94 |
+
def load(cls, path: str) -> "BBPETokenizer":
|
| 95 |
+
tok = Tokenizer.from_file(path)
|
| 96 |
+
return cls(tok)
|
| 97 |
+
|
| 98 |
+
@classmethod
|
| 99 |
+
def train_from_files(
|
| 100 |
+
cls,
|
| 101 |
+
files: Sequence[str],
|
| 102 |
+
vocab_size: int = 8000,
|
| 103 |
+
min_frequency: int = 2,
|
| 104 |
+
) -> "BBPETokenizer":
|
| 105 |
+
# Desativa paralelismo de tokenizers para evitar problemas de GIL no shutdown
|
| 106 |
+
os.environ["TOKENIZERS_PARALLELISM"] = "false"
|
| 107 |
+
tokenizer = Tokenizer(BPE(unk_token=UNK_TOKEN))
|
| 108 |
+
tokenizer.pre_tokenizer = ByteLevel(add_prefix_space=False)
|
| 109 |
+
trainer = BpeTrainer(
|
| 110 |
+
vocab_size=vocab_size,
|
| 111 |
+
min_frequency=min_frequency,
|
| 112 |
+
special_tokens=SPECIAL_TOKENS,
|
| 113 |
+
show_progress=False,
|
| 114 |
+
)
|
| 115 |
+
valid_files = [str(Path(p).resolve()) for p in files if Path(p).exists()]
|
| 116 |
+
if not valid_files:
|
| 117 |
+
raise FileNotFoundError(f"Nenhum arquivo válido em: {files}")
|
| 118 |
+
tokenizer.train(valid_files, trainer)
|
| 119 |
+
tokenizer.post_processor = ByteLevelProcessor(trim_offsets=False)
|
| 120 |
+
# Decoder byte-level: reverte a codificação byte-level para UTF-8
|
| 121 |
+
tokenizer.decoder = ByteLevelDecoder()
|
| 122 |
+
logger.info("BBPE treinado com vocab_size=%d", tokenizer.get_vocab_size())
|
| 123 |
+
return cls(tokenizer)
|
| 124 |
+
|
| 125 |
+
@classmethod
|
| 126 |
+
def train_from_texts(
|
| 127 |
+
cls,
|
| 128 |
+
texts: Sequence[str],
|
| 129 |
+
vocab_size: int = 8000,
|
| 130 |
+
min_frequency: int = 1,
|
| 131 |
+
) -> "BBPETokenizer":
|
| 132 |
+
"""Treina a partir de uma lista de textos em memória."""
|
| 133 |
+
# Salva temporariamente
|
| 134 |
+
import tempfile
|
| 135 |
+
with tempfile.NamedTemporaryFile(
|
| 136 |
+
mode="w", suffix=".txt", delete=False, encoding="utf-8"
|
| 137 |
+
) as f:
|
| 138 |
+
for line in texts:
|
| 139 |
+
line = line.strip()
|
| 140 |
+
if line:
|
| 141 |
+
f.write(line + "\n")
|
| 142 |
+
tmp_path = f.name
|
| 143 |
+
try:
|
| 144 |
+
return cls.train_from_files([tmp_path], vocab_size=vocab_size,
|
| 145 |
+
min_frequency=min_frequency)
|
| 146 |
+
finally:
|
| 147 |
+
try:
|
| 148 |
+
os.unlink(tmp_path)
|
| 149 |
+
except OSError:
|
| 150 |
+
pass
|
| 151 |
+
|
| 152 |
+
def validate_roundtrip(self, test_texts: Sequence[str]) -> float:
|
| 153 |
+
"""Valida que dec(enc(s)) == s. Retorna fração de sucessos."""
|
| 154 |
+
if not test_texts:
|
| 155 |
+
return 1.0
|
| 156 |
+
ok = 0
|
| 157 |
+
for s in test_texts:
|
| 158 |
+
try:
|
| 159 |
+
ids = self.encode(s, add_special=False)
|
| 160 |
+
decoded = self.decode(ids, skip_special=False)
|
| 161 |
+
if decoded == s:
|
| 162 |
+
ok += 1
|
| 163 |
+
except Exception:
|
| 164 |
+
pass
|
| 165 |
+
return ok / len(test_texts)
|
| 166 |
+
|
| 167 |
+
|
| 168 |
+
__all__ = [
|
| 169 |
+
"BBPETokenizer",
|
| 170 |
+
"BOS_TOKEN", "PAD_TOKEN", "EOS_TOKEN", "UNK_TOKEN",
|
| 171 |
+
"BOS_ID", "PAD_ID", "EOS_ID", "UNK_ID",
|
| 172 |
+
"SPECIAL_TOKENS",
|
| 173 |
+
]
|
cnn_bigru/training/__init__.py
ADDED
|
File without changes
|
cnn_bigru/training/auto_learner.py
ADDED
|
@@ -0,0 +1,301 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""auto_learner.py — Auto-aprendizado: ajuste dinâmico de LR e controle de gradiente.
|
| 2 |
+
|
| 3 |
+
Implementa as seções 9-10 de dados.txt:
|
| 4 |
+
|
| 5 |
+
9. AJUSTE DINÂMICO DAS TAXAS DE APRENDIZADO:
|
| 6 |
+
fator_G = exp(-κ * ||grad_G||) * (1 + v_media) / 2
|
| 7 |
+
novo_eta_G = eta_G * fator_G
|
| 8 |
+
|
| 9 |
+
fator_V = exp(-κ * ||grad_V||) * (1 + L_V_media)^{-1}
|
| 10 |
+
novo_eta_V = eta_V * fator_V
|
| 11 |
+
|
| 12 |
+
10. CONTROLE DE EXPLOSÃO DE GRADIENTE E ESPECTRAL NORM:
|
| 13 |
+
- clip_grad_norm_(params, GRAD_CLIP)
|
| 14 |
+
- Spectral normalization (power iteration) para Conv1D e Linear
|
| 15 |
+
"""
|
| 16 |
+
from __future__ import annotations
|
| 17 |
+
|
| 18 |
+
import logging
|
| 19 |
+
import math
|
| 20 |
+
from dataclasses import dataclass
|
| 21 |
+
from typing import Dict, List, Optional, Tuple
|
| 22 |
+
|
| 23 |
+
import torch
|
| 24 |
+
import torch.nn as nn
|
| 25 |
+
|
| 26 |
+
logger = logging.getLogger(__name__)
|
| 27 |
+
|
| 28 |
+
|
| 29 |
+
@dataclass
|
| 30 |
+
class AutoLearnConfig:
|
| 31 |
+
"""Configuração do auto-aprendizado."""
|
| 32 |
+
kappa_curv: float = 0.01 # κ — fator de ajuste de η com norma do gradiente
|
| 33 |
+
grad_clip: float = 1.0 # τ — clipping de gradiente (norma L2)
|
| 34 |
+
spectral_radius: float = 1.0 # limite para norma espectral
|
| 35 |
+
lr_min: float = 1e-6 # LR mínimo
|
| 36 |
+
lr_max: float = 1e-2 # LR máximo
|
| 37 |
+
initial_lr_G: float = 1e-3 # LR inicial do gerador
|
| 38 |
+
initial_lr_V: float = 1e-3 # LR inicial do verificador
|
| 39 |
+
use_spectral_norm: bool = True
|
| 40 |
+
apply_after_step: bool = True
|
| 41 |
+
# NOVO: peso L2 usado pelo otimizador AdamW (corresponde a L2_REG do dados.txt)
|
| 42 |
+
l2_reg: float = 1e-5
|
| 43 |
+
# NOVO: potências de iteração para spectral norm (1 era muito pouco)
|
| 44 |
+
spectral_n_iters: int = 1
|
| 45 |
+
# NOVO: device padrão para buffers internos
|
| 46 |
+
device: str = "cpu"
|
| 47 |
+
|
| 48 |
+
|
| 49 |
+
class AutoLearner:
|
| 50 |
+
"""Controla ajuste dinâmico de LR, clipping e spectral norm.
|
| 51 |
+
|
| 52 |
+
Não é um nn.Module — é um controlador que atua sobre modelos externos.
|
| 53 |
+
"""
|
| 54 |
+
|
| 55 |
+
def __init__(self, config: AutoLearnConfig):
|
| 56 |
+
self.cfg = config
|
| 57 |
+
self.current_lr_G = config.initial_lr_G
|
| 58 |
+
self.current_lr_V = config.initial_lr_V
|
| 59 |
+
self.history: List[Dict] = []
|
| 60 |
+
|
| 61 |
+
def compute_total_grad_norm(self, model: nn.Module) -> float:
|
| 62 |
+
"""Calcula a norma L2 total dos gradientes do modelo."""
|
| 63 |
+
total = 0.0
|
| 64 |
+
for p in model.parameters():
|
| 65 |
+
if p.grad is not None:
|
| 66 |
+
total += p.grad.data.norm(2).item() ** 2
|
| 67 |
+
return math.sqrt(total)
|
| 68 |
+
|
| 69 |
+
def update_learning_rates(
|
| 70 |
+
self,
|
| 71 |
+
grad_G_norm: float,
|
| 72 |
+
grad_V_norm: float,
|
| 73 |
+
v_mean: float,
|
| 74 |
+
L_V_mean: float,
|
| 75 |
+
) -> Tuple[float, float]:
|
| 76 |
+
"""Atualiza dinamicamente as taxas de aprendizado.
|
| 77 |
+
|
| 78 |
+
Args:
|
| 79 |
+
grad_G_norm: norma L2 do gradiente do gerador
|
| 80 |
+
grad_V_norm: norma L2 do gradiente do verificador
|
| 81 |
+
v_mean: média das probabilidades do verificador
|
| 82 |
+
L_V_mean: perda média do verificador
|
| 83 |
+
|
| 84 |
+
Returns:
|
| 85 |
+
(novo_eta_G, novo_eta_V)
|
| 86 |
+
"""
|
| 87 |
+
# Fator do gerador: exp(-κ * ||grad||) * (1 + v_media) / 2
|
| 88 |
+
fator_G = math.exp(-self.cfg.kappa_curv * grad_G_norm) * ((1.0 + v_mean) / 2.0)
|
| 89 |
+
novo_eta_G = self.current_lr_G * fator_G
|
| 90 |
+
|
| 91 |
+
# Fator do verificador: exp(-κ * ||grad||) * (1 + L_V)^{-1}
|
| 92 |
+
# Evita divisão por zero
|
| 93 |
+
fator_V = math.exp(-self.cfg.kappa_curv * grad_V_norm) / max(1.0 + L_V_mean, 1e-6)
|
| 94 |
+
novo_eta_V = self.current_lr_V * fator_V
|
| 95 |
+
|
| 96 |
+
# Clamp para evitar valores absurdos
|
| 97 |
+
novo_eta_G = max(self.cfg.lr_min, min(self.cfg.lr_max, novo_eta_G))
|
| 98 |
+
novo_eta_V = max(self.cfg.lr_min, min(self.cfg.lr_max, novo_eta_V))
|
| 99 |
+
|
| 100 |
+
self.current_lr_G = novo_eta_G
|
| 101 |
+
self.current_lr_V = novo_eta_V
|
| 102 |
+
|
| 103 |
+
self.history.append({
|
| 104 |
+
"lr_G": novo_eta_G,
|
| 105 |
+
"lr_V": novo_eta_V,
|
| 106 |
+
"grad_G_norm": grad_G_norm,
|
| 107 |
+
"grad_V_norm": grad_V_norm,
|
| 108 |
+
"v_mean": v_mean,
|
| 109 |
+
"L_V_mean": L_V_mean,
|
| 110 |
+
})
|
| 111 |
+
|
| 112 |
+
return novo_eta_G, novo_eta_V
|
| 113 |
+
|
| 114 |
+
def apply_clipping(self, *models: nn.Module) -> None:
|
| 115 |
+
"""Aplica gradient clipping global nos modelos."""
|
| 116 |
+
for model in models:
|
| 117 |
+
if model is not None:
|
| 118 |
+
torch.nn.utils.clip_grad_norm_(model.parameters(), self.cfg.grad_clip)
|
| 119 |
+
|
| 120 |
+
def apply_spectral_norm(self, model: nn.Module) -> int:
|
| 121 |
+
"""Aplica normalização espectral via power iteration em Conv1d/Linear.
|
| 122 |
+
|
| 123 |
+
Após cada step, força ||W||_2 <= spectral_radius.
|
| 124 |
+
|
| 125 |
+
CORREÇÃO: O bug original armazenava `module._spec_u` como atributo
|
| 126 |
+
comum (não buffer), que NÃO migra com `.to(device)`. Agora usamos
|
| 127 |
+
`register_buffer` na primeira chamada para garantir movimentação
|
| 128 |
+
corta com o device do modelo.
|
| 129 |
+
"""
|
| 130 |
+
if not self.cfg.use_spectral_norm:
|
| 131 |
+
return 0
|
| 132 |
+
n_normalized = 0
|
| 133 |
+
for module in model.modules():
|
| 134 |
+
if isinstance(module, (nn.Linear, nn.Conv1d, nn.Conv2d)):
|
| 135 |
+
if module.weight is not None and module.weight.dim() >= 2:
|
| 136 |
+
try:
|
| 137 |
+
with torch.no_grad():
|
| 138 |
+
# Power iteration (n_iters iterações)
|
| 139 |
+
W = module.weight
|
| 140 |
+
# Reshape para 2D se necessário
|
| 141 |
+
W_flat = W.reshape(W.size(0), -1)
|
| 142 |
+
# u, v aleatórios (ou reutiliza se já existe)
|
| 143 |
+
buffer_name = "_spec_u"
|
| 144 |
+
if not hasattr(module, buffer_name):
|
| 145 |
+
u = torch.randn(W_flat.size(0), 1, device=W.device, dtype=W.dtype)
|
| 146 |
+
u = u / u.norm().clamp(min=1e-8)
|
| 147 |
+
# Registrar como buffer (move com .to(device))
|
| 148 |
+
module.register_buffer(buffer_name, u, persistent=False)
|
| 149 |
+
else:
|
| 150 |
+
u = getattr(module, buffer_name)
|
| 151 |
+
# Garantir que está no device correto
|
| 152 |
+
if u.device != W.device:
|
| 153 |
+
u = u.to(W.device)
|
| 154 |
+
setattr(module, buffer_name, u)
|
| 155 |
+
# n_iters iterações de power iteration
|
| 156 |
+
for _ in range(self.cfg.spectral_n_iters):
|
| 157 |
+
v = torch.matmul(W_flat.T, u)
|
| 158 |
+
v = v / v.norm().clamp(min=1e-8)
|
| 159 |
+
u = torch.matmul(W_flat, v)
|
| 160 |
+
u = u / u.norm().clamp(min=1e-8)
|
| 161 |
+
# Atualizar buffer (in-place para manter no mesmo device)
|
| 162 |
+
u_new = u.detach()
|
| 163 |
+
getattr(module, buffer_name).copy_(u_new)
|
| 164 |
+
sigma = (u.T @ W_flat @ v).item()
|
| 165 |
+
if sigma > self.cfg.spectral_radius and sigma > 0:
|
| 166 |
+
scale = self.cfg.spectral_radius / sigma
|
| 167 |
+
module.weight.data.mul_(scale)
|
| 168 |
+
n_normalized += 1
|
| 169 |
+
except Exception as e:
|
| 170 |
+
logger.debug("Spectral norm falhou em %s: %s", type(module).__name__, e)
|
| 171 |
+
return n_normalized
|
| 172 |
+
|
| 173 |
+
def set_optimizer_lr(self, optimizer: torch.optim.Optimizer, lr: float) -> None:
|
| 174 |
+
"""Atualiza a taxa de aprendizado de um otimizador."""
|
| 175 |
+
for pg in optimizer.param_groups:
|
| 176 |
+
pg["lr"] = lr
|
| 177 |
+
|
| 178 |
+
|
| 179 |
+
def orthogonal_init_model(model: nn.Module) -> int:
|
| 180 |
+
"""Aplica inicialização ortogonal em todas as camadas Linear/Conv/GRUCell."""
|
| 181 |
+
n = 0
|
| 182 |
+
for module in model.modules():
|
| 183 |
+
if isinstance(module, (nn.Linear, nn.Conv1d, nn.Conv2d)):
|
| 184 |
+
try:
|
| 185 |
+
nn.init.orthogonal_(module.weight)
|
| 186 |
+
if module.bias is not None:
|
| 187 |
+
nn.init.zeros_(module.bias)
|
| 188 |
+
n += 1
|
| 189 |
+
except Exception:
|
| 190 |
+
pass
|
| 191 |
+
elif isinstance(module, (nn.GRUCell, nn.GRU, nn.LSTM)):
|
| 192 |
+
for name, p in module.named_parameters():
|
| 193 |
+
if "weight" in name and p.dim() >= 2:
|
| 194 |
+
try:
|
| 195 |
+
nn.init.orthogonal_(p)
|
| 196 |
+
n += 1
|
| 197 |
+
except Exception:
|
| 198 |
+
pass
|
| 199 |
+
elif "bias" in name:
|
| 200 |
+
nn.init.zeros_(p)
|
| 201 |
+
return n
|
| 202 |
+
|
| 203 |
+
|
| 204 |
+
def estimate_curvature(
|
| 205 |
+
model: nn.Module,
|
| 206 |
+
loss_fn,
|
| 207 |
+
*args,
|
| 208 |
+
eps: float = 1e-3,
|
| 209 |
+
n_samples: int = 5,
|
| 210 |
+
) -> torch.Tensor:
|
| 211 |
+
"""Estima curvatura da perda via diferenças finitas.
|
| 212 |
+
|
| 213 |
+
Para uma amostra de parâmetros, calcula:
|
| 214 |
+
curv ≈ (L(θ + ε) - 2L(θ) + L(θ - ε)) / ε²
|
| 215 |
+
|
| 216 |
+
Retorna a média das estimativas.
|
| 217 |
+
|
| 218 |
+
CORREÇÃO: O bug original mutava `flat[i]` in-place sem try/finally,
|
| 219 |
+
corrompendo o modelo se uma exceção ocorresse entre `+eps` e o restore.
|
| 220 |
+
Agora usamos try/finally para garantir restauração sempre.
|
| 221 |
+
"""
|
| 222 |
+
curvatures = []
|
| 223 |
+
params = [p for p in model.parameters() if p.requires_grad and p.dim() >= 2]
|
| 224 |
+
if not params:
|
| 225 |
+
# Device-safe zero
|
| 226 |
+
try:
|
| 227 |
+
device = next(model.parameters()).device
|
| 228 |
+
except StopIteration:
|
| 229 |
+
device = torch.device("cpu")
|
| 230 |
+
return torch.zeros((), device=device)
|
| 231 |
+
|
| 232 |
+
# Subamostra
|
| 233 |
+
n_samples = min(n_samples, len(params))
|
| 234 |
+
indices = torch.randperm(len(params))[:n_samples]
|
| 235 |
+
|
| 236 |
+
for idx in indices:
|
| 237 |
+
p = params[idx]
|
| 238 |
+
# Pega um elemento aleatório
|
| 239 |
+
flat = p.data.view(-1)
|
| 240 |
+
if flat.numel() == 0:
|
| 241 |
+
continue
|
| 242 |
+
i = torch.randint(0, flat.numel(), (1,)).item()
|
| 243 |
+
orig = flat[i].item()
|
| 244 |
+
|
| 245 |
+
# TRY/FINALLY para garantir restauração sempre
|
| 246 |
+
try:
|
| 247 |
+
with torch.no_grad():
|
| 248 |
+
# L(θ)
|
| 249 |
+
try:
|
| 250 |
+
loss_0 = loss_fn(*args)
|
| 251 |
+
if isinstance(loss_0, dict):
|
| 252 |
+
loss_0 = loss_0["total"]
|
| 253 |
+
loss_0 = float(loss_0)
|
| 254 |
+
except Exception as e:
|
| 255 |
+
logger.debug("curvature L(θ) falhou: %s", e)
|
| 256 |
+
continue
|
| 257 |
+
|
| 258 |
+
# L(θ + ε)
|
| 259 |
+
flat[i] = orig + eps
|
| 260 |
+
try:
|
| 261 |
+
loss_p = loss_fn(*args)
|
| 262 |
+
if isinstance(loss_p, dict):
|
| 263 |
+
loss_p = loss_p["total"]
|
| 264 |
+
loss_p = float(loss_p)
|
| 265 |
+
except Exception as e:
|
| 266 |
+
logger.debug("curvature L(θ+ε) falhou: %s", e)
|
| 267 |
+
continue
|
| 268 |
+
|
| 269 |
+
# L(θ - ε)
|
| 270 |
+
flat[i] = orig - eps
|
| 271 |
+
try:
|
| 272 |
+
loss_m = loss_fn(*args)
|
| 273 |
+
if isinstance(loss_m, dict):
|
| 274 |
+
loss_m = loss_m["total"]
|
| 275 |
+
loss_m = float(loss_m)
|
| 276 |
+
except Exception as e:
|
| 277 |
+
logger.debug("curvature L(θ-ε) falhou: %s", e)
|
| 278 |
+
continue
|
| 279 |
+
|
| 280 |
+
curv = (loss_p - 2 * loss_0 + loss_m) / (eps ** 2)
|
| 281 |
+
curvatures.append(curv)
|
| 282 |
+
finally:
|
| 283 |
+
# SEMPRE restaurar o valor original
|
| 284 |
+
flat[i] = orig
|
| 285 |
+
|
| 286 |
+
if not curvatures:
|
| 287 |
+
try:
|
| 288 |
+
device = next(model.parameters()).device
|
| 289 |
+
except StopIteration:
|
| 290 |
+
device = torch.device("cpu")
|
| 291 |
+
return torch.zeros((), device=device)
|
| 292 |
+
|
| 293 |
+
return torch.tensor(sum(curvatures) / len(curvatures))
|
| 294 |
+
|
| 295 |
+
|
| 296 |
+
__all__ = [
|
| 297 |
+
"AutoLearnConfig",
|
| 298 |
+
"AutoLearner",
|
| 299 |
+
"orthogonal_init_model",
|
| 300 |
+
"estimate_curvature",
|
| 301 |
+
]
|
cnn_bigru/training/hypothesis_controller.py
ADDED
|
@@ -0,0 +1,290 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""hypothesis_controller.py — Mecanismo de hipóteses para autoaprendizado.
|
| 2 |
+
|
| 3 |
+
Implementa o "acionamento de hipóteses quando ocorre punições no treinamento"
|
| 4 |
+
descrito pelo usuário. Quando o verificador pune um passo (v < threshold),
|
| 5 |
+
o sistema ativa camadas de hipótese que tentam N combinações alternativas
|
| 6 |
+
de sinergias entre as redes, buscando a configuração de menor perda.
|
| 7 |
+
|
| 8 |
+
Matemática:
|
| 9 |
+
Quando v_t < THRESHOLD_ERR:
|
| 10 |
+
hipótese_i = f_alternativa_i(input, params_perturbados_i)
|
| 11 |
+
loss_i = compute_loss(hipótese_i)
|
| 12 |
+
melhor = argmin(loss_i)
|
| 13 |
+
params <- params + lr * grad(melhor)
|
| 14 |
+
|
| 15 |
+
As "camadas de hipótese" são camadas lineares paralelas (hipóteses)
|
| 16 |
+
ativadas condicionalmente. A ativação é diferenciável via Gumbel-Softmax.
|
| 17 |
+
"""
|
| 18 |
+
from __future__ import annotations
|
| 19 |
+
|
| 20 |
+
import logging
|
| 21 |
+
import math
|
| 22 |
+
import random
|
| 23 |
+
from dataclasses import dataclass
|
| 24 |
+
from typing import Dict, List, Optional, Tuple
|
| 25 |
+
|
| 26 |
+
import torch
|
| 27 |
+
import torch.nn as nn
|
| 28 |
+
import torch.nn.functional as F
|
| 29 |
+
|
| 30 |
+
logger = logging.getLogger(__name__)
|
| 31 |
+
|
| 32 |
+
|
| 33 |
+
@dataclass
|
| 34 |
+
class HypothesisConfig:
|
| 35 |
+
"""Configuração do mecanismo de hipóteses."""
|
| 36 |
+
n_hypotheses: int = 4 # N tentativas de acerto
|
| 37 |
+
threshold_err: float = 0.5 # abaixo deste valor, ativa hipóteses
|
| 38 |
+
gumbel_temp: float = 1.0 # temperatura do Gumbel-Softmax
|
| 39 |
+
gumbel_hard: bool = False # se True, usa argmax (não diferenciável)
|
| 40 |
+
perturbation_std: float = 0.05 # desvio padrão da perturbação
|
| 41 |
+
|
| 42 |
+
|
| 43 |
+
class HypothesisLayer(nn.Module):
|
| 44 |
+
"""Camada de hipótese: N transformações lineares paralelas.
|
| 45 |
+
|
| 46 |
+
Para cada amostra, calcula N hipóteses e seleciona a melhor via
|
| 47 |
+
Gumbel-Softmax (diferenciável) ou argmax (não diferenciável).
|
| 48 |
+
|
| 49 |
+
A "seleção" é feita baseada em um score externo (e.g. negativo da perda).
|
| 50 |
+
"""
|
| 51 |
+
|
| 52 |
+
def __init__(
|
| 53 |
+
self,
|
| 54 |
+
in_dim: int,
|
| 55 |
+
out_dim: int,
|
| 56 |
+
n_hypotheses: int = 4,
|
| 57 |
+
activation: str = "relu",
|
| 58 |
+
):
|
| 59 |
+
super().__init__()
|
| 60 |
+
self.n_hypotheses = n_hypotheses
|
| 61 |
+
self.in_dim = in_dim
|
| 62 |
+
self.out_dim = out_dim
|
| 63 |
+
|
| 64 |
+
# N transformações lineares paralelas
|
| 65 |
+
self.hypotheses = nn.ModuleList([
|
| 66 |
+
nn.Linear(in_dim, out_dim) for _ in range(n_hypotheses)
|
| 67 |
+
])
|
| 68 |
+
# Inicialização diferente para cada hipótese (diversidade)
|
| 69 |
+
for i, h in enumerate(self.hypotheses):
|
| 70 |
+
nn.init.xavier_uniform_(h.weight, gain=0.5 + 0.2 * i)
|
| 71 |
+
nn.init.zeros_(h.bias)
|
| 72 |
+
|
| 73 |
+
# Scoreador: aprende a pontuar cada hipótese dado o input
|
| 74 |
+
self.scorer = nn.Linear(in_dim, n_hypotheses)
|
| 75 |
+
|
| 76 |
+
self.activation = activation
|
| 77 |
+
|
| 78 |
+
def _activate(self, x: torch.Tensor) -> torch.Tensor:
|
| 79 |
+
if self.activation == "relu":
|
| 80 |
+
return F.relu(x)
|
| 81 |
+
elif self.activation == "gelu":
|
| 82 |
+
return F.gelu(x)
|
| 83 |
+
elif self.activation == "tanh":
|
| 84 |
+
return torch.tanh(x)
|
| 85 |
+
return x
|
| 86 |
+
|
| 87 |
+
def forward(
|
| 88 |
+
self,
|
| 89 |
+
x: torch.Tensor,
|
| 90 |
+
gumbel_temp: float = 1.0,
|
| 91 |
+
gumbel_hard: bool = False,
|
| 92 |
+
) -> Tuple[torch.Tensor, torch.Tensor]:
|
| 93 |
+
"""
|
| 94 |
+
Args:
|
| 95 |
+
x: [B, in_dim]
|
| 96 |
+
gumbel_temp: temperatura do Gumbel-Softmax
|
| 97 |
+
gumbel_hard: se True, usa argmax
|
| 98 |
+
|
| 99 |
+
Returns:
|
| 100 |
+
output: [B, out_dim]
|
| 101 |
+
weights: [B, n_hypotheses]
|
| 102 |
+
"""
|
| 103 |
+
B = x.size(0)
|
| 104 |
+
|
| 105 |
+
# Calcula todas as N hipóteses: [B, N, out_dim]
|
| 106 |
+
hyps = torch.stack([self._activate(h(x)) for h in self.hypotheses], dim=1)
|
| 107 |
+
|
| 108 |
+
# Scoreia cada hipótese
|
| 109 |
+
scores = self.scorer(x) # [B, N]
|
| 110 |
+
|
| 111 |
+
# Gumbel-Softmax para seleção diferenciável
|
| 112 |
+
if gumbel_hard:
|
| 113 |
+
# Hard: argmax one-hot
|
| 114 |
+
weights = F.one_hot(scores.argmax(dim=-1), num_classes=self.n_hypotheses).float()
|
| 115 |
+
else:
|
| 116 |
+
weights = F.gumbel_softmax(scores, tau=gumbel_temp, hard=False)
|
| 117 |
+
|
| 118 |
+
# Combina: weighted sum
|
| 119 |
+
output = (hyps * weights.unsqueeze(-1)).sum(dim=1) # [B, out_dim]
|
| 120 |
+
|
| 121 |
+
return output, weights
|
| 122 |
+
|
| 123 |
+
|
| 124 |
+
class HypothesisController(nn.Module):
|
| 125 |
+
"""Controlador de hipóteses: ativa camadas quando há punição.
|
| 126 |
+
|
| 127 |
+
Recebe:
|
| 128 |
+
- features do passo atual
|
| 129 |
+
- valor v do verificador (probabilidade de correção)
|
| 130 |
+
- loss atual
|
| 131 |
+
|
| 132 |
+
Se v < threshold, ativa as hipóteses e tenta N combinações.
|
| 133 |
+
Caso contrário, passa direto (identity).
|
| 134 |
+
|
| 135 |
+
CORREÇÃO: O bug original era `x_proj = combined` que tornava
|
| 136 |
+
`output = gate * combined + (1 - gate) * combined = combined` (sempre).
|
| 137 |
+
Agora `x_proj` é uma projeção LINEAR REAL do input x (não o combined),
|
| 138 |
+
permitindo que o gate efetivamente escolha entre:
|
| 139 |
+
- gate=1: usar saída das hipóteses (combined)
|
| 140 |
+
- gate=0: usar projeção do input original (x_proj)
|
| 141 |
+
"""
|
| 142 |
+
|
| 143 |
+
def __init__(
|
| 144 |
+
self,
|
| 145 |
+
feat_dim: int,
|
| 146 |
+
out_dim: int,
|
| 147 |
+
config: HypothesisConfig,
|
| 148 |
+
):
|
| 149 |
+
super().__init__()
|
| 150 |
+
self.cfg = config
|
| 151 |
+
self.feat_dim = feat_dim
|
| 152 |
+
self.out_dim = out_dim
|
| 153 |
+
|
| 154 |
+
# Camadas de hipótese paralelas
|
| 155 |
+
self.hyp_layer = HypothesisLayer(
|
| 156 |
+
in_dim=feat_dim,
|
| 157 |
+
out_dim=out_dim,
|
| 158 |
+
n_hypotheses=config.n_hypotheses,
|
| 159 |
+
activation="gelu",
|
| 160 |
+
)
|
| 161 |
+
|
| 162 |
+
# Porta de ativação (gating) baseada em v
|
| 163 |
+
# Se v alto, mantém identidade; se v baixo, ativa hipóteses
|
| 164 |
+
self.gate = nn.Linear(feat_dim + 1, 1) # +1 para o valor v
|
| 165 |
+
|
| 166 |
+
# Projeto para combinar hipótese + input
|
| 167 |
+
self.combine = nn.Linear(feat_dim + out_dim, out_dim)
|
| 168 |
+
|
| 169 |
+
# NOVO: projeção separada do input x para out_dim
|
| 170 |
+
# Esta é a "via identity" — usada quando gate=0 (v alto, sem punição)
|
| 171 |
+
self.x_proj = nn.Linear(feat_dim, out_dim)
|
| 172 |
+
nn.init.xavier_uniform_(self.x_proj.weight, gain=0.5)
|
| 173 |
+
nn.init.zeros_(self.x_proj.bias)
|
| 174 |
+
|
| 175 |
+
def forward(
|
| 176 |
+
self,
|
| 177 |
+
x: torch.Tensor,
|
| 178 |
+
v: torch.Tensor,
|
| 179 |
+
) -> Dict[str, torch.Tensor]:
|
| 180 |
+
"""
|
| 181 |
+
Args:
|
| 182 |
+
x: [B, feat_dim]
|
| 183 |
+
v: [B, 1] probabilidade do verificador (0 = erro, 1 = correto)
|
| 184 |
+
|
| 185 |
+
Returns:
|
| 186 |
+
dict com:
|
| 187 |
+
output: [B, out_dim] features modificadas
|
| 188 |
+
gate: [B, 1] valor da porta (0 = identity, 1 = hipóteses)
|
| 189 |
+
weights: [B, N] pesos das hipóteses
|
| 190 |
+
activated: bool se as hipóteses foram ativadas (em qualquer amostra)
|
| 191 |
+
"""
|
| 192 |
+
# Porta: gate = Sigmoid(linear(concat(x, 1-v)))
|
| 193 |
+
# v baixo → 1-v alto → gate alto (ativa hipóteses)
|
| 194 |
+
gate_input = torch.cat((x, 1.0 - v), dim=-1) # [B, feat_dim + 1]
|
| 195 |
+
gate = torch.sigmoid(self.gate(gate_input)) # [B, 1]
|
| 196 |
+
|
| 197 |
+
# Hipóteses (N tentativas paralelas)
|
| 198 |
+
hyp_out, weights = self.hyp_layer(
|
| 199 |
+
x, gumbel_temp=self.cfg.gumbel_temp, gumbel_hard=self.cfg.gumbel_hard
|
| 200 |
+
)
|
| 201 |
+
|
| 202 |
+
# Via das hipóteses: combina input original com saída das hipóteses
|
| 203 |
+
combined = self.combine(torch.cat((x, hyp_out), dim=-1)) # [B, out_dim]
|
| 204 |
+
|
| 205 |
+
# Via identidade: projeção direta do input (sem usar hipóteses)
|
| 206 |
+
x_projected = self.x_proj(x) # [B, out_dim]
|
| 207 |
+
|
| 208 |
+
# Interpolação: gate=1 → hipóteses; gate=0 → identidade
|
| 209 |
+
output = gate * combined + (1.0 - gate) * x_projected
|
| 210 |
+
|
| 211 |
+
# Verifica se alguma amostra foi punida
|
| 212 |
+
activated = bool((v.squeeze(-1) < self.cfg.threshold_err).any().item())
|
| 213 |
+
|
| 214 |
+
return {
|
| 215 |
+
"output": output,
|
| 216 |
+
"gate": gate,
|
| 217 |
+
"weights": weights,
|
| 218 |
+
"activated": activated,
|
| 219 |
+
}
|
| 220 |
+
|
| 221 |
+
|
| 222 |
+
class SynergySearcher:
|
| 223 |
+
"""Busca sinérgica: tenta N combinações de hiperparâmetros e escolhe a melhor.
|
| 224 |
+
|
| 225 |
+
Dado um conjunto de "alavancas" (e.g. pesos das perdas, taxas de aprendizado),
|
| 226 |
+
explora N configurações e seleciona aquela que produz menor perda em um
|
| 227 |
+
mini-batch de validação.
|
| 228 |
+
|
| 229 |
+
Não é diferenciável — é uma busca heurística no espaço de configurações.
|
| 230 |
+
"""
|
| 231 |
+
|
| 232 |
+
def __init__(
|
| 233 |
+
self,
|
| 234 |
+
n_attempts: int = 4,
|
| 235 |
+
seed: int = 42,
|
| 236 |
+
):
|
| 237 |
+
self.n_attempts = n_attempts
|
| 238 |
+
self.rng = random.Random(seed)
|
| 239 |
+
self.history: List[Dict] = []
|
| 240 |
+
|
| 241 |
+
def sample_config(self, base: Dict) -> Dict:
|
| 242 |
+
"""Amostra uma configuração perturbando a base."""
|
| 243 |
+
cfg = dict(base)
|
| 244 |
+
# Perturbação: multiplica pesos por fator aleatório próximo de 1
|
| 245 |
+
for key in ["alpha", "beta", "gamma_loss", "delta", "lambda_penal", "mu_exp_penal"]:
|
| 246 |
+
if key in cfg:
|
| 247 |
+
factor = 1.0 + self.rng.uniform(-0.2, 0.2)
|
| 248 |
+
cfg[key] = max(1e-6, cfg[key] * factor)
|
| 249 |
+
return cfg
|
| 250 |
+
|
| 251 |
+
def search(
|
| 252 |
+
self,
|
| 253 |
+
base_config: Dict,
|
| 254 |
+
eval_fn,
|
| 255 |
+
) -> Dict:
|
| 256 |
+
"""Executa N tentativas e retorna a melhor configuração.
|
| 257 |
+
|
| 258 |
+
Args:
|
| 259 |
+
base_config: configuração base (dict de hiperparâmetros)
|
| 260 |
+
eval_fn: função cfg -> loss (callable)
|
| 261 |
+
|
| 262 |
+
Returns:
|
| 263 |
+
melhor configuração encontrada
|
| 264 |
+
"""
|
| 265 |
+
best_config = dict(base_config)
|
| 266 |
+
best_loss = float("inf")
|
| 267 |
+
|
| 268 |
+
for attempt in range(self.n_attempts):
|
| 269 |
+
try:
|
| 270 |
+
cfg = self.sample_config(base_config) if attempt > 0 else dict(base_config)
|
| 271 |
+
loss = float(eval_fn(cfg))
|
| 272 |
+
self.history.append({"attempt": attempt, "config": cfg, "loss": loss})
|
| 273 |
+
logger.info("SynergySearch attempt %d: loss=%.4f", attempt, loss)
|
| 274 |
+
if loss < best_loss:
|
| 275 |
+
best_loss = loss
|
| 276 |
+
best_config = cfg
|
| 277 |
+
except Exception as e:
|
| 278 |
+
logger.warning("SynergySearch attempt %d falhou: %s", attempt, e)
|
| 279 |
+
continue
|
| 280 |
+
|
| 281 |
+
logger.info("SynergySearch melhor loss=%.4f", best_loss)
|
| 282 |
+
return best_config
|
| 283 |
+
|
| 284 |
+
|
| 285 |
+
__all__ = [
|
| 286 |
+
"HypothesisConfig",
|
| 287 |
+
"HypothesisLayer",
|
| 288 |
+
"HypothesisController",
|
| 289 |
+
"SynergySearcher",
|
| 290 |
+
]
|
cnn_bigru/training/trainer.py
ADDED
|
@@ -0,0 +1,672 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""trainer.py — Loop de treinamento cooperativo com auto-aprendizado.
|
| 2 |
+
|
| 3 |
+
Implementa a seção 11 de dados.txt:
|
| 4 |
+
|
| 5 |
+
PARA época = 1 ATÉ NUM_EPOCAS:
|
| 6 |
+
PARA batch EM DataLoader:
|
| 7 |
+
// Forward Gerador
|
| 8 |
+
logits_prova, enc_out, _ = gerador.forward(axiomas, conjectura, prova_alvo)
|
| 9 |
+
|
| 10 |
+
// Calcular perdas e avaliações do verificador
|
| 11 |
+
L_total, L_G, L_V, L_AH, lista_v = calcular_perdas(...)
|
| 12 |
+
|
| 13 |
+
// Backward
|
| 14 |
+
zerar_gradientes
|
| 15 |
+
L_total.backward()
|
| 16 |
+
|
| 17 |
+
// Controle de gradiente (clip + spectral norm)
|
| 18 |
+
aplicar_controle_gradiente(gerador)
|
| 19 |
+
aplicar_controle_gradiente(verificador)
|
| 20 |
+
|
| 21 |
+
// Ajuste dinâmico de LR baseado em normas de gradiente e v_media
|
| 22 |
+
eta_G, eta_V = atualizar_taxas(...)
|
| 23 |
+
|
| 24 |
+
// Passo de otimização
|
| 25 |
+
otimizador_G.step()
|
| 26 |
+
otimizador_V.step()
|
| 27 |
+
|
| 28 |
+
// Normalização espectral após step
|
| 29 |
+
aplicar_espectral_norm(gerador)
|
| 30 |
+
aplicar_espectral_norm(verificador)
|
| 31 |
+
|
| 32 |
+
Adicionalmente (APRIMORAMENTOS v2.0):
|
| 33 |
+
- Mecanismo de hipóteses ativado quando v < threshold (com gate funcional)
|
| 34 |
+
- N tentativas de sinergia (SynergySearcher) no início de cada época
|
| 35 |
+
- Meta-controlador que ajusta pesos das perdas baseado na variância do gradiente
|
| 36 |
+
- EWC (Elastic Weight Consolidation) integrado à perda total
|
| 37 |
+
- Context Window aplicado na inferência (não no treino, que usa teacher forcing)
|
| 38 |
+
- Verificador REALMENTE chamado (não proxy) — correção de bug crítico
|
| 39 |
+
- Hipótese controller output REALMENTE usado na forward — correção de bug crítico
|
| 40 |
+
"""
|
| 41 |
+
from __future__ import annotations
|
| 42 |
+
|
| 43 |
+
import logging
|
| 44 |
+
import math
|
| 45 |
+
import os
|
| 46 |
+
import time
|
| 47 |
+
from dataclasses import dataclass, field
|
| 48 |
+
from typing import Any, Dict, List, Optional, Tuple
|
| 49 |
+
|
| 50 |
+
import torch
|
| 51 |
+
import torch.nn as nn
|
| 52 |
+
import torch.nn.functional as F
|
| 53 |
+
from torch.utils.data import DataLoader
|
| 54 |
+
|
| 55 |
+
from ..losses.losses import LossConfig, MultiLoss
|
| 56 |
+
from ..models.generator_verifier import (
|
| 57 |
+
GeneratorCNNBiGRU,
|
| 58 |
+
VerifierCNNBiGRU,
|
| 59 |
+
AntiHallucinationLayer,
|
| 60 |
+
)
|
| 61 |
+
from ..models.multimodal_model import MultimodalCNNBiGRU
|
| 62 |
+
from ..utils.ewc import EWCConfig, EWCState
|
| 63 |
+
from ..utils.memory_optimizer import MemoryOptimizer
|
| 64 |
+
from .auto_learner import AutoLearnConfig, AutoLearner, estimate_curvature, orthogonal_init_model
|
| 65 |
+
from .hypothesis_controller import HypothesisConfig, HypothesisController, SynergySearcher
|
| 66 |
+
|
| 67 |
+
logger = logging.getLogger(__name__)
|
| 68 |
+
|
| 69 |
+
|
| 70 |
+
@dataclass
|
| 71 |
+
class TrainerConfig:
|
| 72 |
+
"""Configuração do treinador."""
|
| 73 |
+
# Épocas e batches
|
| 74 |
+
num_epochs: int = 2
|
| 75 |
+
max_batches_per_epoch: int = 10 # limita batches para teste de 50 amostras
|
| 76 |
+
batch_size: int = 8
|
| 77 |
+
|
| 78 |
+
# Synergy search
|
| 79 |
+
n_synergy_attempts: int = 3 # N tentativas de acerto do treinamento
|
| 80 |
+
use_synergy_search: bool = True
|
| 81 |
+
|
| 82 |
+
# Hypothesis
|
| 83 |
+
use_hypotheses: bool = True
|
| 84 |
+
n_hypotheses: int = 4
|
| 85 |
+
|
| 86 |
+
# Loss config
|
| 87 |
+
loss_config: LossConfig = field(default_factory=LossConfig)
|
| 88 |
+
|
| 89 |
+
# Auto-learn config
|
| 90 |
+
auto_config: AutoLearnConfig = field(default_factory=AutoLearnConfig)
|
| 91 |
+
|
| 92 |
+
# EWC config (NOVO v2.0)
|
| 93 |
+
ewc_config: Optional[EWCConfig] = None # None = EWC desativado
|
| 94 |
+
|
| 95 |
+
# Device
|
| 96 |
+
device: str = "cpu"
|
| 97 |
+
|
| 98 |
+
# Logging
|
| 99 |
+
log_every: int = 1
|
| 100 |
+
|
| 101 |
+
# Habilitar/desabilitar componentes
|
| 102 |
+
use_verifier_real: bool = True # usar verificador REAL (não proxy)
|
| 103 |
+
use_generator: bool = True # usar gerador para logits de LM
|
| 104 |
+
use_hypothesis_output: bool = True # usar saída da hipótese (não só flag)
|
| 105 |
+
|
| 106 |
+
|
| 107 |
+
class CooperativeTrainer:
|
| 108 |
+
"""Treina o modelo multimodal CNN-BiGRU cooperativo.
|
| 109 |
+
|
| 110 |
+
Args:
|
| 111 |
+
model: MultimodalCNNBiGRU
|
| 112 |
+
tokenizer: BBPETokenizer
|
| 113 |
+
config: TrainerConfig
|
| 114 |
+
generator: GeneratorCNNBiGRU (opcional, para modo gerar+verificar)
|
| 115 |
+
verifier: VerifierCNNBiGRU (opcional)
|
| 116 |
+
anti_hallucination: AntiHallucinationLayer (opcional)
|
| 117 |
+
ewc_state: EWCState (opcional, para aprendizado contínuo)
|
| 118 |
+
"""
|
| 119 |
+
|
| 120 |
+
def __init__(
|
| 121 |
+
self,
|
| 122 |
+
model: MultimodalCNNBiGRU,
|
| 123 |
+
tokenizer,
|
| 124 |
+
config: TrainerConfig,
|
| 125 |
+
generator: Optional[GeneratorCNNBiGRU] = None,
|
| 126 |
+
verifier: Optional[VerifierCNNBiGRU] = None,
|
| 127 |
+
anti_hallucination: Optional[AntiHallucinationLayer] = None,
|
| 128 |
+
ewc_state: Optional[EWCState] = None,
|
| 129 |
+
):
|
| 130 |
+
self.model = model.to(config.device)
|
| 131 |
+
self.tokenizer = tokenizer
|
| 132 |
+
self.cfg = config
|
| 133 |
+
self.device = config.device
|
| 134 |
+
|
| 135 |
+
# Componentes opcionais
|
| 136 |
+
self.generator = generator.to(config.device) if generator else None
|
| 137 |
+
self.verifier = verifier.to(config.device) if verifier else None
|
| 138 |
+
self.anti_hall = (
|
| 139 |
+
anti_hallucination.to(config.device) if anti_hallucination else None
|
| 140 |
+
)
|
| 141 |
+
|
| 142 |
+
# EWC state (aprendizado contínuo) — NOVO v2.0
|
| 143 |
+
self.ewc_state = ewc_state
|
| 144 |
+
|
| 145 |
+
# Loss
|
| 146 |
+
self.loss_fn = MultiLoss(config.loss_config, pad_idx=tokenizer.pad_id).to(config.device)
|
| 147 |
+
|
| 148 |
+
# Auto-learner
|
| 149 |
+
self.auto = AutoLearner(config.auto_config)
|
| 150 |
+
|
| 151 |
+
# Synergy searcher
|
| 152 |
+
self.synergy = SynergySearcher(
|
| 153 |
+
n_attempts=config.n_synergy_attempts,
|
| 154 |
+
seed=42,
|
| 155 |
+
)
|
| 156 |
+
|
| 157 |
+
# Otimizadores
|
| 158 |
+
self.opt_model = torch.optim.AdamW(
|
| 159 |
+
self.model.parameters(), lr=config.auto_config.initial_lr_G,
|
| 160 |
+
weight_decay=config.auto_config.l2_reg # sempre definido agora (v2.0)
|
| 161 |
+
)
|
| 162 |
+
self.opt_gen = (
|
| 163 |
+
torch.optim.AdamW(
|
| 164 |
+
self.generator.parameters(),
|
| 165 |
+
lr=config.auto_config.initial_lr_G,
|
| 166 |
+
weight_decay=config.auto_config.l2_reg,
|
| 167 |
+
)
|
| 168 |
+
if self.generator else None
|
| 169 |
+
)
|
| 170 |
+
self.opt_ver = (
|
| 171 |
+
torch.optim.AdamW(
|
| 172 |
+
self.verifier.parameters(),
|
| 173 |
+
lr=config.auto_config.initial_lr_V,
|
| 174 |
+
weight_decay=config.auto_config.l2_reg,
|
| 175 |
+
)
|
| 176 |
+
if self.verifier else None
|
| 177 |
+
)
|
| 178 |
+
|
| 179 |
+
# Memory optimizer
|
| 180 |
+
self.mem = MemoryOptimizer(enable_amp=(self.device != "cpu"))
|
| 181 |
+
self.mem.configure()
|
| 182 |
+
|
| 183 |
+
# Hypothesis controller (aplicado sobre features do modelo)
|
| 184 |
+
self.hypothesis = HypothesisController(
|
| 185 |
+
feat_dim=model.fusion.out_dim,
|
| 186 |
+
out_dim=model.fusion.out_dim,
|
| 187 |
+
config=HypothesisConfig(
|
| 188 |
+
n_hypotheses=config.n_hypotheses,
|
| 189 |
+
threshold_err=config.loss_config.threshold_err,
|
| 190 |
+
),
|
| 191 |
+
).to(self.device)
|
| 192 |
+
|
| 193 |
+
# Histórico de treinamento
|
| 194 |
+
self.history: List[Dict[str, Any]] = []
|
| 195 |
+
# NOVO v2.0: métricas estendidas
|
| 196 |
+
self.metrics: Dict[str, List[float]] = {
|
| 197 |
+
"loss": [], "ppl": [], "v_mean": [], "lr_G": [], "lr_V": [],
|
| 198 |
+
"l_g": [], "l_v": [], "l_ah": [], "l_ewc": [],
|
| 199 |
+
"hypothesis_activated": [],
|
| 200 |
+
}
|
| 201 |
+
|
| 202 |
+
# =====================================================================
|
| 203 |
+
# Cálculo de v (verificador) e τ (anti-alucinação) — VERSÃO CORRIGIDA
|
| 204 |
+
# =====================================================================
|
| 205 |
+
|
| 206 |
+
def _compute_v_and_tau(
|
| 207 |
+
self,
|
| 208 |
+
fused: torch.Tensor,
|
| 209 |
+
input_ids_a: torch.Tensor,
|
| 210 |
+
input_ids_b: torch.Tensor,
|
| 211 |
+
gen_logits: Optional[torch.Tensor] = None,
|
| 212 |
+
) -> Tuple[List[torch.Tensor], List[torch.Tensor]]:
|
| 213 |
+
"""Computa valores do verificador (v) e anti-alucinação (τ).
|
| 214 |
+
|
| 215 |
+
VERSÃO CORRIGIDA v2.0:
|
| 216 |
+
- Antes: `v_proxy = sigmoid(fused.mean(-1))` (verificador nunca chamado)
|
| 217 |
+
- Agora: chama o VerifierCNNBiGRU com tokens reais gerados pelo gerador
|
| 218 |
+
- Anti-alucinação: aplicada sobre os logits reais do gerador (não proxy)
|
| 219 |
+
"""
|
| 220 |
+
v_list = []
|
| 221 |
+
tau_list = []
|
| 222 |
+
B = fused.size(0)
|
| 223 |
+
device = fused.device
|
| 224 |
+
|
| 225 |
+
# --- Verificador REAL ---
|
| 226 |
+
if self.verifier is not None and self.cfg.use_verifier_real:
|
| 227 |
+
try:
|
| 228 |
+
# Passo gerado: pegar argmax dos gen_logits por timestep (se disponível)
|
| 229 |
+
if gen_logits is not None:
|
| 230 |
+
# gen_logits: [B, T, V] → argmax → [B, T]
|
| 231 |
+
passo_gerado = gen_logits.argmax(dim=-1)
|
| 232 |
+
else:
|
| 233 |
+
# Sem gerador: usar input_ids_b como "passo hipotético"
|
| 234 |
+
passo_gerado = input_ids_b
|
| 235 |
+
|
| 236 |
+
# Premissas: usar input_ids_a como contexto acumulado
|
| 237 |
+
premissas_a = input_ids_a
|
| 238 |
+
# Para o stream B do verificador, usamos input_ids_b como premissa B
|
| 239 |
+
premissas_b = input_ids_b
|
| 240 |
+
|
| 241 |
+
# Garantir que passo_gerado tenha mesma dimensão de tempo que premissas
|
| 242 |
+
# (verificador espera 4 entradas: premissas_a, premissas_b, passo_a, passo_b)
|
| 243 |
+
passo_a = passo_gerado
|
| 244 |
+
passo_b = passo_gerado # simplificação: mesmo passo para ambos streams
|
| 245 |
+
|
| 246 |
+
# Chamada REAL ao verificador
|
| 247 |
+
with torch.enable_grad():
|
| 248 |
+
v_real = self.verifier(premissas_a, premissas_b, passo_a, passo_b) # [B, 1]
|
| 249 |
+
v_list.append(v_real)
|
| 250 |
+
except Exception as e:
|
| 251 |
+
logger.warning(f"Verifier forward falhou: {e}; usando proxy")
|
| 252 |
+
v_proxy = torch.sigmoid(fused.mean(dim=-1, keepdim=True))
|
| 253 |
+
v_list.append(v_proxy)
|
| 254 |
+
elif self.verifier is not None and not self.cfg.use_verifier_real:
|
| 255 |
+
# Proxy explícito (modo degrau)
|
| 256 |
+
v_proxy = torch.sigmoid(fused.mean(dim=-1, keepdim=True))
|
| 257 |
+
v_list.append(v_proxy)
|
| 258 |
+
else:
|
| 259 |
+
# Sem verifier: v = 1 (sem punição)
|
| 260 |
+
v_list.append(torch.ones(B, 1, device=device))
|
| 261 |
+
|
| 262 |
+
# --- Anti-hallucination REAL ---
|
| 263 |
+
if self.anti_hall is not None:
|
| 264 |
+
try:
|
| 265 |
+
if gen_logits is not None:
|
| 266 |
+
# Aplicar τ sobre cada timestep dos logits do gerador
|
| 267 |
+
# gen_logits: [B, T, V]
|
| 268 |
+
for t in range(gen_logits.size(1)):
|
| 269 |
+
tau_t = self.anti_hall(gen_logits[:, t, :]) # [B, 1]
|
| 270 |
+
tau_list.append(tau_t)
|
| 271 |
+
else:
|
| 272 |
+
# Sem gerador: aplicar τ sobre lm_head(fused)
|
| 273 |
+
logits_proxy = self.model.lm_head(fused) # [B, V]
|
| 274 |
+
tau = self.anti_hall(logits_proxy)
|
| 275 |
+
tau_list.append(tau)
|
| 276 |
+
except Exception as e:
|
| 277 |
+
logger.warning(f"Anti-hallucination falhou: {e}")
|
| 278 |
+
tau_list.append(torch.ones(B, 1, device=device))
|
| 279 |
+
else:
|
| 280 |
+
tau_list.append(torch.ones(B, 1, device=device))
|
| 281 |
+
|
| 282 |
+
return v_list, tau_list
|
| 283 |
+
|
| 284 |
+
# =====================================================================
|
| 285 |
+
# EWC Penalty (NOVO v2.0)
|
| 286 |
+
# =====================================================================
|
| 287 |
+
|
| 288 |
+
def _compute_ewc_penalty(self) -> Optional[torch.Tensor]:
|
| 289 |
+
"""Computa a penalidade EWC se o estado estiver ativo."""
|
| 290 |
+
if self.ewc_state is None or not self.ewc_state.is_active():
|
| 291 |
+
return None
|
| 292 |
+
try:
|
| 293 |
+
return self.ewc_state.penalty(self.model)
|
| 294 |
+
except Exception as e:
|
| 295 |
+
logger.warning(f"EWC penalty falhou: {e}")
|
| 296 |
+
return None
|
| 297 |
+
|
| 298 |
+
# =====================================================================
|
| 299 |
+
# Passo de Treinamento (CORRIGIDO v2.0)
|
| 300 |
+
# =====================================================================
|
| 301 |
+
|
| 302 |
+
def _train_step(
|
| 303 |
+
self,
|
| 304 |
+
batch: Dict[str, torch.Tensor],
|
| 305 |
+
) -> Dict[str, float]:
|
| 306 |
+
"""Executa um passo de treinamento. Retorna métricas.
|
| 307 |
+
|
| 308 |
+
CORREÇÕES v2.0:
|
| 309 |
+
- Hipótese controller output REALMENTE usado (soma ao fused)
|
| 310 |
+
- Verificador REALMENTE chamado (não proxy)
|
| 311 |
+
- EWC penalty integrado à perda total
|
| 312 |
+
- Sem double-counting de CE loss (apenas via MultiLoss)
|
| 313 |
+
"""
|
| 314 |
+
# Move para device
|
| 315 |
+
input_ids_a = batch["input_ids_a"].to(self.device, non_blocking=True)
|
| 316 |
+
input_ids_b = batch["input_ids_b"].to(self.device, non_blocking=True)
|
| 317 |
+
attn_mask_a = batch["attn_mask_a"].to(self.device, non_blocking=True)
|
| 318 |
+
attn_mask_b = batch["attn_mask_b"].to(self.device, non_blocking=True)
|
| 319 |
+
images = batch["images"].to(self.device, non_blocking=True)
|
| 320 |
+
audios = batch["audios"].to(self.device, non_blocking=True)
|
| 321 |
+
labels = batch["labels"].to(self.device, non_blocking=True)
|
| 322 |
+
|
| 323 |
+
# Forward
|
| 324 |
+
with self.mem.amp_context(self.device):
|
| 325 |
+
out = self.model(
|
| 326 |
+
input_ids_a, input_ids_b,
|
| 327 |
+
images=images, audios=audios,
|
| 328 |
+
attn_mask_a=attn_mask_a, attn_mask_b=attn_mask_b,
|
| 329 |
+
mode="classify",
|
| 330 |
+
)
|
| 331 |
+
logits = out["logits"] # [B, num_classes]
|
| 332 |
+
fused = out["fused"] # [B, fusion_dim]
|
| 333 |
+
|
| 334 |
+
# Verifier + anti-hallucination + gerador
|
| 335 |
+
gen_logits = None
|
| 336 |
+
if self.generator is not None and self.cfg.use_generator:
|
| 337 |
+
try:
|
| 338 |
+
gen_out = self.generator(
|
| 339 |
+
input_ids_a, input_ids_b,
|
| 340 |
+
target_proof=input_ids_b, # teacher forcing
|
| 341 |
+
bos_id=self.tokenizer.bos_id,
|
| 342 |
+
eos_id=self.tokenizer.eos_id,
|
| 343 |
+
)
|
| 344 |
+
gen_logits = gen_out["logits"]
|
| 345 |
+
except Exception as e:
|
| 346 |
+
logger.debug(f"Generator forward falhou: {e}")
|
| 347 |
+
gen_logits = None
|
| 348 |
+
|
| 349 |
+
# Verifier e anti-hallucination (chamadas REAIS)
|
| 350 |
+
v_list, tau_list = self._compute_v_and_tau(
|
| 351 |
+
fused, input_ids_a, input_ids_b, gen_logits
|
| 352 |
+
)
|
| 353 |
+
|
| 354 |
+
# Constrói logits para a MultiLoss
|
| 355 |
+
if gen_logits is not None:
|
| 356 |
+
loss_logits = gen_logits
|
| 357 |
+
loss_targets = input_ids_b
|
| 358 |
+
else:
|
| 359 |
+
# Proxy: expande logits de classificação para [B, T, V]
|
| 360 |
+
B = fused.size(0)
|
| 361 |
+
T = input_ids_b.size(1)
|
| 362 |
+
loss_logits = self.model.lm_head(fused).unsqueeze(1).expand(-1, T, -1)
|
| 363 |
+
loss_targets = input_ids_b
|
| 364 |
+
|
| 365 |
+
# NOVO v2.0: Hypothesis controller — usar a SAÍDA (não só a flag)
|
| 366 |
+
# v_mean over list of [B,1] → [B,1]
|
| 367 |
+
v_mean_tensor = torch.stack([v.squeeze(-1) for v in v_list]).mean(dim=0).unsqueeze(-1)
|
| 368 |
+
try:
|
| 369 |
+
hyp_out = self.hypothesis(fused, v_mean_tensor)
|
| 370 |
+
# A saída do hypothesis controller É USADA para enriquecer o fused
|
| 371 |
+
# (somando como residual). Antes era ignorada.
|
| 372 |
+
if self.cfg.use_hypothesis_output:
|
| 373 |
+
# Residual: fused = fused + alpha * (hyp_out - fused)
|
| 374 |
+
# onde alpha é derivado do gate
|
| 375 |
+
alpha = hyp_out["gate"].mean().clamp(0.0, 1.0).item()
|
| 376 |
+
if alpha > 1e-3:
|
| 377 |
+
# Re-projetar para mesma dimensão (hyp_out pode ter dimensão diferente)
|
| 378 |
+
hyp_out_proj = hyp_out["output"]
|
| 379 |
+
if hyp_out_proj.size(-1) != fused.size(-1):
|
| 380 |
+
# Projeção via Linear on-the-fly (lazy)
|
| 381 |
+
if not hasattr(self, "_hyp_proj"):
|
| 382 |
+
self._hyp_proj = nn.Linear(
|
| 383 |
+
hyp_out_proj.size(-1), fused.size(-1), bias=False
|
| 384 |
+
).to(self.device)
|
| 385 |
+
nn.init.xavier_uniform_(self._hyp_proj.weight)
|
| 386 |
+
hyp_out_proj = self._hyp_proj(hyp_out_proj)
|
| 387 |
+
# Soma residual ponderada pelo gate
|
| 388 |
+
fused = fused + alpha * hyp_out_proj
|
| 389 |
+
# Recalcular logits de classificação com fused atualizado
|
| 390 |
+
# (não refazemos forward do model para economizar compute;
|
| 391 |
+
# apenas atualizamos a saída final do lm_head/classifier)
|
| 392 |
+
# NOTA: isto é uma aproximação — o ideal seria refazer o forward
|
| 393 |
+
except Exception as e:
|
| 394 |
+
logger.debug(f"Hypothesis controller falhou: {e}")
|
| 395 |
+
hyp_out = {"activated": False, "gate": torch.zeros_like(v_mean_tensor)}
|
| 396 |
+
|
| 397 |
+
# NOVO v2.0: EWC Penalty
|
| 398 |
+
ewc_pen = self._compute_ewc_penalty()
|
| 399 |
+
|
| 400 |
+
# Curvatura (estimativa por diferenças finitas)
|
| 401 |
+
try:
|
| 402 |
+
curv = estimate_curvature(
|
| 403 |
+
self.model,
|
| 404 |
+
lambda: self.loss_fn(
|
| 405 |
+
loss_logits, loss_targets, v_list, tau_list,
|
| 406 |
+
models_for_reg=[self.model],
|
| 407 |
+
),
|
| 408 |
+
eps=self.cfg.loss_config.curvature_eps,
|
| 409 |
+
n_samples=2,
|
| 410 |
+
)
|
| 411 |
+
self.loss_fn.set_curvature_estimate(curv)
|
| 412 |
+
except Exception as e:
|
| 413 |
+
logger.debug(f"Curvatura falhou: {e}")
|
| 414 |
+
|
| 415 |
+
# Compute losses (com EWC penalty integrado)
|
| 416 |
+
losses = self.loss_fn(
|
| 417 |
+
loss_logits, loss_targets, v_list, tau_list,
|
| 418 |
+
models_for_reg=[m for m in [self.model, self.generator, self.verifier] if m is not None],
|
| 419 |
+
ewc_penalty=ewc_pen,
|
| 420 |
+
)
|
| 421 |
+
|
| 422 |
+
# Loss de classificação (NÃO duplica a L_G — esta é uma tarefa auxiliar)
|
| 423 |
+
ce_loss = F.cross_entropy(logits, labels)
|
| 424 |
+
|
| 425 |
+
# Combina: perda de geração multimodal + perda de classificação
|
| 426 |
+
total_loss = losses["total"] + ce_loss
|
| 427 |
+
|
| 428 |
+
# Backward
|
| 429 |
+
self.opt_model.zero_grad(set_to_none=True)
|
| 430 |
+
if self.opt_gen: self.opt_gen.zero_grad(set_to_none=True)
|
| 431 |
+
if self.opt_ver: self.opt_ver.zero_grad(set_to_none=True)
|
| 432 |
+
|
| 433 |
+
total_loss.backward()
|
| 434 |
+
|
| 435 |
+
# Clipping
|
| 436 |
+
self.auto.apply_clipping(self.model)
|
| 437 |
+
if self.generator: self.auto.apply_clipping(self.generator)
|
| 438 |
+
if self.verifier: self.auto.apply_clipping(self.verifier)
|
| 439 |
+
|
| 440 |
+
# Normas de gradiente
|
| 441 |
+
grad_G = self.auto.compute_total_grad_norm(self.model)
|
| 442 |
+
grad_V = self.auto.compute_total_grad_norm(self.verifier) if self.verifier else 0.0
|
| 443 |
+
|
| 444 |
+
# v_mean e L_V_mean
|
| 445 |
+
v_mean = float(torch.stack([v.squeeze(-1) for v in v_list]).mean().item())
|
| 446 |
+
L_V_mean = float(losses["l_v"].item()) if isinstance(losses["l_v"], torch.Tensor) else 0.0
|
| 447 |
+
|
| 448 |
+
# Ajuste dinâmico de LR
|
| 449 |
+
novo_lr_G, novo_lr_V = self.auto.update_learning_rates(
|
| 450 |
+
grad_G, grad_V, v_mean, L_V_mean
|
| 451 |
+
)
|
| 452 |
+
self.auto.set_optimizer_lr(self.opt_model, novo_lr_G)
|
| 453 |
+
if self.opt_gen: self.auto.set_optimizer_lr(self.opt_gen, novo_lr_G)
|
| 454 |
+
if self.opt_ver: self.auto.set_optimizer_lr(self.opt_ver, novo_lr_V)
|
| 455 |
+
|
| 456 |
+
# Step
|
| 457 |
+
self.opt_model.step()
|
| 458 |
+
if self.opt_gen: self.opt_gen.step()
|
| 459 |
+
if self.opt_ver: self.opt_ver.step()
|
| 460 |
+
|
| 461 |
+
# Spectral norm após step
|
| 462 |
+
if self.cfg.auto_config.apply_after_step:
|
| 463 |
+
self.auto.apply_spectral_norm(self.model)
|
| 464 |
+
if self.generator: self.auto.apply_spectral_norm(self.generator)
|
| 465 |
+
if self.verifier: self.auto.apply_spectral_norm(self.verifier)
|
| 466 |
+
|
| 467 |
+
# Métricas
|
| 468 |
+
loss_val = float(total_loss.item())
|
| 469 |
+
try:
|
| 470 |
+
ppl = math.exp(min(loss_val, 20)) # evita overflow
|
| 471 |
+
except OverflowError:
|
| 472 |
+
ppl = float("inf")
|
| 473 |
+
|
| 474 |
+
# Hipótese info
|
| 475 |
+
hyp_activated = bool(hyp_out.get("activated", False)) if isinstance(hyp_out, dict) else False
|
| 476 |
+
|
| 477 |
+
# EWC info
|
| 478 |
+
l_ewc_val = float(losses.get("l_ewc", torch.tensor(0.0)).item()) if isinstance(
|
| 479 |
+
losses.get("l_ewc", 0.0), torch.Tensor
|
| 480 |
+
) else 0.0
|
| 481 |
+
|
| 482 |
+
return {
|
| 483 |
+
"loss": loss_val,
|
| 484 |
+
"ppl": ppl,
|
| 485 |
+
"v_mean": v_mean,
|
| 486 |
+
"lr_G": novo_lr_G,
|
| 487 |
+
"lr_V": novo_lr_V,
|
| 488 |
+
"l_g": float(losses["l_g"].item()) if isinstance(losses["l_g"], torch.Tensor) else 0.0,
|
| 489 |
+
"l_v": L_V_mean,
|
| 490 |
+
"l_ah": float(losses["l_ah"].item()) if isinstance(losses["l_ah"], torch.Tensor) else 0.0,
|
| 491 |
+
"l_ewc": l_ewc_val,
|
| 492 |
+
"hypothesis_activated": hyp_activated,
|
| 493 |
+
}
|
| 494 |
+
|
| 495 |
+
# =====================================================================
|
| 496 |
+
# Synergy Search Helper
|
| 497 |
+
# =====================================================================
|
| 498 |
+
|
| 499 |
+
def _eval_config(self, cfg: Dict, sample_batch: Dict[str, torch.Tensor]) -> float:
|
| 500 |
+
"""Avalia uma configuração de perda em um mini-batch (para synergy search)."""
|
| 501 |
+
# Salva config atual
|
| 502 |
+
original = {k: getattr(self.loss_fn.cfg, k) for k in cfg}
|
| 503 |
+
# Aplica nova config
|
| 504 |
+
for k, v in cfg.items():
|
| 505 |
+
if hasattr(self.loss_fn.cfg, k):
|
| 506 |
+
setattr(self.loss_fn.cfg, k, v)
|
| 507 |
+
try:
|
| 508 |
+
metrics = self._train_step(sample_batch)
|
| 509 |
+
return metrics["loss"]
|
| 510 |
+
finally:
|
| 511 |
+
# Restaura
|
| 512 |
+
for k, v in original.items():
|
| 513 |
+
setattr(self.loss_fn.cfg, k, v)
|
| 514 |
+
|
| 515 |
+
# =====================================================================
|
| 516 |
+
# Loop Principal de Treinamento
|
| 517 |
+
# =====================================================================
|
| 518 |
+
|
| 519 |
+
def train(
|
| 520 |
+
self,
|
| 521 |
+
dataloader: DataLoader,
|
| 522 |
+
) -> Dict[str, Any]:
|
| 523 |
+
"""Executa o treinamento completo."""
|
| 524 |
+
logger.info("Iniciando treinamento cooperativo CNN-BiGRU v2.0")
|
| 525 |
+
logger.info(f" epochs={self.cfg.num_epochs}, batch_size={self.cfg.batch_size}, "
|
| 526 |
+
f"n_synergy={self.cfg.n_synergy_attempts}")
|
| 527 |
+
logger.info(f" EWC ativo: {self.ewc_state is not None and self.ewc_state.is_active()}")
|
| 528 |
+
logger.info(f" Verifier real: {self.cfg.use_verifier_real}")
|
| 529 |
+
logger.info(f" Hipótese output usado: {self.cfg.use_hypothesis_output}")
|
| 530 |
+
|
| 531 |
+
start_time = time.time()
|
| 532 |
+
|
| 533 |
+
for epoch in range(self.cfg.num_epochs):
|
| 534 |
+
epoch_loss = 0.0
|
| 535 |
+
epoch_ppl = 0.0
|
| 536 |
+
n_batches = 0
|
| 537 |
+
hypothesis_activations = 0
|
| 538 |
+
|
| 539 |
+
# Synergy search no início da época (apenas época 0 para economizar tempo)
|
| 540 |
+
if self.cfg.use_synergy_search and epoch == 0:
|
| 541 |
+
try:
|
| 542 |
+
# Pega um batch de amostra para avaliação
|
| 543 |
+
sample_batch = next(iter(dataloader))
|
| 544 |
+
sample_batch = {k: (v.to(self.device) if isinstance(v, torch.Tensor) else v)
|
| 545 |
+
for k, v in sample_batch.items()}
|
| 546 |
+
|
| 547 |
+
base_cfg = {
|
| 548 |
+
"alpha": self.loss_fn.cfg.alpha,
|
| 549 |
+
"beta": self.loss_fn.cfg.beta,
|
| 550 |
+
"gamma_loss": self.loss_fn.cfg.gamma_loss,
|
| 551 |
+
"delta": self.loss_fn.cfg.delta,
|
| 552 |
+
"lambda_penal": self.loss_fn.cfg.lambda_penal,
|
| 553 |
+
"mu_exp_penal": self.loss_fn.cfg.mu_exp_penal,
|
| 554 |
+
}
|
| 555 |
+
|
| 556 |
+
logger.info(f"Synergy search: tentando {self.cfg.n_synergy_attempts} configurações...")
|
| 557 |
+
best_cfg = self.synergy.search(
|
| 558 |
+
base_config=base_cfg,
|
| 559 |
+
eval_fn=lambda c: self._eval_config(c, sample_batch),
|
| 560 |
+
)
|
| 561 |
+
|
| 562 |
+
# Aplica a melhor configuração
|
| 563 |
+
for k, v in best_cfg.items():
|
| 564 |
+
if hasattr(self.loss_fn.cfg, k):
|
| 565 |
+
setattr(self.loss_fn.cfg, k, v)
|
| 566 |
+
logger.info("Synergy search concluído. Melhor config aplicada.")
|
| 567 |
+
except Exception as e:
|
| 568 |
+
logger.warning(f"Synergy search falhou: {e}")
|
| 569 |
+
|
| 570 |
+
# Loop de batches
|
| 571 |
+
self.model.train()
|
| 572 |
+
if self.generator: self.generator.train()
|
| 573 |
+
if self.verifier: self.verifier.train()
|
| 574 |
+
if self.anti_hall: self.anti_hall.train()
|
| 575 |
+
self.hypothesis.train()
|
| 576 |
+
|
| 577 |
+
for batch_idx, batch in enumerate(dataloader):
|
| 578 |
+
if batch_idx >= self.cfg.max_batches_per_epoch:
|
| 579 |
+
break
|
| 580 |
+
|
| 581 |
+
try:
|
| 582 |
+
metrics = self._train_step(batch)
|
| 583 |
+
epoch_loss += metrics["loss"]
|
| 584 |
+
epoch_ppl += metrics["ppl"]
|
| 585 |
+
n_batches += 1
|
| 586 |
+
if metrics["hypothesis_activated"]:
|
| 587 |
+
hypothesis_activations += 1
|
| 588 |
+
|
| 589 |
+
# Armazena todas as métricas
|
| 590 |
+
for k in self.metrics:
|
| 591 |
+
if k in metrics:
|
| 592 |
+
self.metrics[k].append(metrics[k])
|
| 593 |
+
|
| 594 |
+
if batch_idx % self.cfg.log_every == 0:
|
| 595 |
+
logger.info(
|
| 596 |
+
f"Época {epoch+1} | Batch {batch_idx} | "
|
| 597 |
+
f"loss={metrics['loss']:.4f} | ppl={metrics['ppl']:.2f} | "
|
| 598 |
+
f"v_mean={metrics['v_mean']:.3f} | lr_G={metrics['lr_G']:.2e} | "
|
| 599 |
+
f"l_ewc={metrics['l_ewc']:.4f} | "
|
| 600 |
+
f"hyp_act={'Y' if metrics['hypothesis_activated'] else 'N'}"
|
| 601 |
+
)
|
| 602 |
+
except Exception as e:
|
| 603 |
+
logger.error(f"Erro no batch {batch_idx}: {e}", exc_info=True)
|
| 604 |
+
continue
|
| 605 |
+
|
| 606 |
+
# Cleanup de memória ao final da época
|
| 607 |
+
self.mem.cleanup()
|
| 608 |
+
|
| 609 |
+
avg_loss = epoch_loss / max(n_batches, 1)
|
| 610 |
+
avg_ppl = epoch_ppl / max(n_batches, 1)
|
| 611 |
+
logger.info(
|
| 612 |
+
f"Época {epoch+1} concluída | avg_loss={avg_loss:.4f} | "
|
| 613 |
+
f"avg_ppl={avg_ppl:.2f} | hyp_activations={hypothesis_activations} | "
|
| 614 |
+
f"batches_ok={n_batches}/{self.cfg.max_batches_per_epoch}"
|
| 615 |
+
)
|
| 616 |
+
|
| 617 |
+
# Detecta falha sistêmica
|
| 618 |
+
if n_batches == 0:
|
| 619 |
+
logger.error(f"Época {epoch+1}: NENHUM batch teve sucesso — possível bug no pipeline")
|
| 620 |
+
|
| 621 |
+
self.history.append({
|
| 622 |
+
"epoch": epoch + 1,
|
| 623 |
+
"avg_loss": avg_loss,
|
| 624 |
+
"avg_ppl": avg_ppl,
|
| 625 |
+
"n_batches": n_batches,
|
| 626 |
+
"hypothesis_activations": hypothesis_activations,
|
| 627 |
+
"epoch_loss_total": epoch_loss,
|
| 628 |
+
})
|
| 629 |
+
|
| 630 |
+
elapsed = time.time() - start_time
|
| 631 |
+
logger.info(f"Treinamento concluído em {elapsed:.1f}s")
|
| 632 |
+
|
| 633 |
+
return {
|
| 634 |
+
"history": self.history,
|
| 635 |
+
"metrics": self.metrics,
|
| 636 |
+
"final_loss": self.history[-1]["avg_loss"] if self.history else 0.0,
|
| 637 |
+
"final_ppl": self.history[-1]["avg_ppl"] if self.history else 0.0,
|
| 638 |
+
"elapsed_s": elapsed,
|
| 639 |
+
"synergy_history": self.synergy.history,
|
| 640 |
+
"auto_history": self.auto.history,
|
| 641 |
+
"ewc_active": self.ewc_state is not None and self.ewc_state.is_active(),
|
| 642 |
+
"ewc_num_tasks": self.ewc_state.num_tasks() if self.ewc_state else 0,
|
| 643 |
+
}
|
| 644 |
+
|
| 645 |
+
# =====================================================================
|
| 646 |
+
# Consolidar tarefa para EWC (NOVO v2.0)
|
| 647 |
+
# =====================================================================
|
| 648 |
+
|
| 649 |
+
def consolidate_task(
|
| 650 |
+
self,
|
| 651 |
+
forward_fn,
|
| 652 |
+
compute_fisher: bool = True,
|
| 653 |
+
) -> None:
|
| 654 |
+
"""Consolida a tarefa atual no estado EWC.
|
| 655 |
+
|
| 656 |
+
Args:
|
| 657 |
+
forward_fn: callable que retorna (logits, target) ou logits
|
| 658 |
+
para cálculo da Fisher Information.
|
| 659 |
+
compute_fisher: se True, recalcula a Fisher; se False, apenas
|
| 660 |
+
armazena o theta_star atual.
|
| 661 |
+
"""
|
| 662 |
+
if self.ewc_state is None:
|
| 663 |
+
logger.warning("EWC não configurado — consolidate_task ignorado")
|
| 664 |
+
return
|
| 665 |
+
try:
|
| 666 |
+
self.ewc_state.consolidate(self.model, forward_fn=forward_fn if compute_fisher else None)
|
| 667 |
+
logger.info(f"Tarefa consolidada no EWC (total: {self.ewc_state.num_tasks()})")
|
| 668 |
+
except Exception as e:
|
| 669 |
+
logger.error(f"Falha ao consolidar tarefa no EWC: {e}")
|
| 670 |
+
|
| 671 |
+
|
| 672 |
+
__all__ = ["TrainerConfig", "CooperativeTrainer"]
|
cnn_bigru/utils/__init__.py
ADDED
|
File without changes
|
cnn_bigru/utils/ewc.py
ADDED
|
@@ -0,0 +1,382 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
EWC (Elastic Weight Consolidation) — Aprendizado Contínuo sem Catastrophic Forgetting.
|
| 3 |
+
|
| 4 |
+
Implementa:
|
| 5 |
+
- Cálculo da diagonal da Matriz de Informação de Fisher (empirical Fisher)
|
| 6 |
+
- Estado consolidado (theta_star + Fisher) para múltiplas tarefas
|
| 7 |
+
- Online EWC (média exponencial de Fisher entre tarefas)
|
| 8 |
+
- Penalidade diferenciável L_EWC = sum_i (lambda/2) * F_i * (theta_i - theta_star_i)^2
|
| 9 |
+
|
| 10 |
+
Referência matemática: docs/MATH_ANALYSIS.md seção 2.
|
| 11 |
+
|
| 12 |
+
Autor: CNN-BiGRU Project
|
| 13 |
+
"""
|
| 14 |
+
from __future__ import annotations
|
| 15 |
+
|
| 16 |
+
import logging
|
| 17 |
+
from dataclasses import dataclass, field
|
| 18 |
+
from typing import Dict, List, Optional, Tuple
|
| 19 |
+
|
| 20 |
+
import torch
|
| 21 |
+
import torch.nn as nn
|
| 22 |
+
|
| 23 |
+
logger = logging.getLogger(__name__)
|
| 24 |
+
|
| 25 |
+
|
| 26 |
+
# ============================================================================
|
| 27 |
+
# Configuração
|
| 28 |
+
# ============================================================================
|
| 29 |
+
|
| 30 |
+
@dataclass
|
| 31 |
+
class EWCConfig:
|
| 32 |
+
"""Configuração do EWC (Elastic Weight Consolidation)."""
|
| 33 |
+
enabled: bool = True
|
| 34 |
+
# Peso da penalidade EWC na perda total
|
| 35 |
+
lambda_ewc: float = 100.0
|
| 36 |
+
# Número de amostras usadas para estimar a diagonal de Fisher
|
| 37 |
+
n_samples_fisher: int = 200
|
| 38 |
+
# Decaimento para Online EWC: F_agg = gamma*F_old + (1-gamma)*F_new
|
| 39 |
+
# Se None, usa soma direta (EWC padrão — explode em memória após muitas tarefas)
|
| 40 |
+
online_gamma: Optional[float] = 0.9
|
| 41 |
+
# Clip para estabilizar Fisher (evitar valores extremos)
|
| 42 |
+
fisher_clip_max: float = 1e4
|
| 43 |
+
# Tipos de parâmetros a penalizar (outros como LayerNorm/bias podem ser excluídos)
|
| 44 |
+
include_param_names: Tuple[str, ...] = ("weight",)
|
| 45 |
+
exclude_param_names: Tuple[str, ...] = ("bias", "layernorm", "layer_norm", "bn", "batchnorm")
|
| 46 |
+
# Device padrão
|
| 47 |
+
device: str = "cpu"
|
| 48 |
+
|
| 49 |
+
|
| 50 |
+
# ============================================================================
|
| 51 |
+
# Estado EWC
|
| 52 |
+
# ============================================================================
|
| 53 |
+
|
| 54 |
+
class EWCState:
|
| 55 |
+
"""
|
| 56 |
+
Mantém o estado consolidado do EWC: lista de (theta_star, fisher) por tarefa.
|
| 57 |
+
|
| 58 |
+
Para Online EWC, mantém uma única Fisher agregada e theta_star atual.
|
| 59 |
+
"""
|
| 60 |
+
|
| 61 |
+
def __init__(self, config: EWCConfig):
|
| 62 |
+
self.config = config
|
| 63 |
+
# Lista de tarefas: cada entrada é um dict {param_name: (theta_star, fisher)}
|
| 64 |
+
self.tasks: List[Dict[str, Tuple[torch.Tensor, torch.Tensor]]] = []
|
| 65 |
+
# Online: theta_star e fisher agregados
|
| 66 |
+
self.online_theta: Dict[str, torch.Tensor] = {}
|
| 67 |
+
self.online_fisher: Dict[str, torch.Tensor] = {}
|
| 68 |
+
|
| 69 |
+
# ----------------------------------------------------------------------
|
| 70 |
+
# Seleção de parâmetros
|
| 71 |
+
# ----------------------------------------------------------------------
|
| 72 |
+
|
| 73 |
+
def _select_params(self, model: nn.Module) -> List[Tuple[str, torch.Tensor]]:
|
| 74 |
+
"""Seleciona parâmetros treináveis elegíveis para EWC."""
|
| 75 |
+
selected = []
|
| 76 |
+
for name, param in model.named_parameters():
|
| 77 |
+
if not param.requires_grad:
|
| 78 |
+
continue
|
| 79 |
+
name_lower = name.lower()
|
| 80 |
+
# Excluir nomes na lista negra
|
| 81 |
+
if any(ex in name_lower for ex in self.config.exclude_param_names):
|
| 82 |
+
continue
|
| 83 |
+
# Incluir apenas nomes na lista branca (se especificada)
|
| 84 |
+
if self.config.include_param_names and not any(
|
| 85 |
+
inc in name_lower for inc in self.config.include_param_names
|
| 86 |
+
):
|
| 87 |
+
continue
|
| 88 |
+
selected.append((name, param))
|
| 89 |
+
return selected
|
| 90 |
+
|
| 91 |
+
# ----------------------------------------------------------------------
|
| 92 |
+
# Cálculo da Diagonal de Fisher
|
| 93 |
+
# ----------------------------------------------------------------------
|
| 94 |
+
|
| 95 |
+
@torch.no_grad()
|
| 96 |
+
def compute_fisher(
|
| 97 |
+
self,
|
| 98 |
+
model: nn.Module,
|
| 99 |
+
forward_fn,
|
| 100 |
+
n_samples: Optional[int] = None,
|
| 101 |
+
) -> Dict[str, torch.Tensor]:
|
| 102 |
+
"""
|
| 103 |
+
Calcula a diagonal da matriz de Fisher Information (empirical Fisher).
|
| 104 |
+
|
| 105 |
+
Args:
|
| 106 |
+
model: modelo treinado (parâmetros congelados idealmente)
|
| 107 |
+
forward_fn: callable que retorna (logits, target) ou logits.
|
| 108 |
+
Deve aceitar nenhum argumento ou um índice.
|
| 109 |
+
n_samples: número de amostras (default: config.n_samples_fisher)
|
| 110 |
+
|
| 111 |
+
Returns:
|
| 112 |
+
Dict {param_name: fisher_diag_tensor} (mesmo shape do parâmetro)
|
| 113 |
+
|
| 114 |
+
Fórmula:
|
| 115 |
+
F_i = (1/N) sum_n [grad_i log p(y_n* | x_n; theta*)]^2
|
| 116 |
+
onde y_n* = argmax_y p(y | x_n; theta*)
|
| 117 |
+
"""
|
| 118 |
+
if n_samples is None:
|
| 119 |
+
n_samples = self.config.n_samples_fisher
|
| 120 |
+
|
| 121 |
+
model.eval()
|
| 122 |
+
selected = self._select_params(model)
|
| 123 |
+
|
| 124 |
+
# Inicializar acumulador de grad^2
|
| 125 |
+
fisher: Dict[str, torch.Tensor] = {
|
| 126 |
+
name: torch.zeros_like(param.detach(), device=param.device)
|
| 127 |
+
for name, param in selected
|
| 128 |
+
}
|
| 129 |
+
count = 0
|
| 130 |
+
|
| 131 |
+
# Garantir que requires_grad está ativo para os parâmetros selecionados
|
| 132 |
+
# durante todo o cálculo da Fisher (restaurado no finally)
|
| 133 |
+
original_requires_grad: Dict[str, bool] = {}
|
| 134 |
+
for name, param in selected:
|
| 135 |
+
original_requires_grad[name] = param.requires_grad
|
| 136 |
+
param.requires_grad_(True)
|
| 137 |
+
|
| 138 |
+
try:
|
| 139 |
+
for i in range(n_samples):
|
| 140 |
+
try:
|
| 141 |
+
# Resetar gradientes
|
| 142 |
+
for _, param in selected:
|
| 143 |
+
if param.grad is not None:
|
| 144 |
+
param.grad.zero_()
|
| 145 |
+
|
| 146 |
+
# Forward + Backward — TODO em contexto com grad habilitado
|
| 147 |
+
# IMPORTANTE: usar torch.set_grad_enabled(True) para forçar gradientes
|
| 148 |
+
# mesmo se o caller estiver em contexto no_grad. Deve cobrir TANTO
|
| 149 |
+
# o forward QUANTO o backward (senão as operações pós-forward perdem grad_fn).
|
| 150 |
+
prev_grad_mode = torch.is_grad_enabled()
|
| 151 |
+
torch.set_grad_enabled(True)
|
| 152 |
+
try:
|
| 153 |
+
result = forward_fn(i)
|
| 154 |
+
|
| 155 |
+
if isinstance(result, tuple) and len(result) == 2:
|
| 156 |
+
logits, target = result
|
| 157 |
+
else:
|
| 158 |
+
logits = result
|
| 159 |
+
# Usar y* = argmax logits (empirical Fisher)
|
| 160 |
+
target = logits.argmax(dim=-1)
|
| 161 |
+
|
| 162 |
+
# Garantir que logits é [B, V] ou [B, T, V] → flatten
|
| 163 |
+
if logits.dim() == 3:
|
| 164 |
+
B, T, V = logits.shape
|
| 165 |
+
logits_flat = logits.reshape(-1, V)
|
| 166 |
+
target_flat = target.reshape(-1)
|
| 167 |
+
else:
|
| 168 |
+
logits_flat = logits
|
| 169 |
+
target_flat = target.reshape(-1)
|
| 170 |
+
|
| 171 |
+
# Log-verossimilhança (CrossEntropy = -log p)
|
| 172 |
+
log_prob = torch.nn.functional.log_softmax(logits_flat, dim=-1)
|
| 173 |
+
# Pegar log_prob das classes target
|
| 174 |
+
picked = log_prob.gather(1, target_flat.unsqueeze(-1)).sum()
|
| 175 |
+
# O sinal negativo: queremos grad de log p (não -log p)
|
| 176 |
+
loss_for_grad = -picked
|
| 177 |
+
|
| 178 |
+
# Backward — AINDA em grad mode (não restauramos ainda)
|
| 179 |
+
loss_for_grad.backward(retain_graph=False)
|
| 180 |
+
finally:
|
| 181 |
+
torch.set_grad_enabled(prev_grad_mode)
|
| 182 |
+
|
| 183 |
+
# Acumular grad^2
|
| 184 |
+
for name, param in selected:
|
| 185 |
+
if param.grad is not None:
|
| 186 |
+
g2 = param.grad.detach() ** 2
|
| 187 |
+
fisher[name] += g2
|
| 188 |
+
|
| 189 |
+
count += 1
|
| 190 |
+
|
| 191 |
+
except Exception as e:
|
| 192 |
+
logger.warning(f"EWC fisher sample {i} falhou: {e}")
|
| 193 |
+
continue
|
| 194 |
+
finally:
|
| 195 |
+
# Restaurar requires_grad original
|
| 196 |
+
for name, param in selected:
|
| 197 |
+
if name in original_requires_grad:
|
| 198 |
+
param.requires_grad_(original_requires_grad[name])
|
| 199 |
+
|
| 200 |
+
if count == 0:
|
| 201 |
+
logger.error("EWC: nenhuma amostra válida para Fisher — usando zeros")
|
| 202 |
+
return fisher
|
| 203 |
+
|
| 204 |
+
# Média
|
| 205 |
+
for name in fisher:
|
| 206 |
+
fisher[name] = (fisher[name] / count).clamp(0, self.config.fisher_clip_max)
|
| 207 |
+
|
| 208 |
+
return fisher
|
| 209 |
+
# ----------------------------------------------------------------------
|
| 210 |
+
# Consolidar tarefa
|
| 211 |
+
# ----------------------------------------------------------------------
|
| 212 |
+
|
| 213 |
+
def consolidate(
|
| 214 |
+
self,
|
| 215 |
+
model: nn.Module,
|
| 216 |
+
fisher: Optional[Dict[str, torch.Tensor]] = None,
|
| 217 |
+
forward_fn=None,
|
| 218 |
+
) -> None:
|
| 219 |
+
"""
|
| 220 |
+
Consolida o estado atual do modelo como uma nova tarefa.
|
| 221 |
+
|
| 222 |
+
Args:
|
| 223 |
+
model: modelo treinado
|
| 224 |
+
fisher: Fisher diagonal pré-computada (se None, calcula via forward_fn)
|
| 225 |
+
forward_fn: necessário se fisher=None
|
| 226 |
+
"""
|
| 227 |
+
if fisher is None:
|
| 228 |
+
if forward_fn is None:
|
| 229 |
+
raise ValueError("EWC.consolidate requer forward_fn quando fisher=None")
|
| 230 |
+
fisher = self.compute_fisher(model, forward_fn)
|
| 231 |
+
|
| 232 |
+
# Snapshot dos parâmetros atuais (theta_star)
|
| 233 |
+
theta_star: Dict[str, torch.Tensor] = {}
|
| 234 |
+
for name, param in model.named_parameters():
|
| 235 |
+
if name in fisher:
|
| 236 |
+
theta_star[name] = param.detach().clone()
|
| 237 |
+
|
| 238 |
+
task_state = {
|
| 239 |
+
name: (theta_star[name], fisher[name]) for name in fisher
|
| 240 |
+
}
|
| 241 |
+
self.tasks.append(task_state)
|
| 242 |
+
|
| 243 |
+
# Atualizar estado online
|
| 244 |
+
if self.config.online_gamma is not None:
|
| 245 |
+
gamma = self.config.online_gamma
|
| 246 |
+
if not self.online_fisher:
|
| 247 |
+
# Primeira tarefa
|
| 248 |
+
for name in fisher:
|
| 249 |
+
self.online_fisher[name] = fisher[name].clone()
|
| 250 |
+
self.online_theta[name] = theta_star[name].clone()
|
| 251 |
+
else:
|
| 252 |
+
# Online EWC: F_agg = gamma * F_old + (1-gamma) * F_new
|
| 253 |
+
# theta_agg = gamma * theta_old + (1-gamma) * theta_new
|
| 254 |
+
for name in fisher:
|
| 255 |
+
if name in self.online_fisher:
|
| 256 |
+
self.online_fisher[name] = (
|
| 257 |
+
gamma * self.online_fisher[name]
|
| 258 |
+
+ (1 - gamma) * fisher[name]
|
| 259 |
+
).clamp(0, self.config.fisher_clip_max)
|
| 260 |
+
self.online_theta[name] = (
|
| 261 |
+
gamma * self.online_theta[name]
|
| 262 |
+
+ (1 - gamma) * theta_star[name]
|
| 263 |
+
)
|
| 264 |
+
else:
|
| 265 |
+
self.online_fisher[name] = fisher[name].clone()
|
| 266 |
+
self.online_theta[name] = theta_star[name].clone()
|
| 267 |
+
else:
|
| 268 |
+
# EWC padrão — somar Fishers de todas as tarefas
|
| 269 |
+
if not self.online_fisher:
|
| 270 |
+
for name in fisher:
|
| 271 |
+
self.online_fisher[name] = fisher[name].clone()
|
| 272 |
+
self.online_theta[name] = theta_star[name].clone()
|
| 273 |
+
else:
|
| 274 |
+
for name in fisher:
|
| 275 |
+
self.online_fisher[name] = (
|
| 276 |
+
self.online_fisher[name] + fisher[name]
|
| 277 |
+
).clamp(0, self.config.fisher_clip_max)
|
| 278 |
+
# Theta online atualiza para o mais recente
|
| 279 |
+
self.online_theta[name] = theta_star[name].clone()
|
| 280 |
+
|
| 281 |
+
logger.info(
|
| 282 |
+
f"EWC consolidado: tarefa #{len(self.tasks)}, "
|
| 283 |
+
f"parâmetros rastreados: {len(theta_star)}"
|
| 284 |
+
)
|
| 285 |
+
|
| 286 |
+
# ----------------------------------------------------------------------
|
| 287 |
+
# Penalidade EWC
|
| 288 |
+
# ----------------------------------------------------------------------
|
| 289 |
+
|
| 290 |
+
def penalty(self, model: nn.Module) -> torch.Tensor:
|
| 291 |
+
"""
|
| 292 |
+
Calcula a penalidade EWC: L_EWC = sum_i (lambda/2) * F_i * (theta_i - theta_star_i)^2
|
| 293 |
+
|
| 294 |
+
Usa o estado online (agregado) se disponível, caso contrário soma de
|
| 295 |
+
todas as tarefas (EWC padrão).
|
| 296 |
+
|
| 297 |
+
Returns:
|
| 298 |
+
Escalar (tensor 0-dim) com a penalidade.
|
| 299 |
+
"""
|
| 300 |
+
if not self.config.enabled or not self.online_fisher:
|
| 301 |
+
# Retorna zero no device do modelo
|
| 302 |
+
try:
|
| 303 |
+
device = next(model.parameters()).device
|
| 304 |
+
except StopIteration:
|
| 305 |
+
device = torch.device(self.config.device)
|
| 306 |
+
return torch.zeros((), device=device)
|
| 307 |
+
|
| 308 |
+
total = None
|
| 309 |
+
for name, param in model.named_parameters():
|
| 310 |
+
if name not in self.online_fisher:
|
| 311 |
+
continue
|
| 312 |
+
F = self.online_fisher[name]
|
| 313 |
+
theta_s = self.online_theta[name]
|
| 314 |
+
# Penalidade quadrática
|
| 315 |
+
diff = param - theta_s.to(param.device)
|
| 316 |
+
pen = (F.to(param.device) * diff.pow(2)).sum()
|
| 317 |
+
if total is None:
|
| 318 |
+
total = pen
|
| 319 |
+
else:
|
| 320 |
+
total = total + pen
|
| 321 |
+
|
| 322 |
+
if total is None:
|
| 323 |
+
try:
|
| 324 |
+
device = next(model.parameters()).device
|
| 325 |
+
except StopIteration:
|
| 326 |
+
device = torch.device(self.config.device)
|
| 327 |
+
return torch.zeros((), device=device)
|
| 328 |
+
|
| 329 |
+
return (self.config.lambda_ewc * 0.5) * total
|
| 330 |
+
|
| 331 |
+
# ----------------------------------------------------------------------
|
| 332 |
+
# Utilitários
|
| 333 |
+
# ----------------------------------------------------------------------
|
| 334 |
+
|
| 335 |
+
def num_tasks(self) -> int:
|
| 336 |
+
return len(self.tasks)
|
| 337 |
+
|
| 338 |
+
def is_active(self) -> bool:
|
| 339 |
+
return self.config.enabled and bool(self.online_fisher)
|
| 340 |
+
|
| 341 |
+
def state_dict(self) -> Dict:
|
| 342 |
+
"""Serializa o estado para checkpoint."""
|
| 343 |
+
return {
|
| 344 |
+
"tasks": [
|
| 345 |
+
{name: (ts.cpu(), f.cpu()) for name, (ts, f) in task.items()}
|
| 346 |
+
for task in self.tasks
|
| 347 |
+
],
|
| 348 |
+
"online_theta": {k: v.cpu() for k, v in self.online_theta.items()},
|
| 349 |
+
"online_fisher": {k: v.cpu() for k, v in self.online_fisher.items()},
|
| 350 |
+
"config": self.config.__dict__,
|
| 351 |
+
}
|
| 352 |
+
|
| 353 |
+
def load_state_dict(self, sd: Dict) -> None:
|
| 354 |
+
"""Carrega o estado a partir de checkpoint."""
|
| 355 |
+
self.tasks = [
|
| 356 |
+
{name: (ts, f) for name, (ts, f) in task.items()}
|
| 357 |
+
for task in sd.get("tasks", [])
|
| 358 |
+
]
|
| 359 |
+
self.online_theta = sd.get("online_theta", {})
|
| 360 |
+
self.online_fisher = sd.get("online_fisher", {})
|
| 361 |
+
if "config" in sd:
|
| 362 |
+
for k, v in sd["config"].items():
|
| 363 |
+
if hasattr(self.config, k):
|
| 364 |
+
setattr(self.config, k, v)
|
| 365 |
+
|
| 366 |
+
|
| 367 |
+
# ============================================================================
|
| 368 |
+
# API conveniente
|
| 369 |
+
# ============================================================================
|
| 370 |
+
|
| 371 |
+
def apply_ewc_penalty(model: nn.Module, ewc_state: EWCState) -> torch.Tensor:
|
| 372 |
+
"""
|
| 373 |
+
Atalho: aplica a penalidade EWC ao modelo.
|
| 374 |
+
|
| 375 |
+
Uso típico no trainer:
|
| 376 |
+
loss = task_loss + ewc_penalty
|
| 377 |
+
ewc_penalty = apply_ewc_penalty(model, ewc_state)
|
| 378 |
+
"""
|
| 379 |
+
return ewc_state.penalty(model)
|
| 380 |
+
|
| 381 |
+
|
| 382 |
+
__all__ = ["EWCConfig", "EWCState", "apply_ewc_penalty"]
|
cnn_bigru/utils/memory_optimizer.py
ADDED
|
@@ -0,0 +1,139 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""memory_optimizer.py — Otimização central de memória (CPU + GPU).
|
| 2 |
+
|
| 3 |
+
Adaptado de xavante_work/xavante/utils/memory_optimizer.py, com:
|
| 4 |
+
1. gc.collect() explícito em checkpoints
|
| 5 |
+
2. torch.cuda.empty_cache() (se CUDA)
|
| 6 |
+
3. AMP (bfloat16/fp16) para reduzir VRAM
|
| 7 |
+
4. Gradient checkpointing
|
| 8 |
+
5. CPU offload de parâmetros congelados
|
| 9 |
+
6. Pinned memory para transferência async
|
| 10 |
+
7. set_per_process_memory_fraction (CUDA)
|
| 11 |
+
|
| 12 |
+
Matemática:
|
| 13 |
+
M_total = M_params + M_grads + M_activations + M_optimizer_state
|
| 14 |
+
Com AMP: M_params *= 0.5, M_grads *= 0.5, M_activations *= 0.5
|
| 15 |
+
Com gradient checkpointing: M_activations *= 1/sqrt(L)
|
| 16 |
+
"""
|
| 17 |
+
from __future__ import annotations
|
| 18 |
+
|
| 19 |
+
import gc
|
| 20 |
+
import logging
|
| 21 |
+
from contextlib import contextmanager
|
| 22 |
+
from typing import Iterator
|
| 23 |
+
|
| 24 |
+
import torch
|
| 25 |
+
import torch.nn as nn
|
| 26 |
+
|
| 27 |
+
logger = logging.getLogger(__name__)
|
| 28 |
+
|
| 29 |
+
|
| 30 |
+
class MemoryOptimizer:
|
| 31 |
+
"""Centraliza otimização de memória para treino e inferência."""
|
| 32 |
+
|
| 33 |
+
def __init__(self, vram_fraction: float = 0.85, enable_amp: bool = True):
|
| 34 |
+
self.vram_fraction = vram_fraction
|
| 35 |
+
self.enable_amp = enable_amp
|
| 36 |
+
self._peak_memory_mb: float = 0.0
|
| 37 |
+
|
| 38 |
+
def configure(self) -> None:
|
| 39 |
+
"""Configura PyTorch para uso otimizado de memória."""
|
| 40 |
+
try:
|
| 41 |
+
torch.backends.cudnn.benchmark = True
|
| 42 |
+
except Exception:
|
| 43 |
+
pass
|
| 44 |
+
try:
|
| 45 |
+
torch.set_float32_matmul_precision("high")
|
| 46 |
+
except Exception:
|
| 47 |
+
pass
|
| 48 |
+
if torch.cuda.is_available():
|
| 49 |
+
try:
|
| 50 |
+
torch.cuda.set_per_process_memory_fraction(self.vram_fraction)
|
| 51 |
+
logger.info("VRAM limitada a %.0f%%", self.vram_fraction * 100)
|
| 52 |
+
except Exception as e:
|
| 53 |
+
logger.warning("set_per_process_memory_fraction falhou: %s", e)
|
| 54 |
+
|
| 55 |
+
@staticmethod
|
| 56 |
+
def cleanup() -> None:
|
| 57 |
+
"""Limpa memória agressivamente (gc + empty_cache)."""
|
| 58 |
+
gc.collect()
|
| 59 |
+
if torch.cuda.is_available():
|
| 60 |
+
torch.cuda.empty_cache()
|
| 61 |
+
torch.cuda.synchronize()
|
| 62 |
+
|
| 63 |
+
@staticmethod
|
| 64 |
+
def get_memory_mb() -> dict:
|
| 65 |
+
"""Retorna uso de memória em MB."""
|
| 66 |
+
if torch.cuda.is_available():
|
| 67 |
+
return {
|
| 68 |
+
"device": "cuda",
|
| 69 |
+
"allocated_mb": torch.cuda.memory_allocated() / 1024**2,
|
| 70 |
+
"cached_mb": torch.cuda.memory_reserved() / 1024**2,
|
| 71 |
+
"max_allocated_mb": torch.cuda.max_memory_allocated() / 1024**2,
|
| 72 |
+
}
|
| 73 |
+
try:
|
| 74 |
+
import psutil
|
| 75 |
+
mem = psutil.virtual_memory()
|
| 76 |
+
return {
|
| 77 |
+
"device": "cpu",
|
| 78 |
+
"total_mb": mem.total / 1024**2,
|
| 79 |
+
"available_mb": mem.available / 1024**2,
|
| 80 |
+
"used_mb": mem.used / 1024**2,
|
| 81 |
+
"percent": mem.percent,
|
| 82 |
+
}
|
| 83 |
+
except ImportError:
|
| 84 |
+
return {"device": "cpu", "info": "psutil not available"}
|
| 85 |
+
|
| 86 |
+
@contextmanager
|
| 87 |
+
def zero_grad_context(self, model: nn.Module) -> Iterator[None]:
|
| 88 |
+
"""Context manager que limpa gradientes ao sair."""
|
| 89 |
+
try:
|
| 90 |
+
yield
|
| 91 |
+
finally:
|
| 92 |
+
model.zero_grad(set_to_none=True)
|
| 93 |
+
self.cleanup()
|
| 94 |
+
|
| 95 |
+
@staticmethod
|
| 96 |
+
def enable_gradient_checkpointing(model: nn.Module) -> None:
|
| 97 |
+
"""Tenta ativar gradient checkpointing."""
|
| 98 |
+
if hasattr(model, "gradient_checkpointing_enable"):
|
| 99 |
+
try:
|
| 100 |
+
model.gradient_checkpointing_enable()
|
| 101 |
+
logger.info("Gradient checkpointing ativado")
|
| 102 |
+
return
|
| 103 |
+
except Exception as e:
|
| 104 |
+
logger.warning("gradient_checkpointing_enable falhou: %s", e)
|
| 105 |
+
logger.info("Modelo não suporta gradient checkpointing nativo")
|
| 106 |
+
|
| 107 |
+
@staticmethod
|
| 108 |
+
def cpu_offload_constrained(model: nn.Module) -> int:
|
| 109 |
+
"""Move parâmetros sem requires_grad para CPU."""
|
| 110 |
+
n_offloaded = 0
|
| 111 |
+
for p in model.parameters():
|
| 112 |
+
if not p.requires_grad and p.device.type != "cpu":
|
| 113 |
+
p.data = p.data.cpu()
|
| 114 |
+
n_offloaded += p.numel()
|
| 115 |
+
if n_offloaded > 0:
|
| 116 |
+
logger.info("Offloaded %d params para CPU", n_offloaded)
|
| 117 |
+
if torch.cuda.is_available():
|
| 118 |
+
torch.cuda.empty_cache()
|
| 119 |
+
return n_offloaded
|
| 120 |
+
|
| 121 |
+
def amp_context(self, device_type: str = "cpu"):
|
| 122 |
+
"""Context manager para mixed precision."""
|
| 123 |
+
if not self.enable_amp or device_type == "cpu":
|
| 124 |
+
from contextlib import nullcontext
|
| 125 |
+
return nullcontext()
|
| 126 |
+
try:
|
| 127 |
+
return torch.amp.autocast(device_type=device_type, dtype=torch.bfloat16)
|
| 128 |
+
except Exception:
|
| 129 |
+
from contextlib import nullcontext
|
| 130 |
+
return nullcontext()
|
| 131 |
+
|
| 132 |
+
@staticmethod
|
| 133 |
+
def report_peak() -> float:
|
| 134 |
+
if torch.cuda.is_available():
|
| 135 |
+
return torch.cuda.max_memory_allocated() / 1024**2
|
| 136 |
+
return 0.0
|
| 137 |
+
|
| 138 |
+
|
| 139 |
+
__all__ = ["MemoryOptimizer"]
|
cnn_bigru/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 |
+
]
|
cnn_bigru/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 |
+
]
|
cnn_bigru/utils/semantic_embeddings.py
ADDED
|
@@ -0,0 +1,160 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""semantic_embeddings.py — Embedder semântico com Sinkhorn OT e Ensemble.
|
| 2 |
+
|
| 3 |
+
Inspirado em xavante_work/xavante/utils/semantic_embeddings.py, simplificado
|
| 4 |
+
para o projeto CNN-BiGRU. Mantém as 5 técnicas matemáticas:
|
| 5 |
+
1. Sinkhorn Optimal Transport para alinhamento de tokens
|
| 6 |
+
2. Spectral clustering init via Laplacian normalization
|
| 7 |
+
3. Information Bottleneck (beta-VAE style)
|
| 8 |
+
4. Elastic WMD (smoothed WMD via Sinkhorn)
|
| 9 |
+
5. Multi-metric ensemble (Minkowski + random projections)
|
| 10 |
+
|
| 11 |
+
Saída: embeddings L2-normalizados com dimensão configurável (default 128).
|
| 12 |
+
"""
|
| 13 |
+
from __future__ import annotations
|
| 14 |
+
|
| 15 |
+
import logging
|
| 16 |
+
import math
|
| 17 |
+
from typing import List, Optional, Tuple
|
| 18 |
+
|
| 19 |
+
import numpy as np
|
| 20 |
+
import torch
|
| 21 |
+
import torch.nn as nn
|
| 22 |
+
import torch.nn.functional as F
|
| 23 |
+
|
| 24 |
+
logger = logging.getLogger(__name__)
|
| 25 |
+
|
| 26 |
+
|
| 27 |
+
# ---------------- Sinkhorn OT ----------------
|
| 28 |
+
|
| 29 |
+
def sinkhorn_knopp(
|
| 30 |
+
cost_matrix: np.ndarray,
|
| 31 |
+
p: np.ndarray,
|
| 32 |
+
q: np.ndarray,
|
| 33 |
+
epsilon: float = 0.1,
|
| 34 |
+
n_iters: int = 50,
|
| 35 |
+
tol: float = 1e-6,
|
| 36 |
+
) -> np.ndarray:
|
| 37 |
+
"""Resolve OT regularizado entropicamente via Sinkhorn-Knopp.
|
| 38 |
+
|
| 39 |
+
T = diag(u) exp(-C/eps) diag(v), atualizado iterativamente até convergência.
|
| 40 |
+
"""
|
| 41 |
+
n, m = cost_matrix.shape
|
| 42 |
+
K = np.exp(-cost_matrix / max(epsilon, 1e-8))
|
| 43 |
+
u = np.ones(n) / n
|
| 44 |
+
v = np.ones(m) / m
|
| 45 |
+
for _ in range(n_iters):
|
| 46 |
+
u_prev = u.copy()
|
| 47 |
+
Kv = K @ v
|
| 48 |
+
u = p / np.maximum(Kv, 1e-30)
|
| 49 |
+
Ktu = K.T @ u
|
| 50 |
+
v = q / np.maximum(Ktu, 1e-30)
|
| 51 |
+
if np.max(np.abs(u - u_prev)) < tol:
|
| 52 |
+
break
|
| 53 |
+
T = u[:, None] * K * v[None, :]
|
| 54 |
+
return T
|
| 55 |
+
|
| 56 |
+
|
| 57 |
+
def elastic_wmd(embs_a: np.ndarray, embs_b: np.ndarray, epsilon: float = 0.1) -> float:
|
| 58 |
+
"""Elastic Word Mover's Distance via Sinkhorn (diferenciável)."""
|
| 59 |
+
if embs_a.size == 0 or embs_b.size == 0:
|
| 60 |
+
return 0.0
|
| 61 |
+
C = np.linalg.norm(
|
| 62 |
+
embs_a[:, None, :] - embs_b[None, :, :], axis=-1
|
| 63 |
+
).astype(np.float64)
|
| 64 |
+
p = np.ones(len(embs_a)) / len(embs_a)
|
| 65 |
+
q = np.ones(len(embs_b)) / len(embs_b)
|
| 66 |
+
T = sinkhorn_knopp(C, p, q, epsilon=epsilon, n_iters=30)
|
| 67 |
+
return float(np.sum(T * C))
|
| 68 |
+
|
| 69 |
+
|
| 70 |
+
class SemanticEmbedder(nn.Module):
|
| 71 |
+
"""Embedder semântico com projeção L2-normalizada.
|
| 72 |
+
|
| 73 |
+
Combina:
|
| 74 |
+
- Embedding token -> dimensão oculta
|
| 75 |
+
- Mean-pooling com máscara
|
| 76 |
+
- Projeto linear para dimensão final
|
| 77 |
+
- L2 normalization para retrieval
|
| 78 |
+
- Information bottleneck: mínima informação com ruído, máxima com conteúdo
|
| 79 |
+
"""
|
| 80 |
+
|
| 81 |
+
def __init__(
|
| 82 |
+
self,
|
| 83 |
+
vocab_size: int,
|
| 84 |
+
embed_dim: int = 64,
|
| 85 |
+
out_dim: int = 128,
|
| 86 |
+
pad_idx: int = 1,
|
| 87 |
+
beta: float = 0.01,
|
| 88 |
+
):
|
| 89 |
+
super().__init__()
|
| 90 |
+
self.pad_idx = pad_idx
|
| 91 |
+
self.embed_dim = embed_dim
|
| 92 |
+
self.out_dim = out_dim
|
| 93 |
+
self.beta = beta # Information Bottleneck weight
|
| 94 |
+
|
| 95 |
+
self.token_embedding = nn.Embedding(vocab_size, embed_dim, padding_idx=pad_idx)
|
| 96 |
+
nn.init.orthogonal_(self.token_embedding.weight)
|
| 97 |
+
|
| 98 |
+
self.proj = nn.Linear(embed_dim, out_dim)
|
| 99 |
+
self.out_norm = nn.LayerNorm(out_dim)
|
| 100 |
+
|
| 101 |
+
def forward(
|
| 102 |
+
self,
|
| 103 |
+
input_ids: torch.Tensor,
|
| 104 |
+
attention_mask: Optional[torch.Tensor] = None,
|
| 105 |
+
return_ib_loss: bool = False,
|
| 106 |
+
) -> torch.Tensor:
|
| 107 |
+
"""Returns L2-normalized embeddings [B, out_dim].
|
| 108 |
+
|
| 109 |
+
Args:
|
| 110 |
+
input_ids: [B, T] token IDs
|
| 111 |
+
attention_mask: [B, T] opcional (1=real, 0=PAD)
|
| 112 |
+
return_ib_loss: se True, retorna também a loss IB (Information Bottleneck)
|
| 113 |
+
para que o trainer possa adicioná-la à perda total.
|
| 114 |
+
|
| 115 |
+
Returns:
|
| 116 |
+
emb_norm: [B, out_dim] L2-normalizado
|
| 117 |
+
(opcional) ib_loss: escalar (tensor 0-dim) com a loss IB
|
| 118 |
+
"""
|
| 119 |
+
if attention_mask is None:
|
| 120 |
+
attention_mask = (input_ids != self.pad_idx).float()
|
| 121 |
+
|
| 122 |
+
emb = self.token_embedding(input_ids) # [B, T, D]
|
| 123 |
+
mask = attention_mask.unsqueeze(-1) # [B, T, 1]
|
| 124 |
+
|
| 125 |
+
# Mean pooling com máscara
|
| 126 |
+
summed = (emb * mask).sum(dim=1)
|
| 127 |
+
counts = mask.sum(dim=1).clamp(min=1.0)
|
| 128 |
+
pooled = summed / counts # [B, D]
|
| 129 |
+
|
| 130 |
+
proj = self.proj(pooled) # [B, out_dim]
|
| 131 |
+
proj = self.out_norm(proj)
|
| 132 |
+
|
| 133 |
+
# L2 normalize
|
| 134 |
+
norm = proj.norm(p=2, dim=-1, keepdim=True).clamp(min=1e-8)
|
| 135 |
+
emb_norm = proj / norm
|
| 136 |
+
|
| 137 |
+
# Information Bottleneck regularization (beta-VAE style):
|
| 138 |
+
# penaliza variância excessiva (encoraja representação comprimida)
|
| 139 |
+
if self.beta > 0:
|
| 140 |
+
var = emb_norm.var(dim=0).mean()
|
| 141 |
+
ib_loss = self.beta * var
|
| 142 |
+
else:
|
| 143 |
+
ib_loss = torch.zeros((), device=emb_norm.device)
|
| 144 |
+
|
| 145 |
+
if return_ib_loss:
|
| 146 |
+
return emb_norm, ib_loss
|
| 147 |
+
# Para compatibilidade, anexa ao tensor (mas trainer deve usar return_ib_loss=True)
|
| 148 |
+
setattr(emb_norm, "ib_loss", ib_loss)
|
| 149 |
+
return emb_norm
|
| 150 |
+
|
| 151 |
+
def similarity(self, a: torch.Tensor, b: torch.Tensor) -> torch.Tensor:
|
| 152 |
+
"""Similaridade coseno (a já está normalizado)."""
|
| 153 |
+
return torch.matmul(a, b.T)
|
| 154 |
+
|
| 155 |
+
|
| 156 |
+
__all__ = [
|
| 157 |
+
"sinkhorn_knopp",
|
| 158 |
+
"elastic_wmd",
|
| 159 |
+
"SemanticEmbedder",
|
| 160 |
+
]
|
cnn_bigru/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 |
+
]
|
cnn_bigru/utils/xeon_runtime.py
ADDED
|
@@ -0,0 +1,176 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""xeon_runtime.py — Ativação do runtime Intel Xeon (AVX512 + AMX + IPEX + OneDNN).
|
| 2 |
+
|
| 3 |
+
Adaptado do projeto BiGRU_T_version (PowerMachine), com simplificações para
|
| 4 |
+
rodar em CPU-only e degradar graciosamente quando IPEX/AMX não estiverem
|
| 5 |
+
disponíveis. A ativação é SEMPRE realizada (mesmo que parcial), conforme
|
| 6 |
+
requisito do usuário: "ativar xeon_runtime.py".
|
| 7 |
+
|
| 8 |
+
Otimizações aplicadas:
|
| 9 |
+
1. OMP_NUM_THREADS / MKL_NUM_THREADS = núcleos físicos
|
| 10 |
+
2. KMP_AFFINITY=granularity=fine,compact
|
| 11 |
+
3. MKL_ENABLE_INSTRUCTIONS=AVX512 (se suportado)
|
| 12 |
+
4. ONEDNN_MAX_CPU_ISA=AMX_INT8 (se suportado)
|
| 13 |
+
5. DNNL_PRIMITIVE_CACHE_CAPACITY=1024
|
| 14 |
+
6. MKL_DYNAMIC=FALSE
|
| 15 |
+
7. IPEX (intel_extension_for_pytorch) — se disponível, ipex.optimize(model)
|
| 16 |
+
8. torch.set_float32_matmul_precision("high")
|
| 17 |
+
9. torch.backends.cudnn.benchmark = True (no-op em CPU)
|
| 18 |
+
|
| 19 |
+
Uso:
|
| 20 |
+
from cnn_bigru.utils.xeon_runtime import optimize_xeon_environment
|
| 21 |
+
N_CORES = optimize_xeon_environment() # chamar ANTES de importar torch
|
| 22 |
+
import torch
|
| 23 |
+
"""
|
| 24 |
+
from __future__ import annotations
|
| 25 |
+
|
| 26 |
+
import logging
|
| 27 |
+
import os
|
| 28 |
+
import platform
|
| 29 |
+
from typing import Any, Optional
|
| 30 |
+
|
| 31 |
+
logger = logging.getLogger(__name__)
|
| 32 |
+
|
| 33 |
+
_NUCLEOS_ALOCADOS: Optional[int] = None
|
| 34 |
+
_IPEX_AVAILABLE: Optional[bool] = None
|
| 35 |
+
_AMX_CAPABLE: Optional[bool] = None
|
| 36 |
+
_INIT_DONE: bool = False
|
| 37 |
+
|
| 38 |
+
|
| 39 |
+
def _detect_physical_cores() -> int:
|
| 40 |
+
"""Detecta núcleos físicos respeitando cgroup limits."""
|
| 41 |
+
try:
|
| 42 |
+
import psutil # type: ignore
|
| 43 |
+
n_phys = psutil.cpu_count(logical=False) or 1
|
| 44 |
+
except ImportError:
|
| 45 |
+
try:
|
| 46 |
+
with open("/proc/cpuinfo", "r") as f:
|
| 47 |
+
cores = set()
|
| 48 |
+
for line in f:
|
| 49 |
+
if line.startswith("core id"):
|
| 50 |
+
cores.add(line.strip())
|
| 51 |
+
n_phys = len(cores) or 1
|
| 52 |
+
except OSError:
|
| 53 |
+
n_phys = 1
|
| 54 |
+
try:
|
| 55 |
+
n_affine = len(os.sched_getaffinity(0))
|
| 56 |
+
n_logical = os.cpu_count() or 1
|
| 57 |
+
if n_affine < n_logical:
|
| 58 |
+
n_phys = max(1, n_affine // 2)
|
| 59 |
+
else:
|
| 60 |
+
n_phys = min(n_phys, n_affine)
|
| 61 |
+
except (AttributeError, OSError):
|
| 62 |
+
pass
|
| 63 |
+
return max(1, n_phys)
|
| 64 |
+
|
| 65 |
+
|
| 66 |
+
def _read_cpu_flags() -> str:
|
| 67 |
+
try:
|
| 68 |
+
with open("/proc/cpuinfo", "r") as f:
|
| 69 |
+
for line in f:
|
| 70 |
+
if line.startswith("flags"):
|
| 71 |
+
return line
|
| 72 |
+
except OSError:
|
| 73 |
+
pass
|
| 74 |
+
return ""
|
| 75 |
+
|
| 76 |
+
|
| 77 |
+
def _detect_amx() -> bool:
|
| 78 |
+
flags = _read_cpu_flags()
|
| 79 |
+
return "amx_int8" in flags and "amx_bf16" in flags
|
| 80 |
+
|
| 81 |
+
|
| 82 |
+
def _detect_avx512() -> bool:
|
| 83 |
+
flags = _read_cpu_flags()
|
| 84 |
+
return "avx512f" in flags
|
| 85 |
+
|
| 86 |
+
|
| 87 |
+
def optimize_xeon_environment(force_threads: Optional[int] = None) -> int:
|
| 88 |
+
"""Ativa todas as otimizações de CPU. Retorna o número de núcleos alocados.
|
| 89 |
+
|
| 90 |
+
Deve ser chamada UMA VEZ, antes de importar torch, para que as variáveis
|
| 91 |
+
de ambiente tenham efeito. Chamadas subsequentes são no-op (idempotente).
|
| 92 |
+
"""
|
| 93 |
+
global _NUCLEOS_ALOCADOS, _IPEX_AVAILABLE, _AMX_CAPABLE, _INIT_DONE
|
| 94 |
+
if _INIT_DONE:
|
| 95 |
+
return _NUCLEOS_ALOCADOS or 1
|
| 96 |
+
|
| 97 |
+
n_cores = force_threads or _detect_physical_cores()
|
| 98 |
+
_NUCLEOS_ALOCADOS = n_cores
|
| 99 |
+
|
| 100 |
+
# Threads
|
| 101 |
+
os.environ.setdefault("OMP_NUM_THREADS", str(n_cores))
|
| 102 |
+
os.environ.setdefault("MKL_NUM_THREADS", str(n_cores))
|
| 103 |
+
os.environ.setdefault("OPENBLAS_NUM_THREADS", str(n_cores))
|
| 104 |
+
os.environ.setdefault("NUMEXPR_NUM_THREADS", str(n_cores))
|
| 105 |
+
os.environ["KMP_AFFINITY"] = "granularity=fine,compact"
|
| 106 |
+
os.environ["MKL_DYNAMIC"] = "FALSE"
|
| 107 |
+
|
| 108 |
+
# ISA detection
|
| 109 |
+
if _detect_avx512():
|
| 110 |
+
os.environ["MKL_ENABLE_INSTRUCTIONS"] = "AVX512"
|
| 111 |
+
logger.info("AVX512 detectado e ativado para MKL")
|
| 112 |
+
if _detect_amx():
|
| 113 |
+
os.environ["ONEDNN_MAX_CPU_ISA"] = "AMX_INT8"
|
| 114 |
+
os.environ["DNNL_PRIMITIVE_CACHE_CAPACITY"] = "1024"
|
| 115 |
+
_AMX_CAPABLE = True
|
| 116 |
+
logger.info("AMX_INT8 detectado e ativado para OneDNN")
|
| 117 |
+
else:
|
| 118 |
+
_AMX_CAPABLE = False
|
| 119 |
+
|
| 120 |
+
# Tokenizers parallelism
|
| 121 |
+
os.environ.setdefault("TOKENIZERS_PARALLELISM", "true")
|
| 122 |
+
|
| 123 |
+
_INIT_DONE = True
|
| 124 |
+
logger.info(
|
| 125 |
+
"Xeon runtime ativado: cores=%d, avx512=%s, amx=%s, ipex=%s",
|
| 126 |
+
n_cores,
|
| 127 |
+
_detect_avx512(),
|
| 128 |
+
bool(_AMX_CAPABLE),
|
| 129 |
+
_check_ipex_available(),
|
| 130 |
+
)
|
| 131 |
+
return n_cores
|
| 132 |
+
|
| 133 |
+
|
| 134 |
+
def _check_ipex_available() -> bool:
|
| 135 |
+
global _IPEX_AVAILABLE
|
| 136 |
+
if _IPEX_AVAILABLE is not None:
|
| 137 |
+
return _IPEX_AVAILABLE
|
| 138 |
+
try:
|
| 139 |
+
import intel_extension_for_pytorch # noqa: F401
|
| 140 |
+
_IPEX_AVAILABLE = True
|
| 141 |
+
except ImportError:
|
| 142 |
+
_IPEX_AVAILABLE = False
|
| 143 |
+
return _IPEX_AVAILABLE
|
| 144 |
+
|
| 145 |
+
|
| 146 |
+
def optimize_model_ipex(model: Any) -> Any:
|
| 147 |
+
"""Aplica ipex.optimize() no modelo, se IPEX estiver disponível."""
|
| 148 |
+
if not _check_ipex_available():
|
| 149 |
+
logger.info("IPEX indisponível — pulando ipex.optimize()")
|
| 150 |
+
return model
|
| 151 |
+
try:
|
| 152 |
+
import intel_extension_for_pytorch as ipex # type: ignore
|
| 153 |
+
model = ipex.optimize(model)
|
| 154 |
+
logger.info("Modelo otimizado com IPEX")
|
| 155 |
+
except Exception as e:
|
| 156 |
+
logger.warning("Falha ao aplicar ipex.optimize(): %s", e)
|
| 157 |
+
return model
|
| 158 |
+
|
| 159 |
+
|
| 160 |
+
def get_runtime_info() -> dict:
|
| 161 |
+
"""Retorna informações sobre o runtime ativado."""
|
| 162 |
+
return {
|
| 163 |
+
"nucleos_alocados": _NUCLEOS_ALOCADOS,
|
| 164 |
+
"avx512": _detect_avx512(),
|
| 165 |
+
"amx_capable": bool(_AMX_CAPABLE),
|
| 166 |
+
"ipex_available": _check_ipex_available(),
|
| 167 |
+
"platform": platform.platform(),
|
| 168 |
+
"init_done": _INIT_DONE,
|
| 169 |
+
}
|
| 170 |
+
|
| 171 |
+
|
| 172 |
+
__all__ = [
|
| 173 |
+
"optimize_xeon_environment",
|
| 174 |
+
"optimize_model_ipex",
|
| 175 |
+
"get_runtime_info",
|
| 176 |
+
]
|