V6.5-V2-dynamic: upload batch (scripts + state + model) [16 files]
Browse files- attention_multimodal.py +114 -0
- hyp_t.py +100 -0
- kohonen_learning_system.py +0 -0
- som_metrics.py +920 -0
- streaming_datasets.py +995 -0
- train_v6_5_v2.py +0 -0
- v6_5_v2_attention_eval.json +5 -5
- v6_5_v2_model_states.pt +3 -0
- v6_5_v2_model_states_after_conhecimento.pt +3 -0
- v6_5_v2_phases_eval.json +0 -0
- v6_5_v2_report.json +35 -30
- v6_5_v2_user_questions.json +7 -7
- vqvae2_hierarchical.py +200 -0
- vqvae2_hierarchical_flexnet.py +513 -0
- xeon_runtime.py +581 -0
attention_multimodal.py
ADDED
|
@@ -0,0 +1,114 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
Xavante - attention_multimodal.py
|
| 3 |
+
Responsabilidade: Multi-head attention multimodal (Teorema 11.1/11.2).
|
| 4 |
+
Suporta attention entre modalidades e dentro de modalidade.
|
| 5 |
+
"""
|
| 6 |
+
from __future__ import annotations
|
| 7 |
+
|
| 8 |
+
import logging
|
| 9 |
+
import math
|
| 10 |
+
from typing import Optional
|
| 11 |
+
|
| 12 |
+
import torch
|
| 13 |
+
import torch.nn as nn
|
| 14 |
+
import torch.nn.functional as F
|
| 15 |
+
|
| 16 |
+
logger = logging.getLogger(__name__)
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
class MultiHeadAttention(nn.Module):
|
| 20 |
+
"""
|
| 21 |
+
MHA padrão com suporte a:
|
| 22 |
+
- Mascara causal
|
| 23 |
+
- Mascara multimodal (mod crossing)
|
| 24 |
+
- Attention 2M tokens via chunked attention (Teorema 6.1)
|
| 25 |
+
"""
|
| 26 |
+
|
| 27 |
+
def __init__(
|
| 28 |
+
self,
|
| 29 |
+
d_model: int,
|
| 30 |
+
n_heads: int = 8,
|
| 31 |
+
dropout: float = 0.0,
|
| 32 |
+
max_chunk: int = 4096,
|
| 33 |
+
):
|
| 34 |
+
super().__init__()
|
| 35 |
+
assert d_model % n_heads == 0
|
| 36 |
+
self.d_model = d_model
|
| 37 |
+
self.n_heads = n_heads
|
| 38 |
+
self.d_head = d_model // n_heads
|
| 39 |
+
self.max_chunk = max_chunk
|
| 40 |
+
self.qkv = nn.Linear(d_model, 3 * d_model, bias=True)
|
| 41 |
+
self.out = nn.Linear(d_model, d_model)
|
| 42 |
+
self.dropout = nn.Dropout(dropout)
|
| 43 |
+
|
| 44 |
+
def forward(
|
| 45 |
+
self,
|
| 46 |
+
x: torch.Tensor,
|
| 47 |
+
mask: Optional[torch.Tensor] = None,
|
| 48 |
+
kv: Optional[torch.Tensor] = None,
|
| 49 |
+
) -> torch.Tensor:
|
| 50 |
+
B, L, D = x.shape
|
| 51 |
+
if kv is None:
|
| 52 |
+
qkv = self.qkv(x)
|
| 53 |
+
q, k, v = qkv.chunk(3, dim=-1)
|
| 54 |
+
else:
|
| 55 |
+
q = self.qkv(x)[:, :, :D]
|
| 56 |
+
kv_proj = self.qkv(kv)
|
| 57 |
+
k = kv_proj[:, :, D : 2 * D]
|
| 58 |
+
v = kv_proj[:, :, 2 * D :]
|
| 59 |
+
# reshape para heads
|
| 60 |
+
q = q.view(B, L, self.n_heads, self.d_head).transpose(1, 2)
|
| 61 |
+
k = k.view(B, -1, self.n_heads, self.d_head).transpose(1, 2)
|
| 62 |
+
v = v.view(B, -1, self.n_heads, self.d_head).transpose(1, 2)
|
| 63 |
+
|
| 64 |
+
# Chunked attention para sequencias longas (Teorema 6.1)
|
| 65 |
+
if L > self.max_chunk:
|
| 66 |
+
# Normalize mask to 4D for chunked path
|
| 67 |
+
if mask is not None and mask.dim() == 2:
|
| 68 |
+
mask = mask.unsqueeze(0).unsqueeze(0)
|
| 69 |
+
elif mask is not None and mask.dim() == 3:
|
| 70 |
+
mask = mask.unsqueeze(1)
|
| 71 |
+
return self._chunked_attention(q, k, v, mask)
|
| 72 |
+
|
| 73 |
+
scale = 1.0 / math.sqrt(self.d_head)
|
| 74 |
+
attn = (q @ k.transpose(-2, -1)) * scale # [B, H, L, L]
|
| 75 |
+
if mask is not None:
|
| 76 |
+
attn = attn.masked_fill(mask == 0, float("-inf"))
|
| 77 |
+
attn = F.softmax(attn, dim=-1)
|
| 78 |
+
attn = self.dropout(attn)
|
| 79 |
+
out = attn @ v # [B, H, L, d_head]
|
| 80 |
+
out = out.transpose(1, 2).contiguous().view(B, L, D)
|
| 81 |
+
return self.out(out)
|
| 82 |
+
|
| 83 |
+
def _chunked_attention(
|
| 84 |
+
self,
|
| 85 |
+
q: torch.Tensor,
|
| 86 |
+
k: torch.Tensor,
|
| 87 |
+
v: torch.Tensor,
|
| 88 |
+
mask: Optional[torch.Tensor],
|
| 89 |
+
) -> torch.Tensor:
|
| 90 |
+
"""Attention em blocos para janelas de 2M tokens (memoria controlada)."""
|
| 91 |
+
B, H, L, d = q.shape
|
| 92 |
+
Lk = k.shape[2]
|
| 93 |
+
chunk = self.max_chunk
|
| 94 |
+
outs = []
|
| 95 |
+
scale = 1.0 / math.sqrt(d)
|
| 96 |
+
for i in range(0, L, chunk):
|
| 97 |
+
qi = q[:, :, i : i + chunk]
|
| 98 |
+
out_chunk = []
|
| 99 |
+
for j in range(0, Lk, chunk):
|
| 100 |
+
kj = k[:, :, j : j + chunk]
|
| 101 |
+
vj = v[:, :, j : j + chunk]
|
| 102 |
+
attn = (qi @ kj.transpose(-2, -1)) * scale
|
| 103 |
+
if mask is not None:
|
| 104 |
+
m_chunk = mask[:, :, i : i + chunk, j : j + chunk]
|
| 105 |
+
attn = attn.masked_fill(m_chunk == 0, float("-inf"))
|
| 106 |
+
attn = F.softmax(attn, dim=-1)
|
| 107 |
+
out_chunk.append(attn @ vj)
|
| 108 |
+
outs.append(torch.cat(out_chunk, dim=2))
|
| 109 |
+
out = torch.cat(outs, dim=2)
|
| 110 |
+
out = out.transpose(1, 2).contiguous().view(B, L, self.n_heads * d)
|
| 111 |
+
return self.out(out)
|
| 112 |
+
|
| 113 |
+
|
| 114 |
+
__all__ = ["MultiHeadAttention"]
|
hyp_t.py
ADDED
|
@@ -0,0 +1,100 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""hyp_t.py — HypT: cabeça de hipótese (correção delta). Implementa o Lema 3.
|
| 2 |
+
|
| 3 |
+
A Hipótese é treinada para cancelar o ruído de quantização W8A8:
|
| 4 |
+
y_final = y_hat + delta
|
| 5 |
+
onde delta = hyp_T(o).
|
| 6 |
+
|
| 7 |
+
Diferenças vs. TrainT:
|
| 8 |
+
- Recebe `o` opcionalmente destacado (stop_grad) para isolar o sinal de
|
| 9 |
+
correção do gradiente principal
|
| 10 |
+
- Mesma arquitetura (TransformerEncoder) — apenas o tratamento de gradiente
|
| 11 |
+
difere
|
| 12 |
+
|
| 13 |
+
Reaproveita a filosofia do HypothesisNetwork do v13.9.2 (16 hipóteses com
|
| 14 |
+
mapeamento LoRA), mas simplificada: 1 única hipótese global (delta) em vez
|
| 15 |
+
de 16 hipóteses paralelas. A multiplexação é feita pelo OrqCell + Lema 1.
|
| 16 |
+
|
| 17 |
+
============================================================================
|
| 18 |
+
V6.4 — pgvector_lookup REMOVIDO (NÃO É MAIS NECESSÁRIO)
|
| 19 |
+
============================================================================
|
| 20 |
+
Conforme directive do usuário:
|
| 21 |
+
"pgvector_lookup não é mais necessário pela lógica do script seguinte"
|
| 22 |
+
|
| 23 |
+
O KohonenLearningSystem (kohonen_refactored/kohonen_learning_system.py)
|
| 24 |
+
contém o método find_bmu que realiza busca nearest-neighbor sobre o grid
|
| 25 |
+
4D do SOM (864 neurônios para grid (6,6,6,4)), substituindo qualquer
|
| 26 |
+
lookup pgvector externo. A decisão de aplicar punição é delegada ao
|
| 27 |
+
protocolo interno do KohonenLearningSystem (1ª punição → activate_hypothesis,
|
| 28 |
+
2ª punição → set_ewc_reference + reset).
|
| 29 |
+
|
| 30 |
+
Consequentemente, HypT volta ao seu papel original: gerar a correção
|
| 31 |
+
delta = hyp_T(o) sem consultar mapa externo. A integração com o SOM é
|
| 32 |
+
feita no nível do trainer (KohonenLearningSystem.process_batch), não no
|
| 33 |
+
nível da hipótese.
|
| 34 |
+
|
| 35 |
+
Os contadores de punição (n_complete_skip, n_incomplete_apply, etc.)
|
| 36 |
+
foram removidos pois não há mais decisão de skip baseada em pgvector.
|
| 37 |
+
"""
|
| 38 |
+
from __future__ import annotations
|
| 39 |
+
|
| 40 |
+
import torch
|
| 41 |
+
import torch.nn as nn
|
| 42 |
+
|
| 43 |
+
|
| 44 |
+
class HypT(nn.Module):
|
| 45 |
+
"""Transformer de hipótese: gera correção delta.
|
| 46 |
+
|
| 47 |
+
Args:
|
| 48 |
+
d_input: dimensão de entrada (saída do OrqCell = d_cache)
|
| 49 |
+
d_model: dimensão interna
|
| 50 |
+
nhead: nº de cabeças
|
| 51 |
+
d_ff: dimensão do FFN
|
| 52 |
+
output_dim: dimensão de saída (igual ao TrainT)
|
| 53 |
+
num_layers: nº de camadas
|
| 54 |
+
dropout: probabilidade de dropout
|
| 55 |
+
|
| 56 |
+
Forward:
|
| 57 |
+
o: (batch, d_input), stop_grad: bool → delta: (batch, output_dim)
|
| 58 |
+
|
| 59 |
+
Se stop_grad=True, aplica o.detach() no início — Lema 3 (isolar correção
|
| 60 |
+
do gradiente principal).
|
| 61 |
+
"""
|
| 62 |
+
|
| 63 |
+
def __init__(
|
| 64 |
+
self,
|
| 65 |
+
d_input: int,
|
| 66 |
+
d_model: int,
|
| 67 |
+
nhead: int,
|
| 68 |
+
d_ff: int,
|
| 69 |
+
output_dim: int,
|
| 70 |
+
num_layers: int = 2,
|
| 71 |
+
dropout: float = 0.1,
|
| 72 |
+
):
|
| 73 |
+
super().__init__()
|
| 74 |
+
self.d_input = d_input
|
| 75 |
+
self.d_model = d_model
|
| 76 |
+
self.output_dim = output_dim
|
| 77 |
+
|
| 78 |
+
self.proj_in = nn.Linear(d_input, d_model)
|
| 79 |
+
encoder_layer = nn.TransformerEncoderLayer(
|
| 80 |
+
d_model, nhead, d_ff, dropout=dropout, batch_first=True
|
| 81 |
+
)
|
| 82 |
+
self.encoder = nn.TransformerEncoder(encoder_layer, num_layers)
|
| 83 |
+
self.proj_out = nn.Linear(d_model, output_dim)
|
| 84 |
+
|
| 85 |
+
def forward(self, o: torch.Tensor, stop_grad: bool = True) -> torch.Tensor:
|
| 86 |
+
"""o: (batch, d_input) → delta: (batch, output_dim)
|
| 87 |
+
|
| 88 |
+
V6.4: SEM consulta pgvector — delta é computado diretamente.
|
| 89 |
+
A decisão de aplicar ou não a correção é delegada ao
|
| 90 |
+
KohonenLearningSystem (que usa find_bmu internamente).
|
| 91 |
+
"""
|
| 92 |
+
# Lema 3: stop_grad_hyp isola o sinal de correção
|
| 93 |
+
if stop_grad:
|
| 94 |
+
o = o.detach()
|
| 95 |
+
|
| 96 |
+
x = o.unsqueeze(1) # (batch, 1, d_input)
|
| 97 |
+
x = self.proj_in(x) # (batch, 1, d_model)
|
| 98 |
+
x = self.encoder(x) # (batch, 1, d_model)
|
| 99 |
+
delta = self.proj_out(x.squeeze(1)) # (batch, output_dim)
|
| 100 |
+
return delta
|
kohonen_learning_system.py
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
som_metrics.py
ADDED
|
@@ -0,0 +1,920 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""som_metrics.py — V6.5-V2-metrics: Métricas de aprendizado e indicadores de
|
| 2 |
+
falha para Mapas Auto-Organizáveis (SOM / Redes de Kohonen).
|
| 3 |
+
|
| 4 |
+
============================================================================
|
| 5 |
+
MÉTRICAS PRINCIPAIS DE APRENDIZADO
|
| 6 |
+
============================================================================
|
| 7 |
+
|
| 8 |
+
1.1. Erro de Quantização (Quantization Error - QE)
|
| 9 |
+
Média das distâncias euclidianas entre cada vetor de entrada e o vetor
|
| 10 |
+
de pesos de seu neurônio vencedor (BMU). Mede a resolução ou ajuste fino
|
| 11 |
+
do mapa aos dados.
|
| 12 |
+
|
| 13 |
+
QE = (1/N) * Σ_i ||x_i - W_BMU(x_i)||₂
|
| 14 |
+
|
| 15 |
+
1.2. Erro Topológico (Topological Error - TE)
|
| 16 |
+
Percentual de vetores de entrada cujos primeiro e segundo neurônios
|
| 17 |
+
mais próximos (BMU e segunda BMU) NÃO são vizinhos diretos na grade
|
| 18 |
+
de saída. Avalia preservação das relações geométricas.
|
| 19 |
+
|
| 20 |
+
TE = (1/N) * Σ_i 𝟙[¬adjacent(BMU_1(x_i), BMU_2(x_i))]
|
| 21 |
+
|
| 22 |
+
1.3. Erro de Kaski-Lagus
|
| 23 |
+
Combinação ponderada que penaliza tanto a falta de resolução do vetor
|
| 24 |
+
de quantização quanto as quebras na continuidade topológica.
|
| 25 |
+
|
| 26 |
+
KL = α * QE_norm + β * TE
|
| 27 |
+
onde α=0.5, β=0.5 (pesos padrão)
|
| 28 |
+
|
| 29 |
+
Variante implementada: KL = QE * (1 + TE) (penaliza QE quando TE é alto)
|
| 30 |
+
|
| 31 |
+
1.4. Variância Explicada (Share of Explained Variance)
|
| 32 |
+
Fração da variabilidade total dos dados de entrada que é mapeada e
|
| 33 |
+
representada pelos vetores de referência (pesos) dos nós da rede.
|
| 34 |
+
|
| 35 |
+
VarExplained = 1 - Var(residual) / Var(total)
|
| 36 |
+
onde residual = x_i - W_BMU(x_i)
|
| 37 |
+
|
| 38 |
+
============================================================================
|
| 39 |
+
INDICADORES DE NÃO APRENDIZADO OU FALHA
|
| 40 |
+
============================================================================
|
| 41 |
+
|
| 42 |
+
2.1. Colapso Topológico (Topological Collapse)
|
| 43 |
+
Ocorre quando todos os vetores de peso convergem para o mesmo ponto
|
| 44 |
+
ou para uma linha reta. Indica perda total da dimensionalidade útil.
|
| 45 |
+
|
| 46 |
+
Detecção:
|
| 47 |
+
- std(weights) por dimensão < threshold (colapso total)
|
| 48 |
+
- rank(weights reshaped) < 2 (todos num ponto)
|
| 49 |
+
- rank < 3 (todos numa linha)
|
| 50 |
+
|
| 51 |
+
2.2. Taxa de Neurônios Mortos (Dead Neurons)
|
| 52 |
+
Percentual de neurônios que NUNCA se tornam BMU para nenhum vetor.
|
| 53 |
+
|
| 54 |
+
DeadRate = (n_neurons - n_active_neurons) / n_neurons
|
| 55 |
+
|
| 56 |
+
Threshold: DeadRate > 0.5 indica problema.
|
| 57 |
+
|
| 58 |
+
2.3. Estagnação do Erro de Quantização
|
| 59 |
+
QE estabiliza em patamar elevado ao longo das épocas. A taxa de
|
| 60 |
+
aprendizado decaiu rápido demais antes da fase de ordenação global.
|
| 61 |
+
|
| 62 |
+
Detecção: slope(QE) ≈ 0 E QE atual > 0.7 * max(QE histórico)
|
| 63 |
+
|
| 64 |
+
2.4. Cruzamento de Vizinhança (Neighborhood Crossing)
|
| 65 |
+
Crescimento repentino ou oscilação crônica do TE nas épocas finais.
|
| 66 |
+
|
| 67 |
+
Detecção:
|
| 68 |
+
- std(TE_final_5_epochs) > 2 * std(TE_initial_5_epochs) (oscilação)
|
| 69 |
+
- TE[-1] > TE[-5] (crescimento no final)
|
| 70 |
+
|
| 71 |
+
============================================================================
|
| 72 |
+
AUTOR: V6.5-V2-metrics
|
| 73 |
+
============================================================================
|
| 74 |
+
"""
|
| 75 |
+
from __future__ import annotations
|
| 76 |
+
|
| 77 |
+
import math
|
| 78 |
+
import logging
|
| 79 |
+
from collections import defaultdict, deque
|
| 80 |
+
from dataclasses import dataclass, field
|
| 81 |
+
from typing import Any, Dict, List, Optional, Tuple
|
| 82 |
+
|
| 83 |
+
import torch
|
| 84 |
+
import torch.nn.functional as F
|
| 85 |
+
|
| 86 |
+
logger = logging.getLogger(__name__)
|
| 87 |
+
|
| 88 |
+
|
| 89 |
+
# ============================================================================
|
| 90 |
+
# Dataclass para armazenar histórico de métricas
|
| 91 |
+
# ============================================================================
|
| 92 |
+
@dataclass
|
| 93 |
+
class SOMMetricHistory:
|
| 94 |
+
"""Histórico deslizante de métricas SOM para detecção de estagnação e
|
| 95 |
+
cruzamento de vizinhança."""
|
| 96 |
+
qe_history: deque = field(default_factory=lambda: deque(maxlen=50))
|
| 97 |
+
te_history: deque = field(default_factory=lambda: deque(maxlen=50))
|
| 98 |
+
kl_history: deque = field(default_factory=lambda: deque(maxlen=50))
|
| 99 |
+
var_explained_history: deque = field(default_factory=lambda: deque(maxlen=50))
|
| 100 |
+
dead_neuron_rate_history: deque = field(default_factory=lambda: deque(maxlen=50))
|
| 101 |
+
epoch: int = 0
|
| 102 |
+
|
| 103 |
+
def record(self, metrics: Dict[str, Any]) -> None:
|
| 104 |
+
"""Registra métricas de uma época.
|
| 105 |
+
|
| 106 |
+
Nota: `dead_neuron_rate` pode ser float (valor direto) ou Dict
|
| 107 |
+
(saída de compute_all_metrics, com chave 'dead_neuron_rate').
|
| 108 |
+
"""
|
| 109 |
+
self.qe_history.append(float(metrics.get("quantization_error", 0.0)))
|
| 110 |
+
self.te_history.append(float(metrics.get("topological_error", 0.0)))
|
| 111 |
+
self.kl_history.append(float(metrics.get("kaski_lagus_error", 0.0)))
|
| 112 |
+
self.var_explained_history.append(
|
| 113 |
+
float(metrics.get("explained_variance_share", 0.0))
|
| 114 |
+
)
|
| 115 |
+
# dead_neuron_rate pode ser float ou Dict
|
| 116 |
+
dnr = metrics.get("dead_neuron_rate", 0.0)
|
| 117 |
+
if isinstance(dnr, dict):
|
| 118 |
+
dnr_val = float(dnr.get("dead_neuron_rate", 0.0))
|
| 119 |
+
else:
|
| 120 |
+
dnr_val = float(dnr)
|
| 121 |
+
self.dead_neuron_rate_history.append(dnr_val)
|
| 122 |
+
self.epoch += 1
|
| 123 |
+
|
| 124 |
+
def to_dict(self) -> Dict[str, Any]:
|
| 125 |
+
return {
|
| 126 |
+
"epoch": int(self.epoch),
|
| 127 |
+
"qe_history": list(self.qe_history),
|
| 128 |
+
"te_history": list(self.te_history),
|
| 129 |
+
"kl_history": list(self.kl_history),
|
| 130 |
+
"var_explained_history": list(self.var_explained_history),
|
| 131 |
+
"dead_neuron_rate_history": list(self.dead_neuron_rate_history),
|
| 132 |
+
}
|
| 133 |
+
|
| 134 |
+
|
| 135 |
+
# ============================================================================
|
| 136 |
+
# Funções auxiliares
|
| 137 |
+
# ============================================================================
|
| 138 |
+
def _grid_neighbors_set(
|
| 139 |
+
idx: Tuple[int, int, int, int],
|
| 140 |
+
grid_shape: Tuple[int, int, int, int],
|
| 141 |
+
) -> set:
|
| 142 |
+
"""Retorna o conjunto de índices vizinhos diretos (6-connectivity em 4D)
|
| 143 |
+
de um neurônio na grade."""
|
| 144 |
+
i, j, k, l = idx
|
| 145 |
+
I, J, K, L = grid_shape
|
| 146 |
+
neighbors = set()
|
| 147 |
+
for di, dj, dk, dl in [
|
| 148 |
+
(1, 0, 0, 0), (-1, 0, 0, 0),
|
| 149 |
+
(0, 1, 0, 0), (0, -1, 0, 0),
|
| 150 |
+
(0, 0, 1, 0), (0, 0, -1, 0),
|
| 151 |
+
(0, 0, 0, 1), (0, 0, 0, -1),
|
| 152 |
+
]:
|
| 153 |
+
ni, nj, nk, nl = i + di, j + dj, k + dk, l + dl
|
| 154 |
+
if 0 <= ni < I and 0 <= nj < J and 0 <= nk < K and 0 <= nl < L:
|
| 155 |
+
neighbors.add((ni, nj, nk, nl))
|
| 156 |
+
return neighbors
|
| 157 |
+
|
| 158 |
+
|
| 159 |
+
def _find_bmu_and_second(
|
| 160 |
+
x: torch.Tensor, weights: torch.Tensor
|
| 161 |
+
) -> Tuple[Tuple[int, int, int, int], Tuple[int, int, int, int], float, float]:
|
| 162 |
+
"""Encontra BMU e 2ª BMU para um vetor x.
|
| 163 |
+
|
| 164 |
+
Args:
|
| 165 |
+
x: tensor [4] — vetor 4D de entrada.
|
| 166 |
+
weights: tensor [I, J, K, L, 4] — pesos do SOM.
|
| 167 |
+
|
| 168 |
+
Returns:
|
| 169 |
+
(bmu_idx, second_bmu_idx, dist_bmu, dist_second)
|
| 170 |
+
"""
|
| 171 |
+
grid_shape = weights.shape[:-1]
|
| 172 |
+
flat_w = weights.view(-1, 4) # (I*J*K*L, 4)
|
| 173 |
+
diff = flat_w - x.view(1, 4)
|
| 174 |
+
dist_sq = torch.sum(diff * diff, dim=-1) # (I*J*K*L,)
|
| 175 |
+
|
| 176 |
+
# Top-2 menores distâncias
|
| 177 |
+
top2 = torch.topk(dist_sq, k=2, largest=False)
|
| 178 |
+
flat_idx_bmu = int(top2.indices[0].item())
|
| 179 |
+
flat_idx_2nd = int(top2.indices[1].item())
|
| 180 |
+
|
| 181 |
+
# Unravel
|
| 182 |
+
L = grid_shape[3]
|
| 183 |
+
K = grid_shape[2]
|
| 184 |
+
J = grid_shape[1]
|
| 185 |
+
|
| 186 |
+
def unravel(flat_idx: int) -> Tuple[int, int, int, int]:
|
| 187 |
+
i = flat_idx // (J * K * L)
|
| 188 |
+
r = flat_idx % (J * K * L)
|
| 189 |
+
j = r // (K * L)
|
| 190 |
+
r = r % (K * L)
|
| 191 |
+
k = r // L
|
| 192 |
+
l = r % L
|
| 193 |
+
return (i, j, k, l)
|
| 194 |
+
|
| 195 |
+
bmu = unravel(flat_idx_bmu)
|
| 196 |
+
second = unravel(flat_idx_2nd)
|
| 197 |
+
dist_bmu = float(math.sqrt(max(top2.values[0].item(), 0.0)))
|
| 198 |
+
dist_second = float(math.sqrt(max(top2.values[1].item(), 0.0)))
|
| 199 |
+
return bmu, second, dist_bmu, dist_second
|
| 200 |
+
|
| 201 |
+
|
| 202 |
+
# ============================================================================
|
| 203 |
+
# Métricas principais (1.1 - 1.4)
|
| 204 |
+
# ============================================================================
|
| 205 |
+
def quantization_error(
|
| 206 |
+
data: torch.Tensor, weights: torch.Tensor
|
| 207 |
+
) -> float:
|
| 208 |
+
"""1.1 Erro de Quantização (QE).
|
| 209 |
+
|
| 210 |
+
QE = (1/N) * Σ_i ||x_i - W_BMU(x_i)||₂
|
| 211 |
+
|
| 212 |
+
Args:
|
| 213 |
+
data: tensor [N, 4] — vetores de entrada.
|
| 214 |
+
weights: tensor [I, J, K, L, 4] — pesos do SOM.
|
| 215 |
+
|
| 216 |
+
Returns:
|
| 217 |
+
QE (float). 0.0 se data vazio. NaN propagado é substituído por 0.0
|
| 218 |
+
(V6.5-V2-metrics-FIX-2: user requirement "investigar QE e KL resultando
|
| 219 |
+
em NAN").
|
| 220 |
+
"""
|
| 221 |
+
if data.numel() == 0:
|
| 222 |
+
return 0.0
|
| 223 |
+
flat_w = weights.view(-1, 4) # (M, 4)
|
| 224 |
+
# distância de cada amostra a todos os neurônios
|
| 225 |
+
# data: (N, 4), flat_w: (M, 4) -> (N, M)
|
| 226 |
+
diff = data.unsqueeze(1) - flat_w.unsqueeze(0)
|
| 227 |
+
dist_sq = torch.sum(diff * diff, dim=-1) # (N, M)
|
| 228 |
+
min_dist_sq, _ = torch.min(dist_sq, dim=1) # (N,)
|
| 229 |
+
# V6.5-V2-metrics-FIX-2: clamp para evitar sqrt de número negativo (float err)
|
| 230 |
+
min_dist_sq = torch.clamp(min_dist_sq, min=0.0)
|
| 231 |
+
qe = torch.mean(torch.sqrt(min_dist_sq))
|
| 232 |
+
qe_val = float(qe.item())
|
| 233 |
+
# NaN check final (defense in depth)
|
| 234 |
+
if qe_val != qe_val: # NaN
|
| 235 |
+
return 0.0
|
| 236 |
+
return qe_val
|
| 237 |
+
|
| 238 |
+
|
| 239 |
+
def topological_error(
|
| 240 |
+
data: torch.Tensor, weights: torch.Tensor
|
| 241 |
+
) -> float:
|
| 242 |
+
"""1.2 Erro Topológico (TE).
|
| 243 |
+
|
| 244 |
+
TE = (1/N) * Σ_i 𝟙[¬adjacent(BMU_1(x_i), BMU_2(x_i))]
|
| 245 |
+
|
| 246 |
+
Dois neurônios são adjacentes se differem em exatamente uma coordenada
|
| 247 |
+
por ±1 (6-connectivity em 4D).
|
| 248 |
+
|
| 249 |
+
Args:
|
| 250 |
+
data: tensor [N, 4]
|
| 251 |
+
weights: tensor [I, J, K, L, 4]
|
| 252 |
+
|
| 253 |
+
Returns:
|
| 254 |
+
TE ∈ [0, 1]. 0.0 se data vazio.
|
| 255 |
+
"""
|
| 256 |
+
if data.numel() == 0:
|
| 257 |
+
return 0.0
|
| 258 |
+
grid_shape = weights.shape[:-1]
|
| 259 |
+
flat_w = weights.view(-1, 4)
|
| 260 |
+
diff = data.unsqueeze(1) - flat_w.unsqueeze(0)
|
| 261 |
+
dist_sq = torch.sum(diff * diff, dim=-1) # (N, M)
|
| 262 |
+
|
| 263 |
+
# Top-2 BMUs para cada amostra
|
| 264 |
+
top2 = torch.topk(dist_sq, k=2, largest=False, dim=1)
|
| 265 |
+
flat_bmu = top2.indices[:, 0] # (N,)
|
| 266 |
+
flat_2nd = top2.indices[:, 1] # (N,)
|
| 267 |
+
|
| 268 |
+
# Unravel
|
| 269 |
+
L = grid_shape[3]
|
| 270 |
+
K = grid_shape[2]
|
| 271 |
+
J = grid_shape[1]
|
| 272 |
+
I = grid_shape[0]
|
| 273 |
+
|
| 274 |
+
def unravel_batch(flat_idx: torch.Tensor) -> torch.Tensor:
|
| 275 |
+
"""Retorna tensor (N, 4) com (i, j, k, l)."""
|
| 276 |
+
i = flat_idx // (J * K * L)
|
| 277 |
+
r = flat_idx % (J * K * L)
|
| 278 |
+
j = r // (K * L)
|
| 279 |
+
r = r % (K * L)
|
| 280 |
+
k = r // L
|
| 281 |
+
l = r % L
|
| 282 |
+
return torch.stack([i, j, k, l], dim=1).long()
|
| 283 |
+
|
| 284 |
+
bmu_idx = unravel_batch(flat_bmu) # (N, 4)
|
| 285 |
+
second_idx = unravel_batch(flat_2nd) # (N, 4)
|
| 286 |
+
|
| 287 |
+
# Adjacência: |Δcoord| soma = 1 (apenas uma coordenada muda por ±1)
|
| 288 |
+
delta = (bmu_idx - second_idx).abs() # (N, 4)
|
| 289 |
+
is_adjacent = (delta.sum(dim=1) == 1) & (delta.max(dim=1).values == 1)
|
| 290 |
+
|
| 291 |
+
te = 1.0 - float(is_adjacent.float().mean().item())
|
| 292 |
+
return te
|
| 293 |
+
|
| 294 |
+
|
| 295 |
+
def kaski_lagus_error(
|
| 296 |
+
data: torch.Tensor, weights: torch.Tensor, alpha: float = 0.5
|
| 297 |
+
) -> float:
|
| 298 |
+
"""1.3 Erro de Kaski-Lagus.
|
| 299 |
+
|
| 300 |
+
Combinação ponderada de QE normalizado e TE.
|
| 301 |
+
|
| 302 |
+
Variante implementada (Kaski & Lagus, 1999):
|
| 303 |
+
KL = α * MQE_norm + (1-α) * TE
|
| 304 |
+
onde MQE_norm = QE / mean(||x||) (normalização por escala dos dados).
|
| 305 |
+
|
| 306 |
+
Uma variante alternativa multiplica:
|
| 307 |
+
KL = QE * (1 + TE)
|
| 308 |
+
que penaliza QE quando TE é alto (usada se use_multiplicative=True).
|
| 309 |
+
|
| 310 |
+
Args:
|
| 311 |
+
data: tensor [N, 4]
|
| 312 |
+
weights: tensor [I, J, K, L, 4]
|
| 313 |
+
alpha: peso do QE normalizado (default 0.5)
|
| 314 |
+
|
| 315 |
+
Returns:
|
| 316 |
+
KL (float). 0.0 se data vazio ou se NaN for detectado
|
| 317 |
+
(V6.5-V2-metrics-FIX-2: user requirement "investigar QE e KL resultando
|
| 318 |
+
em NAN").
|
| 319 |
+
"""
|
| 320 |
+
if data.numel() == 0:
|
| 321 |
+
return 0.0
|
| 322 |
+
qe = quantization_error(data, weights)
|
| 323 |
+
te = topological_error(data, weights)
|
| 324 |
+
# Normalização: escala típica dos dados
|
| 325 |
+
data_norm = float(torch.mean(torch.norm(data, dim=1)).item())
|
| 326 |
+
# V6.5-V2-metrics-FIX-2: se data_norm é NaN ou zero, usa 1.0 como fallback
|
| 327 |
+
if data_norm != data_norm or data_norm < 1e-8: # NaN or ~0
|
| 328 |
+
data_norm = 1.0
|
| 329 |
+
qe_norm = qe / max(data_norm, 1e-8)
|
| 330 |
+
kl = alpha * qe_norm + (1.0 - alpha) * te
|
| 331 |
+
kl_val = float(kl)
|
| 332 |
+
# NaN check final
|
| 333 |
+
if kl_val != kl_val: # NaN
|
| 334 |
+
return 0.0
|
| 335 |
+
return kl_val
|
| 336 |
+
|
| 337 |
+
|
| 338 |
+
def explained_variance_share(
|
| 339 |
+
data: torch.Tensor, weights: torch.Tensor
|
| 340 |
+
) -> float:
|
| 341 |
+
"""1.4 Variância Explicada (Share of Explained Variance).
|
| 342 |
+
|
| 343 |
+
VarExplained = 1 - Var(residual) / Var(total)
|
| 344 |
+
onde residual = x_i - W_BMU(x_i)
|
| 345 |
+
|
| 346 |
+
Implementação: computamos a variância total dos dados (soma das variâncias
|
| 347 |
+
por dimensão) e a variância dos resíduos. A fração explicada é:
|
| 348 |
+
|
| 349 |
+
VE = 1 - Σ_d Var(residual_d) / Σ_d Var(data_d)
|
| 350 |
+
|
| 351 |
+
Returns:
|
| 352 |
+
VE ∈ [0, 1]. Próximo de 1 = boa representação.
|
| 353 |
+
"""
|
| 354 |
+
if data.numel() == 0 or data.shape[0] < 2:
|
| 355 |
+
return 0.0
|
| 356 |
+
flat_w = weights.view(-1, 4)
|
| 357 |
+
diff = data.unsqueeze(1) - flat_w.unsqueeze(0)
|
| 358 |
+
dist_sq = torch.sum(diff * diff, dim=-1) # (N, M)
|
| 359 |
+
_, bmu_flat = torch.min(dist_sq, dim=1) # (N,)
|
| 360 |
+
bmu_w = flat_w[bmu_flat] # (N, 4)
|
| 361 |
+
residual = data - bmu_w # (N, 4)
|
| 362 |
+
|
| 363 |
+
# Variância por dimensão (sem viés, ddof=0)
|
| 364 |
+
var_data = torch.var(data, dim=0, unbiased=False) # (4,)
|
| 365 |
+
var_resid = torch.var(residual, dim=0, unbiased=False) # (4,)
|
| 366 |
+
|
| 367 |
+
total_var = float(torch.sum(var_data).item())
|
| 368 |
+
resid_var = float(torch.sum(var_resid).item())
|
| 369 |
+
if total_var < 1e-12:
|
| 370 |
+
return 0.0
|
| 371 |
+
ve = 1.0 - (resid_var / total_var)
|
| 372 |
+
return float(max(0.0, min(1.0, ve)))
|
| 373 |
+
|
| 374 |
+
|
| 375 |
+
# ============================================================================
|
| 376 |
+
# Indicadores de falha (2.1 - 2.4)
|
| 377 |
+
# ============================================================================
|
| 378 |
+
def topological_collapse(weights: torch.Tensor) -> Dict[str, Any]:
|
| 379 |
+
"""2.1 Colapso Topológico (Topological Collapse).
|
| 380 |
+
|
| 381 |
+
Detecta se os pesos colapsaram para um ponto (rank < 1) ou para uma linha
|
| 382 |
+
(rank < 2). Usa SVD sobre os pesos reshaped (M, 4).
|
| 383 |
+
|
| 384 |
+
Args:
|
| 385 |
+
weights: tensor [I, J, K, L, 4]
|
| 386 |
+
|
| 387 |
+
Returns:
|
| 388 |
+
Dict com:
|
| 389 |
+
- collapsed_to_point: bool (todos iguais)
|
| 390 |
+
- collapsed_to_line: bool (variação só em 1 direção)
|
| 391 |
+
- effective_rank: float (rank efetivo via razão de singular values)
|
| 392 |
+
- singular_values: list[float]
|
| 393 |
+
- severity: str ("none" | "line" | "point")
|
| 394 |
+
"""
|
| 395 |
+
flat_w = weights.view(-1, 4).clone() # (M, 4)
|
| 396 |
+
if flat_w.numel() == 0:
|
| 397 |
+
return {"collapsed_to_point": False, "collapsed_to_line": False,
|
| 398 |
+
"effective_rank": 0.0, "singular_values": [], "severity": "none"}
|
| 399 |
+
|
| 400 |
+
# Centraliza
|
| 401 |
+
flat_w_centered = flat_w - flat_w.mean(dim=0, keepdim=True)
|
| 402 |
+
|
| 403 |
+
# V6.5-V2-metrics FIX: usa SVD com fallback robusto para matrizes
|
| 404 |
+
# singulares/ill-conditioned. Se SVD falhar, usa np.linalg.eigh com
|
| 405 |
+
# regularização (adiciona pequeno ruído diagonal).
|
| 406 |
+
S_list: List[float] = []
|
| 407 |
+
try:
|
| 408 |
+
# Tentativa 1: SVD direto (mais estável que eigh)
|
| 409 |
+
U, S, V = torch.linalg.svd(flat_w_centered, full_matrices=False)
|
| 410 |
+
S_list = [float(s.item()) for s in S]
|
| 411 |
+
except Exception:
|
| 412 |
+
# Tentativa 2: eigendecomposição da matriz de covariância com regularização
|
| 413 |
+
try:
|
| 414 |
+
cov = flat_w_centered.t() @ flat_w_centered
|
| 415 |
+
# Regularização: adiciona pequeno valor diagonal para estabilizar
|
| 416 |
+
n = cov.shape[0]
|
| 417 |
+
reg = 1e-6 * float(cov.diag().max().item()) if n > 0 else 1e-6
|
| 418 |
+
cov_reg = cov + torch.eye(n) * reg
|
| 419 |
+
S_eigh, _ = torch.linalg.eigh(cov_reg)
|
| 420 |
+
S_eigh = S_eigh.flip(0) # descending order
|
| 421 |
+
# Converte autovalores em "singular values" aproximados (sqrt)
|
| 422 |
+
S_list = [float(math.sqrt(max(s.item(), 0.0))) for s in S_eigh]
|
| 423 |
+
except Exception:
|
| 424 |
+
# Tentativa 3: usa norma de Frobenius como aproximação
|
| 425 |
+
s_frob = float(torch.norm(flat_w_centered).item())
|
| 426 |
+
S_list = [s_frob, 0.0, 0.0, 0.0]
|
| 427 |
+
|
| 428 |
+
if not S_list:
|
| 429 |
+
return {"collapsed_to_point": False, "collapsed_to_line": False,
|
| 430 |
+
"effective_rank": 0.0, "singular_values": [], "severity": "none"}
|
| 431 |
+
|
| 432 |
+
s_max = max(S_list[0] if S_list else 0.0, 1e-12)
|
| 433 |
+
|
| 434 |
+
# Effective rank: número de singular values > 1% do maior
|
| 435 |
+
threshold = 0.01 * s_max
|
| 436 |
+
n_significant = sum(1 for s in S_list if s > threshold)
|
| 437 |
+
# Effective rank (float, via entropy)
|
| 438 |
+
s_norm = [s / s_max for s in S_list]
|
| 439 |
+
s_norm = [s for s in s_norm if s > 1e-10]
|
| 440 |
+
if s_norm:
|
| 441 |
+
p = torch.tensor(s_norm)
|
| 442 |
+
p_norm = p / p.sum()
|
| 443 |
+
entropy = -float(torch.sum(p_norm * torch.log(p_norm)).item())
|
| 444 |
+
effective_rank = math.exp(entropy)
|
| 445 |
+
else:
|
| 446 |
+
effective_rank = 0.0
|
| 447 |
+
|
| 448 |
+
collapsed_to_point = n_significant <= 1 and s_max < 1e-6
|
| 449 |
+
collapsed_to_line = n_significant == 1 and s_max >= 1e-6
|
| 450 |
+
|
| 451 |
+
if collapsed_to_point:
|
| 452 |
+
severity = "point"
|
| 453 |
+
elif collapsed_to_line:
|
| 454 |
+
severity = "line"
|
| 455 |
+
else:
|
| 456 |
+
severity = "none"
|
| 457 |
+
|
| 458 |
+
return {
|
| 459 |
+
"collapsed_to_point": bool(collapsed_to_point),
|
| 460 |
+
"collapsed_to_line": bool(collapsed_to_line),
|
| 461 |
+
"effective_rank": float(effective_rank),
|
| 462 |
+
"singular_values": S_list,
|
| 463 |
+
"severity": severity,
|
| 464 |
+
}
|
| 465 |
+
|
| 466 |
+
|
| 467 |
+
def dead_neuron_rate(
|
| 468 |
+
data: torch.Tensor, weights: torch.Tensor
|
| 469 |
+
) -> Dict[str, Any]:
|
| 470 |
+
"""2.2 Taxa de Neurônios Mortos (Dead Neurons).
|
| 471 |
+
|
| 472 |
+
DeadRate = (n_neurons - n_active_neurons) / n_neurons
|
| 473 |
+
|
| 474 |
+
Um neurônio é "ativo" se é BMU para pelo menos uma amostra.
|
| 475 |
+
|
| 476 |
+
Args:
|
| 477 |
+
data: tensor [N, 4]
|
| 478 |
+
weights: tensor [I, J, K, L, 4]
|
| 479 |
+
|
| 480 |
+
Returns:
|
| 481 |
+
Dict com:
|
| 482 |
+
- dead_neuron_rate: float ∈ [0, 1]
|
| 483 |
+
- n_dead: int
|
| 484 |
+
- n_active: int
|
| 485 |
+
- n_total: int
|
| 486 |
+
- dead_indices: list de (i,j,k,l)
|
| 487 |
+
"""
|
| 488 |
+
n_total = weights.shape[0] * weights.shape[1] * weights.shape[2] * weights.shape[3]
|
| 489 |
+
if data.numel() == 0:
|
| 490 |
+
return {"dead_neuron_rate": 1.0, "n_dead": n_total, "n_active": 0,
|
| 491 |
+
"n_total": n_total, "dead_indices": []}
|
| 492 |
+
|
| 493 |
+
flat_w = weights.view(-1, 4)
|
| 494 |
+
diff = data.unsqueeze(1) - flat_w.unsqueeze(0)
|
| 495 |
+
dist_sq = torch.sum(diff * diff, dim=-1) # (N, M)
|
| 496 |
+
_, bmu_flat = torch.min(dist_sq, dim=1) # (N,)
|
| 497 |
+
|
| 498 |
+
active_mask = torch.zeros(n_total, dtype=torch.bool)
|
| 499 |
+
active_mask[bmu_flat] = True
|
| 500 |
+
n_active = int(active_mask.sum().item())
|
| 501 |
+
n_dead = n_total - n_active
|
| 502 |
+
|
| 503 |
+
# Coleta índices mortos
|
| 504 |
+
dead_flat = (~active_mask).nonzero(as_tuple=True)[0]
|
| 505 |
+
L_dim = weights.shape[3]
|
| 506 |
+
K_dim = weights.shape[2]
|
| 507 |
+
J_dim = weights.shape[1]
|
| 508 |
+
dead_indices = []
|
| 509 |
+
for flat_idx in dead_flat.tolist():
|
| 510 |
+
i = flat_idx // (J_dim * K_dim * L_dim)
|
| 511 |
+
r = flat_idx % (J_dim * K_dim * L_dim)
|
| 512 |
+
j = r // (K_dim * L_dim)
|
| 513 |
+
r = r % (K_dim * L_dim)
|
| 514 |
+
k = r // L_dim
|
| 515 |
+
l = r % L_dim
|
| 516 |
+
dead_indices.append((i, j, k, l))
|
| 517 |
+
|
| 518 |
+
return {
|
| 519 |
+
"dead_neuron_rate": float(n_dead / n_total) if n_total > 0 else 0.0,
|
| 520 |
+
"n_dead": int(n_dead),
|
| 521 |
+
"n_active": int(n_active),
|
| 522 |
+
"n_total": int(n_total),
|
| 523 |
+
"dead_indices": dead_indices[:50], # limita para não explodir JSON
|
| 524 |
+
}
|
| 525 |
+
|
| 526 |
+
|
| 527 |
+
def quantization_error_stagnation(
|
| 528 |
+
history: SOMMetricHistory,
|
| 529 |
+
high_threshold_frac: float = 0.7,
|
| 530 |
+
slope_tolerance: float = 1e-4,
|
| 531 |
+
window_size: int = 5,
|
| 532 |
+
) -> Dict[str, Any]:
|
| 533 |
+
"""2.3 Estagnação do Erro de Quantização.
|
| 534 |
+
|
| 535 |
+
Detecta se o QE estabilizou em patamar elevado:
|
| 536 |
+
- |slope(QE)| < slope_tolerance (estável)
|
| 537 |
+
- QE_atual > high_threshold_frac * max(QE_history) (alto)
|
| 538 |
+
|
| 539 |
+
Args:
|
| 540 |
+
history: histórico de métricas SOM.
|
| 541 |
+
high_threshold_frac: fração do QE máximo para considerar "alto".
|
| 542 |
+
slope_tolerance: tolerância de slope para considerar "estável".
|
| 543 |
+
window_size: nº de épocas recentes para regressão linear.
|
| 544 |
+
|
| 545 |
+
Returns:
|
| 546 |
+
Dict com:
|
| 547 |
+
- is_stagnant: bool
|
| 548 |
+
- current_qe: float
|
| 549 |
+
- max_qe: float
|
| 550 |
+
- slope: float
|
| 551 |
+
- severity: str ("none" | "mild" | "severe")
|
| 552 |
+
"""
|
| 553 |
+
qe_list = list(history.qe_history)
|
| 554 |
+
if len(qe_list) < window_size + 2:
|
| 555 |
+
return {"is_stagnant": False, "current_qe": 0.0, "max_qe": 0.0,
|
| 556 |
+
"slope": 0.0, "severity": "none", "n_epochs": len(qe_list)}
|
| 557 |
+
|
| 558 |
+
# Regressão linear sobre as últimas window_size épocas
|
| 559 |
+
recent = qe_list[-window_size:]
|
| 560 |
+
n = len(recent)
|
| 561 |
+
x_mean = (n - 1) / 2.0
|
| 562 |
+
y_mean = sum(recent) / n
|
| 563 |
+
num = sum((i - x_mean) * (y - y_mean) for i, y in enumerate(recent))
|
| 564 |
+
den = sum((i - x_mean) ** 2 for i in range(n))
|
| 565 |
+
slope = num / den if den > 0 else 0.0
|
| 566 |
+
|
| 567 |
+
current_qe = float(qe_list[-1])
|
| 568 |
+
max_qe = float(max(qe_list))
|
| 569 |
+
is_high = current_qe > high_threshold_frac * max_qe
|
| 570 |
+
is_stable = abs(slope) < slope_tolerance
|
| 571 |
+
|
| 572 |
+
is_stagnant = is_stable and is_high
|
| 573 |
+
|
| 574 |
+
if is_stagnant:
|
| 575 |
+
if current_qe > 0.9 * max_qe:
|
| 576 |
+
severity = "severe"
|
| 577 |
+
else:
|
| 578 |
+
severity = "mild"
|
| 579 |
+
else:
|
| 580 |
+
severity = "none"
|
| 581 |
+
|
| 582 |
+
return {
|
| 583 |
+
"is_stagnant": bool(is_stagnant),
|
| 584 |
+
"current_qe": current_qe,
|
| 585 |
+
"max_qe": max_qe,
|
| 586 |
+
"slope": float(slope),
|
| 587 |
+
"severity": severity,
|
| 588 |
+
"n_epochs": len(qe_list),
|
| 589 |
+
}
|
| 590 |
+
|
| 591 |
+
|
| 592 |
+
def neighborhood_crossing(
|
| 593 |
+
history: SOMMetricHistory,
|
| 594 |
+
initial_window: int = 5,
|
| 595 |
+
final_window: int = 5,
|
| 596 |
+
growth_threshold: float = 1.5,
|
| 597 |
+
) -> Dict[str, Any]:
|
| 598 |
+
"""2.4 Cruzamento de Vizinhança (Neighborhood Crossing).
|
| 599 |
+
|
| 600 |
+
Detecta crescimento repentino ou oscilação crônica do TE nas épocas
|
| 601 |
+
finais:
|
| 602 |
+
- std(TE_final) > growth_threshold * std(TE_initial) (oscilação cresceu)
|
| 603 |
+
- TE[-1] > TE[-5] (crescimento no final)
|
| 604 |
+
|
| 605 |
+
Args:
|
| 606 |
+
history: histórico de métricas SOM.
|
| 607 |
+
initial_window: nº de épocas iniciais para baseline.
|
| 608 |
+
final_window: nº de épocas finais para comparação.
|
| 609 |
+
growth_threshold: razão de std para considerar "crescimento crônico".
|
| 610 |
+
|
| 611 |
+
Returns:
|
| 612 |
+
Dict com:
|
| 613 |
+
- detected: bool
|
| 614 |
+
- te_initial_mean: float
|
| 615 |
+
- te_final_mean: float
|
| 616 |
+
- te_initial_std: float
|
| 617 |
+
- te_final_std: float
|
| 618 |
+
- std_ratio: float
|
| 619 |
+
- is_oscillating: bool
|
| 620 |
+
- is_growing: bool
|
| 621 |
+
- severity: str ("none" | "mild" | "severe")
|
| 622 |
+
"""
|
| 623 |
+
te_list = list(history.te_history)
|
| 624 |
+
if len(te_list) < initial_window + final_window:
|
| 625 |
+
return {"detected": False, "te_initial_mean": 0.0, "te_final_mean": 0.0,
|
| 626 |
+
"te_initial_std": 0.0, "te_final_std": 0.0, "std_ratio": 0.0,
|
| 627 |
+
"is_oscillating": False, "is_growing": False, "severity": "none",
|
| 628 |
+
"n_epochs": len(te_list)}
|
| 629 |
+
|
| 630 |
+
initial = te_list[:initial_window]
|
| 631 |
+
final = te_list[-final_window:]
|
| 632 |
+
|
| 633 |
+
init_mean = sum(initial) / len(initial)
|
| 634 |
+
final_mean = sum(final) / len(final)
|
| 635 |
+
init_var = sum((x - init_mean) ** 2 for x in initial) / len(initial)
|
| 636 |
+
final_var = sum((x - final_mean) ** 2 for x in final) / len(final)
|
| 637 |
+
init_std = math.sqrt(init_var)
|
| 638 |
+
final_std = math.sqrt(final_var)
|
| 639 |
+
|
| 640 |
+
std_ratio = final_std / max(init_std, 1e-8)
|
| 641 |
+
is_oscillating = std_ratio > growth_threshold
|
| 642 |
+
is_growing = final_mean > init_mean and (final_mean - init_mean) > 0.05
|
| 643 |
+
|
| 644 |
+
detected = is_oscillating or is_growing
|
| 645 |
+
|
| 646 |
+
if detected:
|
| 647 |
+
if std_ratio > 2.0 * growth_threshold or is_growing:
|
| 648 |
+
severity = "severe"
|
| 649 |
+
else:
|
| 650 |
+
severity = "mild"
|
| 651 |
+
else:
|
| 652 |
+
severity = "none"
|
| 653 |
+
|
| 654 |
+
return {
|
| 655 |
+
"detected": bool(detected),
|
| 656 |
+
"te_initial_mean": float(init_mean),
|
| 657 |
+
"te_final_mean": float(final_mean),
|
| 658 |
+
"te_initial_std": float(init_std),
|
| 659 |
+
"te_final_std": float(final_std),
|
| 660 |
+
"std_ratio": float(std_ratio),
|
| 661 |
+
"is_oscillating": bool(is_oscillating),
|
| 662 |
+
"is_growing": bool(is_growing),
|
| 663 |
+
"severity": severity,
|
| 664 |
+
"n_epochs": len(te_list),
|
| 665 |
+
}
|
| 666 |
+
|
| 667 |
+
|
| 668 |
+
# ============================================================================
|
| 669 |
+
# Função agregadora: compute_all_metrics
|
| 670 |
+
# ============================================================================
|
| 671 |
+
def compute_all_metrics(
|
| 672 |
+
data: torch.Tensor,
|
| 673 |
+
weights: torch.Tensor,
|
| 674 |
+
history: Optional[SOMMetricHistory] = None,
|
| 675 |
+
) -> Dict[str, Any]:
|
| 676 |
+
"""Computa todas as métricas SOM de uma vez.
|
| 677 |
+
|
| 678 |
+
Args:
|
| 679 |
+
data: tensor [N, 4] — buffer de vetores 4D.
|
| 680 |
+
weights: tensor [I, J, K, L, 4] — pesos do SOM.
|
| 681 |
+
history: histórico de métricas (opcional, para estagnação e cruzamento).
|
| 682 |
+
|
| 683 |
+
Returns:
|
| 684 |
+
Dict com:
|
| 685 |
+
- quantization_error: float
|
| 686 |
+
- topological_error: float
|
| 687 |
+
- kaski_lagus_error: float
|
| 688 |
+
- explained_variance_share: float
|
| 689 |
+
- topological_collapse: Dict
|
| 690 |
+
- dead_neuron_rate: Dict
|
| 691 |
+
- quantization_error_stagnation: Dict (se history fornecido)
|
| 692 |
+
- neighborhood_crossing: Dict (se history fornecido)
|
| 693 |
+
- overall_health: str ("healthy" | "warning" | "critical")
|
| 694 |
+
- failure_indicators: list[str]
|
| 695 |
+
"""
|
| 696 |
+
# V6.5-V2-metrics FIX: wrap entire computation in try/except para evitar
|
| 697 |
+
# crashes fatais quando o SOM tem pesos degenerados (ex: todos iguais
|
| 698 |
+
# após muito treinamento sem diversidade). Em caso de erro, retorna
|
| 699 |
+
# métricas safe-default em vez de crashar o treino.
|
| 700 |
+
|
| 701 |
+
# V6.5-V2-metrics-FIX-2 (user requirement: "investigar QE e KL resultando
|
| 702 |
+
# em NAN"): detecção EXPLÍCITA de NaN/Inf nos tensores de entrada ANTES
|
| 703 |
+
# de computar as métricas. Se os pesos do SOM contêm NaN (devido a
|
| 704 |
+
# gradientes explosivos durante train_hypotheses ou apply_best_delta),
|
| 705 |
+
# TODAS as métricas (QE, KL, etc.) seriam NaN, propagando o erro para
|
| 706 |
+
# o histórico e quebrando a detecção de estagnação/cruzamento.
|
| 707 |
+
# Esta barreira inicial retorna safe-defaults imediatamente quando
|
| 708 |
+
# NaN/Inf é detectado, evitando contaminação do histórico.
|
| 709 |
+
try:
|
| 710 |
+
if data is not None and data.numel() > 0:
|
| 711 |
+
if bool(torch.isnan(data).any().item()) or bool(torch.isinf(data).any().item()):
|
| 712 |
+
n_nan = int(torch.isnan(data).sum().item())
|
| 713 |
+
n_inf = int(torch.isinf(data).sum().item())
|
| 714 |
+
logger.warning(
|
| 715 |
+
f"[som_metrics] NaN/Inf detected in INPUT DATA: "
|
| 716 |
+
f"{n_nan} NaN, {n_inf} Inf out of {data.numel()} elements. "
|
| 717 |
+
f"Returning safe-default metrics (skipping computation)."
|
| 718 |
+
)
|
| 719 |
+
return _nan_safe_metrics(
|
| 720 |
+
reason=f"nan_in_data (nan={n_nan}, inf={n_inf})",
|
| 721 |
+
data_shape=tuple(data.shape),
|
| 722 |
+
weights_shape=tuple(weights.shape) if weights is not None else None,
|
| 723 |
+
)
|
| 724 |
+
if weights is not None and weights.numel() > 0:
|
| 725 |
+
if bool(torch.isnan(weights).any().item()) or bool(torch.isinf(weights).any().item()):
|
| 726 |
+
n_nan = int(torch.isnan(weights).sum().item())
|
| 727 |
+
n_inf = int(torch.isinf(weights).sum().item())
|
| 728 |
+
logger.warning(
|
| 729 |
+
f"[som_metrics] NaN/Inf detected in SOM WEIGHTS: "
|
| 730 |
+
f"{n_nan} NaN, {n_inf} Inf out of {weights.numel()} elements. "
|
| 731 |
+
f"Returning safe-default metrics (skipping computation). "
|
| 732 |
+
f"This usually indicates exploding gradients in "
|
| 733 |
+
f"train_hypotheses() or apply_best_delta_and_consolidate()."
|
| 734 |
+
)
|
| 735 |
+
return _nan_safe_metrics(
|
| 736 |
+
reason=f"nan_in_weights (nan={n_nan}, inf={n_inf})",
|
| 737 |
+
data_shape=tuple(data.shape) if data is not None else None,
|
| 738 |
+
weights_shape=tuple(weights.shape),
|
| 739 |
+
)
|
| 740 |
+
except Exception as nan_check_err:
|
| 741 |
+
logger.warning(
|
| 742 |
+
f"[som_metrics] NaN pre-check failed (proceeding to compute): {nan_check_err}"
|
| 743 |
+
)
|
| 744 |
+
|
| 745 |
+
try:
|
| 746 |
+
result = _compute_all_metrics_impl(data, weights, history)
|
| 747 |
+
# V6.5-V2-metrics-FIX-2: pós-computação, verifica se QE ou KL é NaN.
|
| 748 |
+
# Se sim, registra warning e substitui por 0.0 (para não quebrar logs JSON
|
| 749 |
+
# e nem contaminar o histórico de estagnação/cruzamento).
|
| 750 |
+
qe_val = result.get("quantization_error", 0.0)
|
| 751 |
+
kl_val = result.get("kaski_lagus_error", 0.0)
|
| 752 |
+
if isinstance(qe_val, float) and (qe_val != qe_val): # NaN check
|
| 753 |
+
logger.warning(
|
| 754 |
+
f"[som_metrics] QE computed as NaN — replacing with 0.0 "
|
| 755 |
+
f"(indicates degenerate SOM state). Investigate train_hypotheses "
|
| 756 |
+
f"or apply_best_delta for gradient explosion."
|
| 757 |
+
)
|
| 758 |
+
result["quantization_error"] = 0.0
|
| 759 |
+
result["qe_was_nan"] = True
|
| 760 |
+
# KL depende de QE, então também será NaN
|
| 761 |
+
result["kaski_lagus_error"] = 0.0
|
| 762 |
+
result["kl_was_nan"] = True
|
| 763 |
+
elif isinstance(kl_val, float) and (kl_val != kl_val): # NaN check
|
| 764 |
+
logger.warning(
|
| 765 |
+
f"[som_metrics] KL computed as NaN (QE was fine) — replacing with 0.0. "
|
| 766 |
+
f"Usually due to zero-norm data in kaski_lagus normalization."
|
| 767 |
+
)
|
| 768 |
+
result["kaski_lagus_error"] = 0.0
|
| 769 |
+
result["kl_was_nan"] = True
|
| 770 |
+
return result
|
| 771 |
+
except Exception as e:
|
| 772 |
+
logger.warning(
|
| 773 |
+
f"[som_metrics] compute_all_metrics failed (using safe defaults): {e}"
|
| 774 |
+
)
|
| 775 |
+
return _nan_safe_metrics(reason=f"compute_error: {str(e)[:200]}", data_shape=None, weights_shape=None)
|
| 776 |
+
|
| 777 |
+
|
| 778 |
+
def _nan_safe_metrics(
|
| 779 |
+
reason: str,
|
| 780 |
+
data_shape: Optional[Tuple[int, ...]] = None,
|
| 781 |
+
weights_shape: Optional[Tuple[int, ...]] = None,
|
| 782 |
+
) -> Dict[str, Any]:
|
| 783 |
+
"""V6.5-V2-metrics-FIX-2 — Retorna métricas safe-default quando NaN é detectado.
|
| 784 |
+
|
| 785 |
+
User requirement: "Aprimorar a FASE2 PUNITIVA investigar QE e KL resultando em NAN".
|
| 786 |
+
|
| 787 |
+
Este helper centraliza a resposta a NaN/Inf em qualquer ponto do pipeline
|
| 788 |
+
de métricas, garantindo que:
|
| 789 |
+
1. O treino NÃO crashe.
|
| 790 |
+
2. O histórico de métricas (SOMMetricHistory) NÃO seja contaminado com NaN
|
| 791 |
+
(o que quebraria detecção de estagnação e cruzamento de vizinhança).
|
| 792 |
+
3. O log indique claramente a fonte do NaN (data vs weights vs compute_error).
|
| 793 |
+
4. failure_indicators registre o evento para análise pós-treino.
|
| 794 |
+
"""
|
| 795 |
+
return {
|
| 796 |
+
"quantization_error": 0.0,
|
| 797 |
+
"topological_error": 0.0,
|
| 798 |
+
"kaski_lagus_error": 0.0,
|
| 799 |
+
"explained_variance_share": 0.0,
|
| 800 |
+
"topological_collapse": {
|
| 801 |
+
"collapsed_to_point": False, "collapsed_to_line": False,
|
| 802 |
+
"effective_rank": 0.0, "singular_values": [],
|
| 803 |
+
"severity": "none", "error": reason,
|
| 804 |
+
},
|
| 805 |
+
"dead_neuron_rate": {
|
| 806 |
+
"dead_neuron_rate": 1.0, "n_dead": 0, "n_active": 0,
|
| 807 |
+
"n_total": 0, "dead_indices": [], "error": reason,
|
| 808 |
+
},
|
| 809 |
+
"qe_stagnation": {
|
| 810 |
+
"is_stagnant": False, "current_qe": 0.0, "max_qe": 0.0,
|
| 811 |
+
"slope": 0.0, "severity": "none", "n_epochs": 0,
|
| 812 |
+
},
|
| 813 |
+
"neighborhood_crossing": {
|
| 814 |
+
"detected": False, "te_initial_mean": 0.0, "te_final_mean": 0.0,
|
| 815 |
+
"te_initial_std": 0.0, "te_final_std": 0.0, "std_ratio": 0.0,
|
| 816 |
+
"is_oscillating": False, "is_growing": False, "severity": "none",
|
| 817 |
+
"n_epochs": 0,
|
| 818 |
+
},
|
| 819 |
+
"overall_health": "critical",
|
| 820 |
+
"failure_indicators": [f"nan_detected: {reason[:120]}"],
|
| 821 |
+
"n_failure_indicators": 1,
|
| 822 |
+
"nan_detected": True,
|
| 823 |
+
"nan_reason": reason,
|
| 824 |
+
"data_shape": list(data_shape) if data_shape else None,
|
| 825 |
+
"weights_shape": list(weights_shape) if weights_shape else None,
|
| 826 |
+
}
|
| 827 |
+
|
| 828 |
+
|
| 829 |
+
def _compute_all_metrics_impl(
|
| 830 |
+
data: torch.Tensor,
|
| 831 |
+
weights: torch.Tensor,
|
| 832 |
+
history: Optional[SOMMetricHistory] = None,
|
| 833 |
+
) -> Dict[str, Any]:
|
| 834 |
+
"""Implementação interna de compute_all_metrics (sem try/except wrapper)."""
|
| 835 |
+
metrics: Dict[str, Any] = {}
|
| 836 |
+
|
| 837 |
+
# 1.1 - 1.4: Métricas principais
|
| 838 |
+
metrics["quantization_error"] = quantization_error(data, weights)
|
| 839 |
+
metrics["topological_error"] = topological_error(data, weights)
|
| 840 |
+
metrics["kaski_lagus_error"] = kaski_lagus_error(data, weights)
|
| 841 |
+
metrics["explained_variance_share"] = explained_variance_share(data, weights)
|
| 842 |
+
|
| 843 |
+
# 2.1: Colapso topológico
|
| 844 |
+
collapse = topological_collapse(weights)
|
| 845 |
+
metrics["topological_collapse"] = collapse
|
| 846 |
+
|
| 847 |
+
# 2.2: Dead neurons
|
| 848 |
+
dead = dead_neuron_rate(data, weights)
|
| 849 |
+
metrics["dead_neuron_rate"] = dead
|
| 850 |
+
|
| 851 |
+
# 2.3: Estagnação do QE (requer histórico)
|
| 852 |
+
if history is not None and len(history.qe_history) > 5:
|
| 853 |
+
metrics["qe_stagnation"] = quantization_error_stagnation(history)
|
| 854 |
+
else:
|
| 855 |
+
metrics["qe_stagnation"] = {
|
| 856 |
+
"is_stagnant": False, "current_qe": metrics["quantization_error"],
|
| 857 |
+
"max_qe": metrics["quantization_error"], "slope": 0.0,
|
| 858 |
+
"severity": "none", "n_epochs": 0,
|
| 859 |
+
}
|
| 860 |
+
|
| 861 |
+
# 2.4: Cruzamento de vizinhança (requer histórico)
|
| 862 |
+
if history is not None and len(history.te_history) > 10:
|
| 863 |
+
metrics["neighborhood_crossing"] = neighborhood_crossing(history)
|
| 864 |
+
else:
|
| 865 |
+
metrics["neighborhood_crossing"] = {
|
| 866 |
+
"detected": False, "te_initial_mean": 0.0, "te_final_mean": 0.0,
|
| 867 |
+
"te_initial_std": 0.0, "te_final_std": 0.0, "std_ratio": 0.0,
|
| 868 |
+
"is_oscillating": False, "is_growing": False, "severity": "none",
|
| 869 |
+
"n_epochs": 0,
|
| 870 |
+
}
|
| 871 |
+
|
| 872 |
+
# Avaliação de saúde geral
|
| 873 |
+
failure_indicators: List[str] = []
|
| 874 |
+
if collapse["collapsed_to_point"]:
|
| 875 |
+
failure_indicators.append("topological_collapse_point")
|
| 876 |
+
elif collapse["collapsed_to_line"]:
|
| 877 |
+
failure_indicators.append("topological_collapse_line")
|
| 878 |
+
if dead["dead_neuron_rate"] > 0.5:
|
| 879 |
+
failure_indicators.append(f"dead_neuron_rate_high_{dead['dead_neuron_rate']:.2f}")
|
| 880 |
+
if metrics["qe_stagnation"]["is_stagnant"]:
|
| 881 |
+
failure_indicators.append(f"qe_stagnant_{metrics['qe_stagnation']['severity']}")
|
| 882 |
+
if metrics["neighborhood_crossing"]["detected"]:
|
| 883 |
+
failure_indicators.append(
|
| 884 |
+
f"neighborhood_crossing_{metrics['neighborhood_crossing']['severity']}"
|
| 885 |
+
)
|
| 886 |
+
if metrics["explained_variance_share"] < 0.3:
|
| 887 |
+
failure_indicators.append(
|
| 888 |
+
f"low_explained_variance_{metrics['explained_variance_share']:.2f}"
|
| 889 |
+
)
|
| 890 |
+
if metrics["topological_error"] > 0.5:
|
| 891 |
+
failure_indicators.append(
|
| 892 |
+
f"high_topological_error_{metrics['topological_error']:.2f}"
|
| 893 |
+
)
|
| 894 |
+
|
| 895 |
+
if len(failure_indicators) >= 2:
|
| 896 |
+
overall_health = "critical"
|
| 897 |
+
elif len(failure_indicators) == 1:
|
| 898 |
+
overall_health = "warning"
|
| 899 |
+
else:
|
| 900 |
+
overall_health = "healthy"
|
| 901 |
+
|
| 902 |
+
metrics["failure_indicators"] = failure_indicators
|
| 903 |
+
metrics["overall_health"] = overall_health
|
| 904 |
+
metrics["n_failure_indicators"] = len(failure_indicators)
|
| 905 |
+
|
| 906 |
+
return metrics
|
| 907 |
+
|
| 908 |
+
|
| 909 |
+
__all__ = [
|
| 910 |
+
"SOMMetricHistory",
|
| 911 |
+
"quantization_error",
|
| 912 |
+
"topological_error",
|
| 913 |
+
"kaski_lagus_error",
|
| 914 |
+
"explained_variance_share",
|
| 915 |
+
"topological_collapse",
|
| 916 |
+
"dead_neuron_rate",
|
| 917 |
+
"quantization_error_stagnation",
|
| 918 |
+
"neighborhood_crossing",
|
| 919 |
+
"compute_all_metrics",
|
| 920 |
+
]
|
streaming_datasets.py
ADDED
|
@@ -0,0 +1,995 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""streaming_datasets_v13_9.py — Streaming dataset loader for v13.9.1 training.
|
| 2 |
+
|
| 3 |
+
Carrega os 9 datasets especificados pelo usuário via streaming (IterableDataset),
|
| 4 |
+
um por vez, exaustivamente dentro do limite max_samples_per_dataset.
|
| 5 |
+
|
| 6 |
+
Datasets:
|
| 7 |
+
1. CEIA-POSITIVO/ultrachat_br_clustred_balanced_v1
|
| 8 |
+
2. Madras1/corpus-ptbr-v2
|
| 9 |
+
3. rhaymison/multmodal_175k_portuguese
|
| 10 |
+
4. TucanoBR/GigaVerbo
|
| 11 |
+
5. nvidia/OpenMathReasoning
|
| 12 |
+
6. MathLLMs/MathVision
|
| 13 |
+
7. nvidia/OpenMathInstruct-2
|
| 14 |
+
8. dominguesm/restore-punctuation-ptbr-dataset
|
| 15 |
+
9. carolina-c4ai/corpus-carolina
|
| 16 |
+
|
| 17 |
+
Memory-efficient: streaming=True, no full materialization.
|
| 18 |
+
Returns ProcessedSample with raw_text only (tokenizer handles encoding).
|
| 19 |
+
"""
|
| 20 |
+
from __future__ import annotations
|
| 21 |
+
|
| 22 |
+
import logging
|
| 23 |
+
import os
|
| 24 |
+
import time
|
| 25 |
+
import traceback
|
| 26 |
+
from dataclasses import dataclass, field
|
| 27 |
+
from typing import Any, Dict, Iterator, List, Optional, Tuple
|
| 28 |
+
|
| 29 |
+
import torch
|
| 30 |
+
|
| 31 |
+
logger = logging.getLogger(__name__)
|
| 32 |
+
|
| 33 |
+
# Datasets V13.9.1 — exatamente os 9 especificados pelo usuário
|
| 34 |
+
# V6.5-final: adicionado dominguesm/Canarim-Instruct-PTBR-Dataset
|
| 35 |
+
# V6.5-7ds: adicionados adalbertojunior/punctuation-ptbr,
|
| 36 |
+
# iara-project/news-articles-ptbr-dataset, manoela/noticias_ptbr
|
| 37 |
+
# V6.5-V2: adicionado BrunoN-Dev/corpus-ptbr-v1 como 8º dataset de CONHECIMENTO
|
| 38 |
+
# e usado como dataset de TREINAMENTO COM PUNIÇÃO
|
| 39 |
+
DEFAULT_DATASETS = [
|
| 40 |
+
"CEIA-POSITIVO/ultrachat_br_clustred_balanced_v1",
|
| 41 |
+
"Madras1/corpus-ptbr-v2",
|
| 42 |
+
"rhaymison/multmodal_175k_portuguese",
|
| 43 |
+
"TucanoBR/GigaVerbo",
|
| 44 |
+
"nvidia/OpenMathReasoning",
|
| 45 |
+
"MathLLMs/MathVision",
|
| 46 |
+
"nvidia/OpenMathInstruct-2",
|
| 47 |
+
"dominguesm/restore-punctuation-ptbr-dataset",
|
| 48 |
+
"carolina-c4ai/corpus-carolina",
|
| 49 |
+
"dominguesm/Canarim-Instruct-PTBR-Dataset",
|
| 50 |
+
"adalbertojunior/punctuation-ptbr",
|
| 51 |
+
"iara-project/news-articles-ptbr-dataset",
|
| 52 |
+
"manoela/noticias_ptbr",
|
| 53 |
+
"BrunoN-Dev/corpus-ptbr-v1",
|
| 54 |
+
]
|
| 55 |
+
|
| 56 |
+
# V6.5-V2 — Sequência de 8 datasets para a fase CONHECIMENTO (na ordem exata
|
| 57 |
+
# especificada pelo usuário). Meta mínima: 8000 samples ou mais.
|
| 58 |
+
CONHECIMENTO_DATASETS_V2 = [
|
| 59 |
+
"dominguesm/restore-punctuation-ptbr-dataset",
|
| 60 |
+
"carolina-c4ai/corpus-carolina",
|
| 61 |
+
"CEIA-POSITIVO/ultrachat_br_clustred_balanced_v1",
|
| 62 |
+
"dominguesm/Canarim-Instruct-PTBR-Dataset",
|
| 63 |
+
"adalbertojunior/punctuation-ptbr",
|
| 64 |
+
"iara-project/news-articles-ptbr-dataset",
|
| 65 |
+
"manoela/noticias_ptbr",
|
| 66 |
+
"BrunoN-Dev/corpus-ptbr-v1",
|
| 67 |
+
]
|
| 68 |
+
|
| 69 |
+
# V6.5-V2 — Dataset para TREINAMENTO COM PUNIÇÃO (500 samples).
|
| 70 |
+
# O usuário especificou: "dataset para TREINAMENTO COM PUNIÇÃO ATIVA
|
| 71 |
+
# 'BrunoN-Dev/corpus-ptbr-v1'".
|
| 72 |
+
PUNICAO_DATASET_V2 = "BrunoN-Dev/corpus-ptbr-v1"
|
| 73 |
+
|
| 74 |
+
# Dataset format descriptors — Lista priorizada de campos de texto
|
| 75 |
+
DATASET_FORMATS: Dict[str, Dict[str, Any]] = {
|
| 76 |
+
"CEIA-POSITIVO/ultrachat_br_clustred_balanced_v1": {
|
| 77 |
+
"type": "chat",
|
| 78 |
+
"text_fields": ["conversa", "text", "content", "conversation", "messages"],
|
| 79 |
+
"label_fields": [],
|
| 80 |
+
"split": "train",
|
| 81 |
+
"config": None,
|
| 82 |
+
},
|
| 83 |
+
"Madras1/corpus-ptbr-v2": {
|
| 84 |
+
"type": "text",
|
| 85 |
+
"text_fields": ["text", "content", "document", "body"],
|
| 86 |
+
"label_fields": [],
|
| 87 |
+
"split": "train",
|
| 88 |
+
"config": None,
|
| 89 |
+
},
|
| 90 |
+
"rhaymison/multmodal_175k_portuguese": {
|
| 91 |
+
"type": "multimodal",
|
| 92 |
+
"text_fields": ["description", "text", "question", "prompt", "instruction", "input"],
|
| 93 |
+
"image_fields": ["image", "image_url"],
|
| 94 |
+
"label_fields": ["answer", "response", "output"],
|
| 95 |
+
"split": "train",
|
| 96 |
+
"config": None,
|
| 97 |
+
},
|
| 98 |
+
"TucanoBR/GigaVerbo": {
|
| 99 |
+
"type": "text",
|
| 100 |
+
"text_fields": ["text", "content", "document", "body"],
|
| 101 |
+
"label_fields": [],
|
| 102 |
+
"split": "train",
|
| 103 |
+
"config": None,
|
| 104 |
+
},
|
| 105 |
+
"nvidia/OpenMathReasoning": {
|
| 106 |
+
"type": "math",
|
| 107 |
+
"text_fields": ["problem", "question", "input", "prompt"],
|
| 108 |
+
"label_fields": ["solution", "answer", "output", "response"],
|
| 109 |
+
# BUG FIX: OpenMathReasoning has splits ['cot', 'tir', 'genselect', 'additional_problems']
|
| 110 |
+
"split": "cot", # Use 'cot' (chain-of-thought) as primary split
|
| 111 |
+
"config": None,
|
| 112 |
+
},
|
| 113 |
+
"MathLLMs/MathVision": {
|
| 114 |
+
"type": "math_multimodal",
|
| 115 |
+
"text_fields": ["question", "problem", "text", "query"],
|
| 116 |
+
"image_fields": ["image", "image_url"],
|
| 117 |
+
"label_fields": ["answer", "solution", "response"],
|
| 118 |
+
# BUG FIX: MathVision has only 'test' and 'testmini' splits (no train)
|
| 119 |
+
"split": "test",
|
| 120 |
+
"config": None,
|
| 121 |
+
},
|
| 122 |
+
"nvidia/OpenMathInstruct-2": {
|
| 123 |
+
"type": "math",
|
| 124 |
+
"text_fields": ["problem", "question", "input", "prompt"],
|
| 125 |
+
# V3 FIX: campos reais do dataset são generated_solution e expected_answer
|
| 126 |
+
"label_fields": ["generated_solution", "expected_answer", "solution", "answer", "output", "response"],
|
| 127 |
+
"split": "train",
|
| 128 |
+
"config": None,
|
| 129 |
+
},
|
| 130 |
+
"dominguesm/restore-punctuation-ptbr-dataset": {
|
| 131 |
+
"type": "punctuation",
|
| 132 |
+
"text_fields": ["text", "original", "unpunctuated", "input"],
|
| 133 |
+
"label_fields": ["punctuated", "restored", "target", "output"],
|
| 134 |
+
"split": "train",
|
| 135 |
+
"config": None,
|
| 136 |
+
},
|
| 137 |
+
"carolina-c4ai/corpus-carolina": {
|
| 138 |
+
"type": "text",
|
| 139 |
+
"text_fields": ["text", "content", "document", "body", "xml"],
|
| 140 |
+
"label_fields": [],
|
| 141 |
+
"split": "train",
|
| 142 |
+
"config": None,
|
| 143 |
+
# V13.9.2-carolina: dataset original usa script Python (não suportado em
|
| 144 |
+
# datasets 5.0+). Workaround: stream_xml_gz direto via lxml.iterparse,
|
| 145 |
+
# sem usar o script Python. Implementado em _stream_carolina_direct().
|
| 146 |
+
},
|
| 147 |
+
# ── V13.9.2-finetune-v2: novos datasets PT-BR ───────────────────────────
|
| 148 |
+
# Adicionados para fine-tuning contínuo do modelo v13.9.2 (preserva
|
| 149 |
+
# arquitetura e parâmetros). Nenhuma mudança estrutural — apenas novos
|
| 150 |
+
# descritores de formato para que stream_dataset() reconheça os datasets.
|
| 151 |
+
#
|
| 152 |
+
# V13.9.2-finetune-v3: padronização para o formato unificado
|
| 153 |
+
# "### Instruction:\n...\n\n### Response:\n..." conforme requisitado pelo
|
| 154 |
+
# usuário. Apenas os datasets com estrutura instruction/response (orion,
|
| 155 |
+
# cnmoro) usam o template; strak2005/bratao continuam como texto puro
|
| 156 |
+
# (corpus sem estrutura de instrução). Nenhuma mudança nos 9 datasets
|
| 157 |
+
# base — retrocompatibilidade total.
|
| 158 |
+
"orion-research/translations-en_US-pt_BR": {
|
| 159 |
+
"type": "translation",
|
| 160 |
+
# string = EN (instruction), string_translation = PT-BR (response).
|
| 161 |
+
# V13.9.2-finetune-v3: separa instruction/response para o template
|
| 162 |
+
# unificado em vez de concatenar como texto único.
|
| 163 |
+
"text_fields": ["string"],
|
| 164 |
+
"label_fields": ["string_translation"],
|
| 165 |
+
"split": "train",
|
| 166 |
+
"config": None,
|
| 167 |
+
"format_template": "instruction_response",
|
| 168 |
+
"instruction_prefix": "Translate the following text to Portuguese:",
|
| 169 |
+
},
|
| 170 |
+
"cnmoro/Instruct-PTBR-10M": {
|
| 171 |
+
"type": "instruct",
|
| 172 |
+
# NB: campos em MAIÚSCULAS no schema do dataset
|
| 173 |
+
# V13.9.2-finetune-v3: INSTRUCTION -> instruction, RESPONSE -> response
|
| 174 |
+
# via template unificado.
|
| 175 |
+
"text_fields": ["INSTRUCTION"],
|
| 176 |
+
"label_fields": ["RESPONSE"],
|
| 177 |
+
"split": "train",
|
| 178 |
+
"config": None,
|
| 179 |
+
"format_template": "instruction_response",
|
| 180 |
+
# V13.9.2-finetune-v2: dataset é um único parquet de 10 GB — streaming
|
| 181 |
+
# via datasets.load_dataset() é extremamente lento para materializar
|
| 182 |
+
# a primeira amostra. Fallback via HF datasets-server rows API (HTTP),
|
| 183 |
+
# que retorna batches de 100 rows rapidamente. Implementado em
|
| 184 |
+
# _stream_via_rows_api().
|
| 185 |
+
"loader": "rows_api",
|
| 186 |
+
},
|
| 187 |
+
"strak2005/corpus-ptbr-v1": {
|
| 188 |
+
"type": "text",
|
| 189 |
+
"text_fields": ["text", "content"],
|
| 190 |
+
"label_fields": [],
|
| 191 |
+
"split": "train",
|
| 192 |
+
"config": None,
|
| 193 |
+
# Fallback: se strak2005 cair, tenta bratao (mirror idêntico)
|
| 194 |
+
"fallback": "bratao/corpus-ptbr-v1",
|
| 195 |
+
# V13.9.2-finetune-v2: dataset tem 18 shards parquet de ~1GB cada.
|
| 196 |
+
# Streaming via datasets.load_dataset() não materializa primeira
|
| 197 |
+
# amostra em tempo aceitável. Fallback: baixa 1 shard localmente e
|
| 198 |
+
# itera com pyarrow. Implementado em _stream_via_parquet_first_shard().
|
| 199 |
+
"loader": "parquet_first_shard",
|
| 200 |
+
},
|
| 201 |
+
"bratao/corpus-ptbr-v1": {
|
| 202 |
+
"type": "text",
|
| 203 |
+
"text_fields": ["text", "content"],
|
| 204 |
+
"label_fields": [],
|
| 205 |
+
"split": "train",
|
| 206 |
+
"config": None,
|
| 207 |
+
"loader": "parquet_first_shard",
|
| 208 |
+
},
|
| 209 |
+
# ── V6.5: novos datasets para esgotar ────────────────────────────────
|
| 210 |
+
# User requirement: "esgotar 'dominguesm/restore-punctuation-ptbr-dataset'
|
| 211 |
+
# e 'carolina-c4ai/corpus-carolina' e 'nvidia/OpenMathInstruct-2' e
|
| 212 |
+
# 'nvidia/OpenMathReasoning' e 'CEIA-POSITIVO/ultrachat_br_clustred_balanced_v1'
|
| 213 |
+
# e 'Dexavator/English-PTBR'"
|
| 214 |
+
#
|
| 215 |
+
# Dexavator/English-PTBR: dataset de tradução EN->PT-BR.
|
| 216 |
+
# Estrutura típica: {"english": "...", "portuguese": "..."} ou
|
| 217 |
+
# {"en": "...", "pt": "..."}.
|
| 218 |
+
"Dexavator/English-PTBR": {
|
| 219 |
+
"type": "translation",
|
| 220 |
+
"text_fields": ["english", "en", "text", "source"],
|
| 221 |
+
"label_fields": ["portuguese", "pt", "target", "translation"],
|
| 222 |
+
"split": "train",
|
| 223 |
+
"config": None,
|
| 224 |
+
"format_template": "instruction_response",
|
| 225 |
+
"instruction_prefix": "Translate the following text to Portuguese:",
|
| 226 |
+
},
|
| 227 |
+
# ── V6.5-final: Canarim-Instruct-PTBR-Dataset ────────────────────────
|
| 228 |
+
# User requirement: "esgotar 'dominguesm/Canarim-Instruct-PTBR-Dataset'"
|
| 229 |
+
# Canarim-Instruct: dataset PT-BR de instruções (~430k samples).
|
| 230 |
+
# Estrutura típica: {"instruction": "...", "input": "...", "output": "..."}
|
| 231 |
+
# ou {"text": "...", "conversation": [{"role": "user", "content": "..."}, ...]}
|
| 232 |
+
"dominguesm/Canarim-Instruct-PTBR-Dataset": {
|
| 233 |
+
"type": "instruct",
|
| 234 |
+
"text_fields": ["instruction", "input", "text", "prompt", "question"],
|
| 235 |
+
"label_fields": ["output", "response", "answer"],
|
| 236 |
+
"split": "train",
|
| 237 |
+
"config": None,
|
| 238 |
+
"format_template": "instruction_response",
|
| 239 |
+
"instruction_prefix": "Instrução:",
|
| 240 |
+
},
|
| 241 |
+
# ── V6.5-7ds: novos 3 datasets PT-BR ─────────────────────────────────
|
| 242 |
+
# User requirement (latest): "streaming até esgotar nesta sequência
|
| 243 |
+
# 'dominguesm/restore-punctuation-ptbr-dataset' e
|
| 244 |
+
# 'carolina-c4ai/corpus-carolina' e
|
| 245 |
+
# 'CEIA-POSITIVO/ultrachat_br_clustred_balanced_v1' e
|
| 246 |
+
# 'dominguesm/Canarim-Instruct-PTBR-Dataset' e
|
| 247 |
+
# 'adalbertojunior/punctuation-ptbr' e
|
| 248 |
+
# 'iara-project/news-articles-ptbr-dataset' e
|
| 249 |
+
# 'manoela/noticias_ptbr'"
|
| 250 |
+
#
|
| 251 |
+
# adalbertojunior/punctuation-ptbr:
|
| 252 |
+
# Dataset de restauração de pontuação PT-BR. Schema:
|
| 253 |
+
# {id: str, tokens: list[str], pos_tags, chunk_tags, ner_tags}
|
| 254 |
+
# O campo "tokens" é uma lista de palavras (sem pontuação). Precisamos
|
| 255 |
+
# fazer join com espaços para obter texto bruto. Como o dataset usa
|
| 256 |
+
# script Python (datasets 5.0+ não suporta), usamos rows_api loader.
|
| 257 |
+
# Splits: train/validation/test com config "punctuation-ptbr".
|
| 258 |
+
"adalbertojunior/punctuation-ptbr": {
|
| 259 |
+
"type": "punctuation_tokens",
|
| 260 |
+
"text_fields": ["tokens"], # list[str] — join com espaços
|
| 261 |
+
"label_fields": [],
|
| 262 |
+
"split": "train",
|
| 263 |
+
"config": "punctuation-ptbr", # necessário para rows_api
|
| 264 |
+
"loader": "rows_api",
|
| 265 |
+
"join_tokens": True, # flag especial para _normalize_sample
|
| 266 |
+
},
|
| 267 |
+
# iara-project/news-articles-ptbr-dataset:
|
| 268 |
+
# Dataset de notícias PT-BR da Folha de S.Paulo (via IARA project).
|
| 269 |
+
# Schema: {title, text, date, category, category_natural_language, link}
|
| 270 |
+
# Streaming normal via parquet — primeira amostra materializa rápido.
|
| 271 |
+
# "text" é o corpo da notícia, "title" é o título. Concatenamos ambos
|
| 272 |
+
# via format_template="title_text_concat" para enriquecer contexto.
|
| 273 |
+
# O campo "category" é usado como label (categoria da notícia).
|
| 274 |
+
"iara-project/news-articles-ptbr-dataset": {
|
| 275 |
+
"type": "news",
|
| 276 |
+
# text_fields em ordem de prioridade: "text" (corpo) primeiro
|
| 277 |
+
"text_fields": ["text", "title", "content", "body"],
|
| 278 |
+
"label_fields": ["category", "category_natural_language", "title"],
|
| 279 |
+
"split": "train",
|
| 280 |
+
"config": None,
|
| 281 |
+
# V6.5-7ds: title_text_concat usa title como label (prefixo)
|
| 282 |
+
# e text como corpo. Aqui label_fields pega "category" primeiro
|
| 283 |
+
# para dar contexto temático.
|
| 284 |
+
"format_template": "title_text_concat",
|
| 285 |
+
},
|
| 286 |
+
# manoela/noticias_ptbr:
|
| 287 |
+
# Mirror do iara-project/news-articles-ptbr-dataset (mesmo parquet,
|
| 288 |
+
# mesmo schema). Streaming normal via parquet.
|
| 289 |
+
"manoela/noticias_ptbr": {
|
| 290 |
+
"type": "news",
|
| 291 |
+
"text_fields": ["text", "title", "content", "body"],
|
| 292 |
+
"label_fields": ["category", "category_natural_language", "title"],
|
| 293 |
+
"split": "train",
|
| 294 |
+
"config": None,
|
| 295 |
+
"format_template": "title_text_concat",
|
| 296 |
+
},
|
| 297 |
+
# ── V6.5-V2: BrunoN-Dev/corpus-ptbr-v1 ───────────────────────────────
|
| 298 |
+
# User requirement: "dataset para TREINAMENTO COM PUNIÇÃO ATIVA
|
| 299 |
+
# 'BrunoN-Dev/corpus-ptbr-v1'" e o 8º dataset de CONHECIMENTO.
|
| 300 |
+
#
|
| 301 |
+
# BrunoN-Dev/corpus-ptbr-v1: corpus PT-BR genérico com 18 shards parquet
|
| 302 |
+
# (~1GB cada). Schema: {"text": "...", "meta": "..."}.
|
| 303 |
+
# Streaming via datasets.load_dataset(streaming=True) é lento para
|
| 304 |
+
# materializar a primeira amostra — usamos o loader parquet_first_shard
|
| 305 |
+
# (baixa apenas o primeiro shard e itera localmente com pyarrow).
|
| 306 |
+
# Sem label explícita — usamos hash do texto (paridade do nº de palavras)
|
| 307 |
+
# para gerar label binário (0/1) — não é dado sintético, é derivação
|
| 308 |
+
# algorítmica do conteúdo real do dataset.
|
| 309 |
+
"BrunoN-Dev/corpus-ptbr-v1": {
|
| 310 |
+
"type": "text",
|
| 311 |
+
"text_fields": ["text", "content", "body", "document"],
|
| 312 |
+
"label_fields": [],
|
| 313 |
+
"split": "train",
|
| 314 |
+
"config": None,
|
| 315 |
+
"loader": "parquet_first_shard",
|
| 316 |
+
# V6.5-V2: gera label binário determinístico via hash do texto
|
| 317 |
+
# (paridade do número de palavras) — não é dado sintético, é
|
| 318 |
+
# derivação algorítmica do conteúdo real do dataset.
|
| 319 |
+
"label_strategy": "hash_parity",
|
| 320 |
+
},
|
| 321 |
+
}
|
| 322 |
+
|
| 323 |
+
|
| 324 |
+
# Lista de arquivos XML.gz do Carolina que tentaremos baixar (ordem: menores
|
| 325 |
+
# taxonomias primeiro para minimizar banda e RAM). Selecionamos 2 arquivos
|
| 326 |
+
# por taxonomia para garantir diversidade sem explodir a memória.
|
| 327 |
+
CAROLINA_XML_FILES = [
|
| 328 |
+
# taxonomia "datasets_and_other_corpora" (DAT = pt-BR notícias)
|
| 329 |
+
"corpus/datasets_and_other_corpora/pt-BR/DATa.xml.gz",
|
| 330 |
+
"corpus/datasets_and_other_corpora/pt-BR/DATb.xml.gz",
|
| 331 |
+
# taxonomia "wik" (wikis)
|
| 332 |
+
"corpus/wik/WIKa.xml.gz",
|
| 333 |
+
"corpus/wik/WIKb.xml.gz",
|
| 334 |
+
# taxonomia "uni" (university_domains)
|
| 335 |
+
"corpus/uni/UNIa.xml.gz",
|
| 336 |
+
"corpus/uni/UNIb.xml.gz",
|
| 337 |
+
# taxonomia "pub" (public_domain_works)
|
| 338 |
+
"corpus/pub/PUBa.xml.gz",
|
| 339 |
+
"corpus/pub/PUBb.xml.gz",
|
| 340 |
+
]
|
| 341 |
+
|
| 342 |
+
|
| 343 |
+
def _stream_carolina_direct(
|
| 344 |
+
max_samples: int,
|
| 345 |
+
hf_token: Optional[str] = None,
|
| 346 |
+
timeout_per_file: int = 60,
|
| 347 |
+
) -> Iterator[Dict[str, Any]]:
|
| 348 |
+
"""Workaround V13.9.2-carolina: carrega Carolina sem usar o script Python.
|
| 349 |
+
|
| 350 |
+
Baixa arquivos XML.gz diretamente do repo e faz parse streaming com
|
| 351 |
+
lxml.etree.iterparse (compatível com datasets 5.0+). Cada <TEI> vira um
|
| 352 |
+
documento de texto (concatenação dos <p> dentro de <body>).
|
| 353 |
+
|
| 354 |
+
Args:
|
| 355 |
+
max_samples: número máximo de amostras a retornar
|
| 356 |
+
hf_token: token HF opcional (dataset é público, mas token evita rate limit)
|
| 357 |
+
timeout_per_file: tempo máximo (segundos) por arquivo XML
|
| 358 |
+
|
| 359 |
+
Yields:
|
| 360 |
+
dict com chave "text" (compatível com _normalize_sample)
|
| 361 |
+
"""
|
| 362 |
+
from huggingface_hub import hf_hub_download
|
| 363 |
+
import gzip
|
| 364 |
+
from lxml import etree
|
| 365 |
+
|
| 366 |
+
TEI_NS = "{http://www.tei-c.org/ns/1.0}"
|
| 367 |
+
BODY_P_TAG = f".//{TEI_NS}body/{TEI_NS}p"
|
| 368 |
+
|
| 369 |
+
count = 0
|
| 370 |
+
for xml_path in CAROLINA_XML_FILES:
|
| 371 |
+
if count >= max_samples:
|
| 372 |
+
return
|
| 373 |
+
try:
|
| 374 |
+
local_path = hf_hub_download(
|
| 375 |
+
repo_id="carolina-c4ai/corpus-carolina",
|
| 376 |
+
filename=xml_path,
|
| 377 |
+
repo_type="dataset",
|
| 378 |
+
token=hf_token,
|
| 379 |
+
)
|
| 380 |
+
except Exception as e:
|
| 381 |
+
logger.warning(f"Carolina: não foi possível baixar {xml_path}: {e}")
|
| 382 |
+
continue
|
| 383 |
+
|
| 384 |
+
try:
|
| 385 |
+
with gzip.open(local_path, "rb") as gz:
|
| 386 |
+
# iterparse streaming: não carrega árvore inteira na memória
|
| 387 |
+
for _, tei in etree.iterparse(
|
| 388 |
+
gz, huge_tree=True, encoding="utf-8", tag=f"{TEI_NS}TEI"
|
| 389 |
+
):
|
| 390 |
+
if count >= max_samples:
|
| 391 |
+
tei.clear()
|
| 392 |
+
return
|
| 393 |
+
# Extrai texto dos <p> dentro de <body>
|
| 394 |
+
parts = []
|
| 395 |
+
for p in tei.findall(BODY_P_TAG):
|
| 396 |
+
if p.text:
|
| 397 |
+
parts.append(p.text)
|
| 398 |
+
text = " ".join(parts).strip()
|
| 399 |
+
if text and len(text) >= 10:
|
| 400 |
+
yield {"text": text, "meta": ""}
|
| 401 |
+
tei.clear() # libera memória da árvore TEI
|
| 402 |
+
count += 1
|
| 403 |
+
except Exception as e:
|
| 404 |
+
logger.warning(f"Carolina: erro processando {xml_path}: {e}")
|
| 405 |
+
continue
|
| 406 |
+
|
| 407 |
+
logger.info(f"Carolina: streamou {count} amostras no total")
|
| 408 |
+
|
| 409 |
+
|
| 410 |
+
# ═══════════════════════════════════════════════════════════════════════════
|
| 411 |
+
# V13.9.2-finetune-v2: Loaders alternativos para datasets parquet grandes
|
| 412 |
+
# onde datasets.load_dataset(streaming=True) é inviável (primeira amostra
|
| 413 |
+
# demora minutos para materializar).
|
| 414 |
+
# ═══════════════════════════════════════════════════════════════════════════
|
| 415 |
+
|
| 416 |
+
def _stream_via_rows_api(
|
| 417 |
+
dataset_name: str,
|
| 418 |
+
max_samples: int,
|
| 419 |
+
hf_token: Optional[str] = None,
|
| 420 |
+
batch_size: int = 100,
|
| 421 |
+
max_offset: int = 100_000,
|
| 422 |
+
config: Optional[str] = None,
|
| 423 |
+
split: str = "train",
|
| 424 |
+
) -> Iterator[Dict[str, Any]]:
|
| 425 |
+
"""Carrega amostras via HF datasets-server /rows API (HTTP).
|
| 426 |
+
|
| 427 |
+
Usado para datasets cujo parquet é muito grande para streaming eficiente
|
| 428 |
+
(ex: cnmoro/Instruct-PTBR-10M com 10 GB em arquivo único) OU datasets
|
| 429 |
+
que usam script Python como loader (não suportado em datasets 5.0+,
|
| 430 |
+
ex: adalbertojunior/punctuation-ptbr). A rows API retorna batches de
|
| 431 |
+
até 100 rows via HTTP, sem precisar baixar o parquet.
|
| 432 |
+
|
| 433 |
+
Args:
|
| 434 |
+
dataset_name: repo_id do dataset (ex: "cnmoro/Instruct-PTBR-10M")
|
| 435 |
+
max_samples: nº máximo de amostras a retornar
|
| 436 |
+
hf_token: token HF (opcional — dataset público)
|
| 437 |
+
batch_size: nº de rows por requisição (máx 100)
|
| 438 |
+
max_offset: offset máximo a tentar (sai do loop se excedido)
|
| 439 |
+
config: nome da config do dataset (default: "default"). NECESSÁRIO
|
| 440 |
+
para datasets cuja config não é "default" (ex:
|
| 441 |
+
adalbertojunior/punctuation-ptbr usa "punctuation-ptbr").
|
| 442 |
+
split: nome do split (default: "train").
|
| 443 |
+
|
| 444 |
+
Yields:
|
| 445 |
+
dict com campos do dataset (compatível com _normalize_sample)
|
| 446 |
+
"""
|
| 447 |
+
import json as _json
|
| 448 |
+
import urllib.request
|
| 449 |
+
import urllib.parse
|
| 450 |
+
|
| 451 |
+
base_url = "https://datasets-server.huggingface.co/rows"
|
| 452 |
+
config_name = config or "default"
|
| 453 |
+
count = 0
|
| 454 |
+
offset = 0
|
| 455 |
+
|
| 456 |
+
while count < max_samples and offset < max_offset:
|
| 457 |
+
params = urllib.parse.urlencode({
|
| 458 |
+
"dataset": dataset_name,
|
| 459 |
+
"config": config_name,
|
| 460 |
+
"split": split,
|
| 461 |
+
"offset": offset,
|
| 462 |
+
"length": min(batch_size, max_samples - count),
|
| 463 |
+
})
|
| 464 |
+
url = f"{base_url}?{params}"
|
| 465 |
+
try:
|
| 466 |
+
req = urllib.request.Request(url)
|
| 467 |
+
if hf_token:
|
| 468 |
+
req.add_header("Authorization", f"Bearer {hf_token}")
|
| 469 |
+
with urllib.request.urlopen(req, timeout=30) as r:
|
| 470 |
+
data = _json.loads(r.read())
|
| 471 |
+
except Exception as e:
|
| 472 |
+
logger.warning(f"rows_api {dataset_name} config={config_name} offset={offset}: {e}")
|
| 473 |
+
break
|
| 474 |
+
|
| 475 |
+
rows = data.get("rows", [])
|
| 476 |
+
if not rows:
|
| 477 |
+
break
|
| 478 |
+
|
| 479 |
+
for row in rows:
|
| 480 |
+
if count >= max_samples:
|
| 481 |
+
return
|
| 482 |
+
row_data = row.get("row", {})
|
| 483 |
+
# Remove campos internos como __index_level_0__
|
| 484 |
+
clean = {k: v for k, v in row_data.items()
|
| 485 |
+
if not k.startswith("__")}
|
| 486 |
+
yield clean
|
| 487 |
+
count += 1
|
| 488 |
+
|
| 489 |
+
offset += len(rows)
|
| 490 |
+
if len(rows) < batch_size:
|
| 491 |
+
break # chegou ao fim do dataset
|
| 492 |
+
|
| 493 |
+
logger.info(f"rows_api {dataset_name} config={config_name}: streamou {count} amostras")
|
| 494 |
+
|
| 495 |
+
|
| 496 |
+
def _stream_via_parquet_first_shard(
|
| 497 |
+
dataset_name: str,
|
| 498 |
+
max_samples: int,
|
| 499 |
+
hf_token: Optional[str] = None,
|
| 500 |
+
) -> Iterator[Dict[str, Any]]:
|
| 501 |
+
"""Baixa primeiro shard parquet do dataset e itera localmente com pyarrow.
|
| 502 |
+
|
| 503 |
+
Usado para datasets com múltiplos shards parquet grandes (ex:
|
| 504 |
+
strak2005/corpus-ptbr-v1 com 18 shards de ~1GB). Streaming via
|
| 505 |
+
datasets.load_dataset() não materializa primeira amostra em tempo
|
| 506 |
+
aceitável. Esta função baixa o primeiro shard (~26s para 1GB) e itera
|
| 507 |
+
localmente — rápido após download.
|
| 508 |
+
|
| 509 |
+
Args:
|
| 510 |
+
dataset_name: repo_id do dataset
|
| 511 |
+
max_samples: nº máximo de amostras
|
| 512 |
+
hf_token: token HF opcional
|
| 513 |
+
|
| 514 |
+
Yields:
|
| 515 |
+
dict com campos do dataset
|
| 516 |
+
"""
|
| 517 |
+
from huggingface_hub import HfApi, hf_hub_download
|
| 518 |
+
import pyarrow.parquet as pq
|
| 519 |
+
|
| 520 |
+
api = HfApi(token=hf_token)
|
| 521 |
+
|
| 522 |
+
# Lista arquivos parquet no repo
|
| 523 |
+
try:
|
| 524 |
+
info = api.dataset_info(dataset_name)
|
| 525 |
+
except Exception as e:
|
| 526 |
+
logger.warning(f"parquet_shard {dataset_name}: dataset_info failed: {e}")
|
| 527 |
+
return
|
| 528 |
+
|
| 529 |
+
parquet_files = sorted([
|
| 530 |
+
s.rfilename for s in info.siblings
|
| 531 |
+
if s.rfilename.endswith(".parquet")
|
| 532 |
+
])
|
| 533 |
+
if not parquet_files:
|
| 534 |
+
logger.warning(f"parquet_shard {dataset_name}: nenhum .parquet encontrado")
|
| 535 |
+
return
|
| 536 |
+
|
| 537 |
+
# Baixa apenas o primeiro shard
|
| 538 |
+
first_shard = parquet_files[0]
|
| 539 |
+
logger.info(f"parquet_shard {dataset_name}: baixando {first_shard}...")
|
| 540 |
+
try:
|
| 541 |
+
local_path = hf_hub_download(
|
| 542 |
+
repo_id=dataset_name,
|
| 543 |
+
filename=first_shard,
|
| 544 |
+
repo_type="dataset",
|
| 545 |
+
token=hf_token,
|
| 546 |
+
)
|
| 547 |
+
except Exception as e:
|
| 548 |
+
logger.warning(f"parquet_shard {dataset_name}: download failed: {e}")
|
| 549 |
+
return
|
| 550 |
+
|
| 551 |
+
# Itera com pyarrow.iter_batches — método eficiente que lê batches pequenos
|
| 552 |
+
# do parquet sem materializar o row group inteiro na memória.
|
| 553 |
+
# BUG FIX V13.9.2-finetune-v2: read_row_group() materializa TODAS as ~1M
|
| 554 |
+
# rows de uma vez (lento + ~1GB RAM). iter_batches(batch_size=10) lê só
|
| 555 |
+
# 10 rows por vez (rápido + ~10KB RAM).
|
| 556 |
+
try:
|
| 557 |
+
pf = pq.ParquetFile(local_path)
|
| 558 |
+
count = 0
|
| 559 |
+
# batch_read_size: nº de rows por batch (pequeno = menos RAM)
|
| 560 |
+
for batch in pf.iter_batches(batch_size=100):
|
| 561 |
+
if count >= max_samples:
|
| 562 |
+
break
|
| 563 |
+
cols = batch.schema.names
|
| 564 |
+
# to_pylist() é eficiente para batches pequenos
|
| 565 |
+
pydict = batch.to_pydict()
|
| 566 |
+
n_rows = len(pydict[cols[0]]) if cols else 0
|
| 567 |
+
for i in range(n_rows):
|
| 568 |
+
if count >= max_samples:
|
| 569 |
+
break
|
| 570 |
+
row = {col: pydict[col][i] for col in cols}
|
| 571 |
+
yield row
|
| 572 |
+
count += 1
|
| 573 |
+
except Exception as e:
|
| 574 |
+
logger.warning(f"parquet_shard {dataset_name}: pyarrow iter failed: {e}")
|
| 575 |
+
return
|
| 576 |
+
|
| 577 |
+
logger.info(f"parquet_shard {dataset_name}: streamou {count} amostras do shard {first_shard}")
|
| 578 |
+
|
| 579 |
+
|
| 580 |
+
@dataclass
|
| 581 |
+
class ProcessedSample:
|
| 582 |
+
"""Uma amostra processada: texto bruto + metadados."""
|
| 583 |
+
raw_text: str
|
| 584 |
+
dataset_name: str
|
| 585 |
+
sample_idx: int
|
| 586 |
+
text_field: str = ""
|
| 587 |
+
label_text: str = ""
|
| 588 |
+
|
| 589 |
+
|
| 590 |
+
def _extract_field(sample: Dict[str, Any], candidate_fields: List[str]) -> Optional[str]:
|
| 591 |
+
"""Extrai texto da primeira field disponível no sample.
|
| 592 |
+
|
| 593 |
+
Suporta:
|
| 594 |
+
- strings diretas
|
| 595 |
+
- listas de strings (join com espaço — útil para datasets tokenizados
|
| 596 |
+
como adalbertojunior/punctuation-ptbr onde "tokens": ["hoje","a",...])
|
| 597 |
+
- listas de dicts (chat messages — concatena role: content)
|
| 598 |
+
- dicts aninhados (pega text/content/body)
|
| 599 |
+
"""
|
| 600 |
+
for field_name in candidate_fields:
|
| 601 |
+
if field_name in sample:
|
| 602 |
+
val = sample[field_name]
|
| 603 |
+
if val is None:
|
| 604 |
+
continue
|
| 605 |
+
if isinstance(val, str):
|
| 606 |
+
if val.strip():
|
| 607 |
+
return val
|
| 608 |
+
elif isinstance(val, list):
|
| 609 |
+
# V6.5-7ds: se for lista de strings, join com espaços
|
| 610 |
+
# (ex: adalbertojunior/punctuation-ptbr tokens field)
|
| 611 |
+
string_parts = []
|
| 612 |
+
chat_parts = []
|
| 613 |
+
all_strings = True
|
| 614 |
+
for item in val:
|
| 615 |
+
if isinstance(item, str):
|
| 616 |
+
if item.strip():
|
| 617 |
+
string_parts.append(item)
|
| 618 |
+
elif isinstance(item, (int, float)):
|
| 619 |
+
# tokens podem vir como ints (token IDs) — skip
|
| 620 |
+
all_strings = False
|
| 621 |
+
break
|
| 622 |
+
elif isinstance(item, dict):
|
| 623 |
+
all_strings = False
|
| 624 |
+
role = item.get("role", "")
|
| 625 |
+
content = item.get("content", "")
|
| 626 |
+
if isinstance(content, str) and content.strip():
|
| 627 |
+
chat_parts.append(f"{role}: {content}")
|
| 628 |
+
else:
|
| 629 |
+
all_strings = False
|
| 630 |
+
if all_strings and string_parts:
|
| 631 |
+
joined = " ".join(string_parts).strip()
|
| 632 |
+
if joined:
|
| 633 |
+
return joined
|
| 634 |
+
if chat_parts:
|
| 635 |
+
return "\n".join(chat_parts)
|
| 636 |
+
elif isinstance(val, dict):
|
| 637 |
+
# Pode ser nested
|
| 638 |
+
text = val.get("text") or val.get("content") or val.get("body")
|
| 639 |
+
if isinstance(text, str) and text.strip():
|
| 640 |
+
return text
|
| 641 |
+
return None
|
| 642 |
+
|
| 643 |
+
|
| 644 |
+
def _normalize_sample(sample: Dict[str, Any], dataset_name: str, idx: int) -> Optional[ProcessedSample]:
|
| 645 |
+
"""Normaliza uma amostra bruta em ProcessedSample."""
|
| 646 |
+
fmt = DATASET_FORMATS.get(dataset_name, {})
|
| 647 |
+
text_fields = fmt.get("text_fields", ["text"])
|
| 648 |
+
label_fields = fmt.get("label_fields", [])
|
| 649 |
+
|
| 650 |
+
text = _extract_field(sample, text_fields)
|
| 651 |
+
if not text or len(text.strip()) < 10: # skip very short
|
| 652 |
+
return None
|
| 653 |
+
|
| 654 |
+
label = ""
|
| 655 |
+
if label_fields:
|
| 656 |
+
label = _extract_field(sample, label_fields) or ""
|
| 657 |
+
|
| 658 |
+
# Se tem label, aplica template apropriado
|
| 659 |
+
template = fmt.get("format_template")
|
| 660 |
+
full_text = text
|
| 661 |
+
|
| 662 |
+
# V13.9.2-finetune-v3: para o template instruction_response, exige ambos
|
| 663 |
+
# os lados (instruction + response). Samples sem response são descartados
|
| 664 |
+
# para evitar treinar em texto parcial.
|
| 665 |
+
if template == "instruction_response" and not label:
|
| 666 |
+
return None
|
| 667 |
+
|
| 668 |
+
if label:
|
| 669 |
+
# V13.9.2-finetune-v3: format_template unificado "### Instruction:/### Response:"
|
| 670 |
+
# quando o dataset declarar explicitamente. Caso contrário, mantém o
|
| 671 |
+
# comportamento legado (concatenação simples text\nlabel) para os 9
|
| 672 |
+
# datasets base — retrocompatibilidade total.
|
| 673 |
+
if template == "instruction_response":
|
| 674 |
+
instr_prefix = fmt.get("instruction_prefix", "")
|
| 675 |
+
if instr_prefix:
|
| 676 |
+
full_text = (
|
| 677 |
+
f"### Instruction:\n{instr_prefix}\n{text}\n\n"
|
| 678 |
+
f"### Response:\n{label}"
|
| 679 |
+
)
|
| 680 |
+
else:
|
| 681 |
+
full_text = (
|
| 682 |
+
f"### Instruction:\n{text}\n\n"
|
| 683 |
+
f"### Response:\n{label}"
|
| 684 |
+
)
|
| 685 |
+
elif template == "title_text_concat":
|
| 686 |
+
# V6.5-7ds: para datasets de notícias (iara-project, manoela),
|
| 687 |
+
# o campo "text" (corpo da notícia) é o conteúdo principal e
|
| 688 |
+
# "category" (categoria) é o contexto temático. Formatamos como
|
| 689 |
+
# "TEXTO: <corpo>\nCATEGORIA: <categoria>".
|
| 690 |
+
full_text = f"TEXTO: {text}\nCATEGORIA: {label}"
|
| 691 |
+
else:
|
| 692 |
+
# Comportamento legado (9 datasets base): concatenação simples
|
| 693 |
+
full_text = f"{text}\n{label}"
|
| 694 |
+
|
| 695 |
+
return ProcessedSample(
|
| 696 |
+
raw_text=full_text,
|
| 697 |
+
dataset_name=dataset_name,
|
| 698 |
+
sample_idx=idx,
|
| 699 |
+
text_field=text_fields[0] if text_fields else "",
|
| 700 |
+
label_text=label,
|
| 701 |
+
)
|
| 702 |
+
|
| 703 |
+
|
| 704 |
+
def _load_dataset_streaming(
|
| 705 |
+
dataset_name: str,
|
| 706 |
+
split: str = "train",
|
| 707 |
+
config: Optional[str] = None,
|
| 708 |
+
hf_token: Optional[str] = None,
|
| 709 |
+
):
|
| 710 |
+
"""Carrega dataset em modo streaming.
|
| 711 |
+
|
| 712 |
+
BUG FIX V13.9.1: Removido trust_remote_code (deprecated).
|
| 713 |
+
Se o dataset não carregar com split especificado, tenta 'train' como fallback.
|
| 714 |
+
BUG FIX V13.9.2: Adicionado fallback para datasets com script (carolina-c4ai).
|
| 715 |
+
"""
|
| 716 |
+
from datasets import load_dataset
|
| 717 |
+
|
| 718 |
+
# Lista de splits para tentar, em ordem
|
| 719 |
+
splits_to_try = [split]
|
| 720 |
+
if split != "train":
|
| 721 |
+
splits_to_try.append("train")
|
| 722 |
+
# Adiciona outros splits comuns
|
| 723 |
+
for s in ["test", "validation", "cot", "tir"]:
|
| 724 |
+
if s not in splits_to_try:
|
| 725 |
+
splits_to_try.append(s)
|
| 726 |
+
|
| 727 |
+
fmt = DATASET_FORMATS.get(dataset_name, {})
|
| 728 |
+
fallback_name = fmt.get("fallback")
|
| 729 |
+
datasets_to_try = [dataset_name]
|
| 730 |
+
if fallback_name:
|
| 731 |
+
datasets_to_try.append(fallback_name)
|
| 732 |
+
|
| 733 |
+
for try_ds_name in datasets_to_try:
|
| 734 |
+
for try_split in splits_to_try:
|
| 735 |
+
try:
|
| 736 |
+
if config:
|
| 737 |
+
ds = load_dataset(
|
| 738 |
+
try_ds_name, config, split=try_split, streaming=True,
|
| 739 |
+
token=hf_token,
|
| 740 |
+
)
|
| 741 |
+
else:
|
| 742 |
+
ds = load_dataset(
|
| 743 |
+
try_ds_name, split=try_split, streaming=True,
|
| 744 |
+
token=hf_token,
|
| 745 |
+
)
|
| 746 |
+
logger.info(f"Loaded {try_ds_name} split={try_split}")
|
| 747 |
+
return ds
|
| 748 |
+
except Exception as e:
|
| 749 |
+
err_str = str(e).lower()
|
| 750 |
+
# If "Bad split" error, try next split
|
| 751 |
+
if "bad split" in err_str or "available splits" in err_str:
|
| 752 |
+
logger.info(f" Split {try_split} not available for {try_ds_name}, trying next")
|
| 753 |
+
continue
|
| 754 |
+
# If "scripts no longer supported", try next dataset (fallback)
|
| 755 |
+
if "scripts are no longer supported" in err_str and fallback_name:
|
| 756 |
+
logger.info(f" {try_ds_name} uses script — trying fallback {fallback_name}")
|
| 757 |
+
break
|
| 758 |
+
# For other errors, log and try next split
|
| 759 |
+
logger.warning(f"Failed to load {try_ds_name} split={try_split}: {str(e)[:120]}")
|
| 760 |
+
continue
|
| 761 |
+
|
| 762 |
+
logger.warning(f"All split attempts failed for {dataset_name} (will be skipped)")
|
| 763 |
+
return None
|
| 764 |
+
|
| 765 |
+
|
| 766 |
+
def stream_dataset(
|
| 767 |
+
dataset_name: str,
|
| 768 |
+
max_samples: int = 500,
|
| 769 |
+
hf_token: Optional[str] = None,
|
| 770 |
+
) -> Iterator[ProcessedSample]:
|
| 771 |
+
"""Faz streaming de um dataset, retornando até max_samples amostras.
|
| 772 |
+
|
| 773 |
+
Args:
|
| 774 |
+
dataset_name: nome do dataset no HuggingFace
|
| 775 |
+
max_samples: número máximo de amostras a retornar
|
| 776 |
+
hf_token: token HF opcional
|
| 777 |
+
|
| 778 |
+
Yields:
|
| 779 |
+
ProcessedSample
|
| 780 |
+
"""
|
| 781 |
+
fmt = DATASET_FORMATS.get(dataset_name)
|
| 782 |
+
if fmt is None:
|
| 783 |
+
logger.error(f"Unknown dataset: {dataset_name}")
|
| 784 |
+
return
|
| 785 |
+
|
| 786 |
+
# ── V13.9.2-carolina: bypass do script Python ─────────────────────────
|
| 787 |
+
# O dataset carolina-c4ai/corpus-carolina usa um script Python como loader,
|
| 788 |
+
# que não é suportado por datasets 5.0+. Streamamos os XML.gz diretamente.
|
| 789 |
+
if dataset_name == "carolina-c4ai/corpus-carolina":
|
| 790 |
+
logger.info(f"Carolina: usando carregador direto (bypass do script Python)")
|
| 791 |
+
count = 0
|
| 792 |
+
try:
|
| 793 |
+
for raw_sample in _stream_carolina_direct(
|
| 794 |
+
max_samples=max_samples, hf_token=hf_token
|
| 795 |
+
):
|
| 796 |
+
processed = _normalize_sample(raw_sample, dataset_name, count)
|
| 797 |
+
if processed is not None:
|
| 798 |
+
yield processed
|
| 799 |
+
count += 1
|
| 800 |
+
except Exception as e:
|
| 801 |
+
logger.warning(f"Carolina direct stream error at sample {count}: {e}")
|
| 802 |
+
return
|
| 803 |
+
logger.info(f"Streamed {count} samples from {dataset_name} (direct XML)")
|
| 804 |
+
return
|
| 805 |
+
|
| 806 |
+
# ── V13.9.2-finetune-v2: loaders alternativos para parquets grandes ───
|
| 807 |
+
# Alguns datasets (cnmoro, strak2005/bratao) têm parquets tão grandes que
|
| 808 |
+
# datasets.load_dataset(streaming=True) demora minutos para materializar
|
| 809 |
+
# a primeira amostra. Usamos loaders alternativos especificados no campo
|
| 810 |
+
# "loader" do formato.
|
| 811 |
+
loader = fmt.get("loader")
|
| 812 |
+
if loader == "rows_api":
|
| 813 |
+
logger.info(f"{dataset_name}: usando rows_api loader (datasets-server HTTP)")
|
| 814 |
+
count = 0
|
| 815 |
+
try:
|
| 816 |
+
for raw_sample in _stream_via_rows_api(
|
| 817 |
+
dataset_name, max_samples=max_samples, hf_token=hf_token,
|
| 818 |
+
config=fmt.get("config"),
|
| 819 |
+
split=fmt.get("split", "train"),
|
| 820 |
+
):
|
| 821 |
+
processed = _normalize_sample(raw_sample, dataset_name, count)
|
| 822 |
+
if processed is not None:
|
| 823 |
+
yield processed
|
| 824 |
+
count += 1
|
| 825 |
+
except Exception as e:
|
| 826 |
+
logger.warning(f"rows_api stream error at sample {count}: {e}")
|
| 827 |
+
return
|
| 828 |
+
logger.info(f"Streamed {count} samples from {dataset_name} (rows_api)")
|
| 829 |
+
return
|
| 830 |
+
|
| 831 |
+
if loader == "parquet_first_shard":
|
| 832 |
+
logger.info(f"{dataset_name}: usando parquet_first_shard loader")
|
| 833 |
+
count = 0
|
| 834 |
+
try:
|
| 835 |
+
for raw_sample in _stream_via_parquet_first_shard(
|
| 836 |
+
dataset_name, max_samples=max_samples, hf_token=hf_token
|
| 837 |
+
):
|
| 838 |
+
processed = _normalize_sample(raw_sample, dataset_name, count)
|
| 839 |
+
if processed is not None:
|
| 840 |
+
yield processed
|
| 841 |
+
count += 1
|
| 842 |
+
except Exception as e:
|
| 843 |
+
logger.warning(f"parquet_shard stream error at sample {count}: {e}")
|
| 844 |
+
# Tenta fallback (ex: strak2005 -> bratao)
|
| 845 |
+
fallback_name = fmt.get("fallback")
|
| 846 |
+
if fallback_name and fallback_name != dataset_name:
|
| 847 |
+
logger.info(f"Tentando fallback: {fallback_name}")
|
| 848 |
+
fallback_fmt = DATASET_FORMATS.get(fallback_name, {})
|
| 849 |
+
# Recursão respeitando o loader do fallback
|
| 850 |
+
yield from stream_dataset(
|
| 851 |
+
fallback_name, max_samples=max_samples, hf_token=hf_token
|
| 852 |
+
)
|
| 853 |
+
return
|
| 854 |
+
logger.info(f"Streamed {count} samples from {dataset_name} (parquet_first_shard)")
|
| 855 |
+
return
|
| 856 |
+
|
| 857 |
+
ds = _load_dataset_streaming(
|
| 858 |
+
dataset_name,
|
| 859 |
+
split=fmt.get("split", "train"),
|
| 860 |
+
config=fmt.get("config"),
|
| 861 |
+
hf_token=hf_token,
|
| 862 |
+
)
|
| 863 |
+
if ds is None:
|
| 864 |
+
return
|
| 865 |
+
|
| 866 |
+
count = 0
|
| 867 |
+
try:
|
| 868 |
+
for idx, sample in enumerate(ds):
|
| 869 |
+
if count >= max_samples:
|
| 870 |
+
break
|
| 871 |
+
processed = _normalize_sample(sample, dataset_name, idx)
|
| 872 |
+
if processed is not None:
|
| 873 |
+
yield processed
|
| 874 |
+
count += 1
|
| 875 |
+
except Exception as e:
|
| 876 |
+
logger.warning(f"Error streaming {dataset_name} at sample {idx}: {e}")
|
| 877 |
+
return
|
| 878 |
+
|
| 879 |
+
logger.info(f"Streamed {count} samples from {dataset_name}")
|
| 880 |
+
|
| 881 |
+
|
| 882 |
+
def stream_all_datasets(
|
| 883 |
+
datasets: Optional[List[str]] = None,
|
| 884 |
+
max_samples_per_dataset: int = 500,
|
| 885 |
+
hf_token: Optional[str] = None,
|
| 886 |
+
) -> Iterator[ProcessedSample]:
|
| 887 |
+
"""Faz streaming sequencial de múltiplos datasets.
|
| 888 |
+
|
| 889 |
+
Args:
|
| 890 |
+
datasets: lista de nomes (default: DEFAULT_DATASETS)
|
| 891 |
+
max_samples_per_dataset: limite por dataset
|
| 892 |
+
hf_token: token HF opcional
|
| 893 |
+
|
| 894 |
+
Yields:
|
| 895 |
+
ProcessedSample
|
| 896 |
+
"""
|
| 897 |
+
if datasets is None:
|
| 898 |
+
datasets = DEFAULT_DATASETS
|
| 899 |
+
|
| 900 |
+
for ds_name in datasets:
|
| 901 |
+
logger.info(f"Starting dataset: {ds_name}")
|
| 902 |
+
yield from stream_dataset(
|
| 903 |
+
ds_name,
|
| 904 |
+
max_samples=max_samples_per_dataset,
|
| 905 |
+
hf_token=hf_token,
|
| 906 |
+
)
|
| 907 |
+
|
| 908 |
+
|
| 909 |
+
def collate_samples_to_tensors(
|
| 910 |
+
samples: List[ProcessedSample],
|
| 911 |
+
tokenizer,
|
| 912 |
+
max_seq_len: int = 64,
|
| 913 |
+
pad_token_id: int = 0,
|
| 914 |
+
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
| 915 |
+
"""Converte lista de ProcessedSample em tensors para treino.
|
| 916 |
+
|
| 917 |
+
BUG FIX V13.9.1: Tokenizer.encode() retorna 'Encoding' object, não list.
|
| 918 |
+
Use .ids attribute to get token IDs.
|
| 919 |
+
|
| 920 |
+
Returns:
|
| 921 |
+
input_ids: (B, T)
|
| 922 |
+
attention_mask: (B, T)
|
| 923 |
+
labels: (B, T) — same as input_ids (LM training)
|
| 924 |
+
"""
|
| 925 |
+
batch_input_ids = []
|
| 926 |
+
batch_attn_mask = []
|
| 927 |
+
|
| 928 |
+
for sample in samples:
|
| 929 |
+
# Tokenize
|
| 930 |
+
encoding = tokenizer.encode(sample.raw_text, add_special_tokens=True)
|
| 931 |
+
|
| 932 |
+
# BUG FIX: Handle both `tokenizers.Encoding` and `transformers` tokenizer outputs
|
| 933 |
+
if hasattr(encoding, 'ids'):
|
| 934 |
+
tokens = encoding.ids # tokenizers library
|
| 935 |
+
elif isinstance(encoding, list):
|
| 936 |
+
tokens = encoding # already a list of ints
|
| 937 |
+
elif hasattr(encoding, 'input_ids'):
|
| 938 |
+
tokens = encoding.input_ids # transformers BatchEncoding
|
| 939 |
+
else:
|
| 940 |
+
tokens = list(encoding) # fallback
|
| 941 |
+
|
| 942 |
+
# Truncate to max_seq_len
|
| 943 |
+
if len(tokens) > max_seq_len:
|
| 944 |
+
tokens = tokens[:max_seq_len]
|
| 945 |
+
# Pad
|
| 946 |
+
attn_mask = [1] * len(tokens) + [0] * (max_seq_len - len(tokens))
|
| 947 |
+
tokens = list(tokens) + [pad_token_id] * (max_seq_len - len(tokens))
|
| 948 |
+
batch_input_ids.append(tokens)
|
| 949 |
+
batch_attn_mask.append(attn_mask)
|
| 950 |
+
|
| 951 |
+
input_ids = torch.tensor(batch_input_ids, dtype=torch.long)
|
| 952 |
+
attention_mask = torch.tensor(batch_attn_mask, dtype=torch.long)
|
| 953 |
+
labels = input_ids.clone()
|
| 954 |
+
|
| 955 |
+
return input_ids, attention_mask, labels
|
| 956 |
+
|
| 957 |
+
|
| 958 |
+
# ═══════════════════════════════════════════════════════════════════════════
|
| 959 |
+
# Self-test (NÃO enviar para HuggingFace)
|
| 960 |
+
# ═══════════════════════════════════════════════════════════════════════════
|
| 961 |
+
if __name__ == "__main__":
|
| 962 |
+
print("=== Streaming Datasets V13.9.1 Self-Test ===\n")
|
| 963 |
+
|
| 964 |
+
# Test 1: Test field extraction
|
| 965 |
+
print("Test 1: Field extraction")
|
| 966 |
+
sample_chat = {
|
| 967 |
+
"messages": [
|
| 968 |
+
{"role": "user", "content": "Olá"},
|
| 969 |
+
{"role": "assistant", "content": "Como posso ajudar?"},
|
| 970 |
+
]
|
| 971 |
+
}
|
| 972 |
+
text = _extract_field(sample_chat, ["messages", "text"])
|
| 973 |
+
print(f" Chat text: {text!r}")
|
| 974 |
+
assert text and "Olá" in text, "Field extraction failed for chat"
|
| 975 |
+
|
| 976 |
+
sample_text = {"text": "Hello world"}
|
| 977 |
+
text = _extract_field(sample_text, ["text", "content"])
|
| 978 |
+
assert text == "Hello world"
|
| 979 |
+
print(f" Text field: {text!r}")
|
| 980 |
+
print(f" OK\n")
|
| 981 |
+
|
| 982 |
+
# Test 2: Test list of datasets
|
| 983 |
+
print("Test 2: Dataset list")
|
| 984 |
+
print(f" Configured datasets ({len(DEFAULT_DATASETS)}):")
|
| 985 |
+
for i, ds in enumerate(DEFAULT_DATASETS, 1):
|
| 986 |
+
print(f" {i}. {ds}")
|
| 987 |
+
assert len(DEFAULT_DATASETS) == 14
|
| 988 |
+
# V6.5-V2: verifica sequência CONHECIMENTO (8 datasets) + PUNICAO (1)
|
| 989 |
+
assert len(CONHECIMENTO_DATASETS_V2) == 8
|
| 990 |
+
assert PUNICAO_DATASET_V2 == "BrunoN-Dev/corpus-ptbr-v1"
|
| 991 |
+
print(f" V2 CONHECIMENTO datasets: {len(CONHECIMENTO_DATASETS_V2)}")
|
| 992 |
+
print(f" V2 PUNICAO dataset: {PUNICAO_DATASET_V2}")
|
| 993 |
+
print(f" OK\n")
|
| 994 |
+
|
| 995 |
+
print("=== ALL STREAMING TESTS PASSED ===")
|
train_v6_5_v2.py
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
v6_5_v2_attention_eval.json
CHANGED
|
@@ -3,18 +3,18 @@
|
|
| 3 |
"user_requirement": "verificar se o mecanismo de atenção está ativo e acessado logicamente funcional",
|
| 4 |
"metrics": {
|
| 5 |
"active": true,
|
| 6 |
-
"n_calls":
|
| 7 |
"n_errors": 0,
|
| 8 |
-
"last_norm_in":
|
| 9 |
-
"last_norm_out":
|
| 10 |
"last_attn_activated": true,
|
| 11 |
-
"last_attn_diff_norm":
|
| 12 |
"n_heads": 8,
|
| 13 |
"logic_functional": true
|
| 14 |
},
|
| 15 |
"active": true,
|
| 16 |
"logic_functional": true,
|
| 17 |
-
"n_calls":
|
| 18 |
"n_errors": 0,
|
| 19 |
"n_heads": 8,
|
| 20 |
"assessment": "PASS"
|
|
|
|
| 3 |
"user_requirement": "verificar se o mecanismo de atenção está ativo e acessado logicamente funcional",
|
| 4 |
"metrics": {
|
| 5 |
"active": true,
|
| 6 |
+
"n_calls": 10136,
|
| 7 |
"n_errors": 0,
|
| 8 |
+
"last_norm_in": 78.78106689453125,
|
| 9 |
+
"last_norm_out": 83.25476837158203,
|
| 10 |
"last_attn_activated": true,
|
| 11 |
+
"last_attn_diff_norm": 81.70692443847656,
|
| 12 |
"n_heads": 8,
|
| 13 |
"logic_functional": true
|
| 14 |
},
|
| 15 |
"active": true,
|
| 16 |
"logic_functional": true,
|
| 17 |
+
"n_calls": 10136,
|
| 18 |
"n_errors": 0,
|
| 19 |
"n_heads": 8,
|
| 20 |
"assessment": "PASS"
|
v6_5_v2_model_states.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:e6ec0c41f5e815e0451cd8627febd2dbe76395a1e0e3cb4a11cb58f98ddb29fe
|
| 3 |
+
size 56004222
|
v6_5_v2_model_states_after_conhecimento.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:1b943657505dfb8f42b12875272a1395932d836755aea0cf52ba1c4ce7a901d2
|
| 3 |
+
size 53528857
|
v6_5_v2_phases_eval.json
CHANGED
|
The diff for this file is too large to render.
See raw diff
|
|
|
v6_5_v2_report.json
CHANGED
|
@@ -1,6 +1,6 @@
|
|
| 1 |
{
|
| 2 |
"version": "V6.5-V2",
|
| 3 |
-
"timestamp": "2026-08-
|
| 4 |
"config": {
|
| 5 |
"som_grid": [
|
| 6 |
6,
|
|
@@ -9,9 +9,9 @@
|
|
| 9 |
4
|
| 10 |
],
|
| 11 |
"n_neurons": 864,
|
| 12 |
-
"hidden_dim":
|
| 13 |
-
"vocab_size":
|
| 14 |
-
"n_hypotheses":
|
| 15 |
"n_trials": 3,
|
| 16 |
"hyp_train_steps": 30,
|
| 17 |
"stream_batch_size": 100,
|
|
@@ -54,10 +54,10 @@
|
|
| 54 |
"init_done": true
|
| 55 |
},
|
| 56 |
"fp16_benchmark": {
|
| 57 |
-
"best_time_ms":
|
| 58 |
-
"avg_time_ms":
|
| 59 |
-
"best_tflops": 1.
|
| 60 |
-
"avg_tflops": 1.
|
| 61 |
"matrix_size": 4000.0
|
| 62 |
},
|
| 63 |
"v2_verification": {
|
|
@@ -86,13 +86,13 @@
|
|
| 86 |
"fase_1_conhecimento_summary": {
|
| 87 |
"total_samples": 8000,
|
| 88 |
"meta_atingida": true,
|
| 89 |
-
"elapsed_s":
|
| 90 |
"storage_critical_stopped": false
|
| 91 |
},
|
| 92 |
"fase_2_punicão_summary": {
|
| 93 |
"total_samples": 2000,
|
| 94 |
"meta_atingida": true,
|
| 95 |
-
"elapsed_s":
|
| 96 |
"punishment_events": 120,
|
| 97 |
"hypotheses_trainings": 60,
|
| 98 |
"delta_applications": 60,
|
|
@@ -101,7 +101,7 @@
|
|
| 101 |
"attention_eval_summary": {
|
| 102 |
"active": true,
|
| 103 |
"logic_functional": true,
|
| 104 |
-
"n_calls":
|
| 105 |
"assessment": "PASS"
|
| 106 |
},
|
| 107 |
"predict_fix_summary": {
|
|
@@ -117,36 +117,41 @@
|
|
| 117 |
"n_with_think": 3,
|
| 118 |
"answer_rate": 1.0,
|
| 119 |
"think_rate": 1.0,
|
| 120 |
-
"avg_latency_ms":
|
| 121 |
"avg_reasoning_length": 1026.0
|
| 122 |
},
|
| 123 |
"model_states_saved_to": "/home/z/my-project/BiGRU_T_version/v6_5_v2_model_states.pt",
|
| 124 |
"save_info": {
|
| 125 |
"saved": true,
|
| 126 |
"path": "/home/z/my-project/BiGRU_T_version/v6_5_v2_model_states.pt",
|
| 127 |
-
"size_mb":
|
|
|
|
|
|
|
|
|
|
| 128 |
"reason": "end_of_training_v2",
|
| 129 |
"step": 700,
|
| 130 |
"total_samples": 10000,
|
| 131 |
"n_tensors": 7,
|
| 132 |
-
"n_buffer_tail":
|
| 133 |
},
|
| 134 |
"final_v2_metrics": {
|
| 135 |
"version": "V2-dynamic",
|
| 136 |
-
"n_hypotheses":
|
| 137 |
-
"n_hypotheses_active":
|
| 138 |
-
"max_n_hypotheses":
|
| 139 |
"n_trials": 6,
|
| 140 |
"hyp_train_steps": 30,
|
| 141 |
"hyp_lr": 0.0001,
|
| 142 |
"delta_scale": 0.0010000000474974513,
|
| 143 |
-
"n_generators":
|
| 144 |
"punishment_count": 0,
|
| 145 |
"success_count": 0,
|
| 146 |
"training_ready": false,
|
| 147 |
"classifier_trained": true,
|
| 148 |
-
"ewc_reference_set":
|
| 149 |
"buffer_size": 256,
|
|
|
|
|
|
|
| 150 |
"dynamic_adaptation": {
|
| 151 |
"loss_history_len": 8,
|
| 152 |
"loss_stats": {
|
|
@@ -161,7 +166,7 @@
|
|
| 161 |
"n_adaptations": 20,
|
| 162 |
"limits": {
|
| 163 |
"min_n_hypotheses": 4,
|
| 164 |
-
"max_n_hypotheses":
|
| 165 |
"min_n_trials": 1,
|
| 166 |
"max_n_trials": 6,
|
| 167 |
"min_hyp_train_steps": 10,
|
|
@@ -172,12 +177,12 @@
|
|
| 172 |
"trigger": "punishment",
|
| 173 |
"step": 115,
|
| 174 |
"before": {
|
| 175 |
-
"n_hypotheses":
|
| 176 |
"n_trials": 6,
|
| 177 |
"hyp_train_steps": 30
|
| 178 |
},
|
| 179 |
"after": {
|
| 180 |
-
"n_hypotheses":
|
| 181 |
"n_trials": 6,
|
| 182 |
"hyp_train_steps": 30
|
| 183 |
},
|
|
@@ -196,12 +201,12 @@
|
|
| 196 |
"trigger": "auto",
|
| 197 |
"step": 116,
|
| 198 |
"before": {
|
| 199 |
-
"n_hypotheses":
|
| 200 |
"n_trials": 6,
|
| 201 |
"hyp_train_steps": 30
|
| 202 |
},
|
| 203 |
"after": {
|
| 204 |
-
"n_hypotheses":
|
| 205 |
"n_trials": 6,
|
| 206 |
"hyp_train_steps": 30
|
| 207 |
},
|
|
@@ -220,12 +225,12 @@
|
|
| 220 |
"trigger": "punishment",
|
| 221 |
"step": 117,
|
| 222 |
"before": {
|
| 223 |
-
"n_hypotheses":
|
| 224 |
"n_trials": 6,
|
| 225 |
"hyp_train_steps": 30
|
| 226 |
},
|
| 227 |
"after": {
|
| 228 |
-
"n_hypotheses":
|
| 229 |
"n_trials": 6,
|
| 230 |
"hyp_train_steps": 30
|
| 231 |
},
|
|
@@ -244,12 +249,12 @@
|
|
| 244 |
"trigger": "auto",
|
| 245 |
"step": 118,
|
| 246 |
"before": {
|
| 247 |
-
"n_hypotheses":
|
| 248 |
"n_trials": 6,
|
| 249 |
"hyp_train_steps": 30
|
| 250 |
},
|
| 251 |
"after": {
|
| 252 |
-
"n_hypotheses":
|
| 253 |
"n_trials": 6,
|
| 254 |
"hyp_train_steps": 30
|
| 255 |
},
|
|
@@ -268,12 +273,12 @@
|
|
| 268 |
"trigger": "punishment",
|
| 269 |
"step": 119,
|
| 270 |
"before": {
|
| 271 |
-
"n_hypotheses":
|
| 272 |
"n_trials": 6,
|
| 273 |
"hyp_train_steps": 30
|
| 274 |
},
|
| 275 |
"after": {
|
| 276 |
-
"n_hypotheses":
|
| 277 |
"n_trials": 6,
|
| 278 |
"hyp_train_steps": 30
|
| 279 |
},
|
|
|
|
| 1 |
{
|
| 2 |
"version": "V6.5-V2",
|
| 3 |
+
"timestamp": "2026-08-09T01:05:21.476394",
|
| 4 |
"config": {
|
| 5 |
"som_grid": [
|
| 6 |
6,
|
|
|
|
| 9 |
4
|
| 10 |
],
|
| 11 |
"n_neurons": 864,
|
| 12 |
+
"hidden_dim": 512,
|
| 13 |
+
"vocab_size": 8192,
|
| 14 |
+
"n_hypotheses": 8,
|
| 15 |
"n_trials": 3,
|
| 16 |
"hyp_train_steps": 30,
|
| 17 |
"stream_batch_size": 100,
|
|
|
|
| 54 |
"init_done": true
|
| 55 |
},
|
| 56 |
"fp16_benchmark": {
|
| 57 |
+
"best_time_ms": 85.25625100082834,
|
| 58 |
+
"avg_time_ms": 97.09998300058942,
|
| 59 |
+
"best_tflops": 1.5013561879322652,
|
| 60 |
+
"avg_tflops": 1.3182288610619326,
|
| 61 |
"matrix_size": 4000.0
|
| 62 |
},
|
| 63 |
"v2_verification": {
|
|
|
|
| 86 |
"fase_1_conhecimento_summary": {
|
| 87 |
"total_samples": 8000,
|
| 88 |
"meta_atingida": true,
|
| 89 |
+
"elapsed_s": 337.3461060523987,
|
| 90 |
"storage_critical_stopped": false
|
| 91 |
},
|
| 92 |
"fase_2_punicão_summary": {
|
| 93 |
"total_samples": 2000,
|
| 94 |
"meta_atingida": true,
|
| 95 |
+
"elapsed_s": 221.32905673980713,
|
| 96 |
"punishment_events": 120,
|
| 97 |
"hypotheses_trainings": 60,
|
| 98 |
"delta_applications": 60,
|
|
|
|
| 101 |
"attention_eval_summary": {
|
| 102 |
"active": true,
|
| 103 |
"logic_functional": true,
|
| 104 |
+
"n_calls": 10136,
|
| 105 |
"assessment": "PASS"
|
| 106 |
},
|
| 107 |
"predict_fix_summary": {
|
|
|
|
| 117 |
"n_with_think": 3,
|
| 118 |
"answer_rate": 1.0,
|
| 119 |
"think_rate": 1.0,
|
| 120 |
+
"avg_latency_ms": 6.689548492431641,
|
| 121 |
"avg_reasoning_length": 1026.0
|
| 122 |
},
|
| 123 |
"model_states_saved_to": "/home/z/my-project/BiGRU_T_version/v6_5_v2_model_states.pt",
|
| 124 |
"save_info": {
|
| 125 |
"saved": true,
|
| 126 |
"path": "/home/z/my-project/BiGRU_T_version/v6_5_v2_model_states.pt",
|
| 127 |
+
"size_mb": 56.004222,
|
| 128 |
+
"size_gb": 0.05215799622237682,
|
| 129 |
+
"size_status": "OK",
|
| 130 |
+
"size_within_1gb_limit": true,
|
| 131 |
"reason": "end_of_training_v2",
|
| 132 |
"step": 700,
|
| 133 |
"total_samples": 10000,
|
| 134 |
"n_tensors": 7,
|
| 135 |
+
"n_buffer_tail": 64
|
| 136 |
},
|
| 137 |
"final_v2_metrics": {
|
| 138 |
"version": "V2-dynamic",
|
| 139 |
+
"n_hypotheses": 8,
|
| 140 |
+
"n_hypotheses_active": 8,
|
| 141 |
+
"max_n_hypotheses": 16,
|
| 142 |
"n_trials": 6,
|
| 143 |
"hyp_train_steps": 30,
|
| 144 |
"hyp_lr": 0.0001,
|
| 145 |
"delta_scale": 0.0010000000474974513,
|
| 146 |
+
"n_generators": 16,
|
| 147 |
"punishment_count": 0,
|
| 148 |
"success_count": 0,
|
| 149 |
"training_ready": false,
|
| 150 |
"classifier_trained": true,
|
| 151 |
+
"ewc_reference_set": false,
|
| 152 |
"buffer_size": 256,
|
| 153 |
+
"total_hyp_steps_executed": 1800,
|
| 154 |
+
"n_train_hyp_calls": 60,
|
| 155 |
"dynamic_adaptation": {
|
| 156 |
"loss_history_len": 8,
|
| 157 |
"loss_stats": {
|
|
|
|
| 166 |
"n_adaptations": 20,
|
| 167 |
"limits": {
|
| 168 |
"min_n_hypotheses": 4,
|
| 169 |
+
"max_n_hypotheses": 16,
|
| 170 |
"min_n_trials": 1,
|
| 171 |
"max_n_trials": 6,
|
| 172 |
"min_hyp_train_steps": 10,
|
|
|
|
| 177 |
"trigger": "punishment",
|
| 178 |
"step": 115,
|
| 179 |
"before": {
|
| 180 |
+
"n_hypotheses": 8,
|
| 181 |
"n_trials": 6,
|
| 182 |
"hyp_train_steps": 30
|
| 183 |
},
|
| 184 |
"after": {
|
| 185 |
+
"n_hypotheses": 8,
|
| 186 |
"n_trials": 6,
|
| 187 |
"hyp_train_steps": 30
|
| 188 |
},
|
|
|
|
| 201 |
"trigger": "auto",
|
| 202 |
"step": 116,
|
| 203 |
"before": {
|
| 204 |
+
"n_hypotheses": 8,
|
| 205 |
"n_trials": 6,
|
| 206 |
"hyp_train_steps": 30
|
| 207 |
},
|
| 208 |
"after": {
|
| 209 |
+
"n_hypotheses": 8,
|
| 210 |
"n_trials": 6,
|
| 211 |
"hyp_train_steps": 30
|
| 212 |
},
|
|
|
|
| 225 |
"trigger": "punishment",
|
| 226 |
"step": 117,
|
| 227 |
"before": {
|
| 228 |
+
"n_hypotheses": 8,
|
| 229 |
"n_trials": 6,
|
| 230 |
"hyp_train_steps": 30
|
| 231 |
},
|
| 232 |
"after": {
|
| 233 |
+
"n_hypotheses": 8,
|
| 234 |
"n_trials": 6,
|
| 235 |
"hyp_train_steps": 30
|
| 236 |
},
|
|
|
|
| 249 |
"trigger": "auto",
|
| 250 |
"step": 118,
|
| 251 |
"before": {
|
| 252 |
+
"n_hypotheses": 8,
|
| 253 |
"n_trials": 6,
|
| 254 |
"hyp_train_steps": 30
|
| 255 |
},
|
| 256 |
"after": {
|
| 257 |
+
"n_hypotheses": 8,
|
| 258 |
"n_trials": 6,
|
| 259 |
"hyp_train_steps": 30
|
| 260 |
},
|
|
|
|
| 273 |
"trigger": "punishment",
|
| 274 |
"step": 119,
|
| 275 |
"before": {
|
| 276 |
+
"n_hypotheses": 8,
|
| 277 |
"n_trials": 6,
|
| 278 |
"hyp_train_steps": 30
|
| 279 |
},
|
| 280 |
"after": {
|
| 281 |
+
"n_hypotheses": 8,
|
| 282 |
"n_trials": 6,
|
| 283 |
"hyp_train_steps": 30
|
| 284 |
},
|
v6_5_v2_user_questions.json
CHANGED
|
@@ -25,10 +25,10 @@
|
|
| 25 |
"has_answer": true,
|
| 26 |
"has_decompose": true,
|
| 27 |
"think_preview": "Analisando a query: 'Luva de Pedreiro Távila'\nIdentificando o tipo de problema e requisitos.\nDeterminando se ferramentas são necessárias.",
|
| 28 |
-
"answer_preview": "prediction=short_text | BMU=(
|
| 29 |
"raw_response_preview": "<think>\nAnalisando a query: 'Luva de Pedreiro Távila'\nIdentificando o tipo de problema e requisitos.\nDeterminando se ferramentas são necessárias.\n</think>\n<plan>\nPlano de resolução:\n1. Decompor o problema em sub-tarefas\n2. Identificar ferramentas necessárias (disponíveis: som_query, buffer_stats)\n3. Executar sub-tarefas em sequência\n4. Monitorar resultados\n5. Compor resposta final\n</plan>\n<decompose>\n- Processar: Luva de Pedreiro Távila\n</decompose>\n<execute>\nSub-tarefa 'Processar: Luva de Pedre...",
|
| 30 |
"n_tags": 4,
|
| 31 |
-
"latency_ms":
|
| 32 |
},
|
| 33 |
{
|
| 34 |
"query": "Lula reserva valor",
|
|
@@ -43,10 +43,10 @@
|
|
| 43 |
"has_answer": true,
|
| 44 |
"has_decompose": true,
|
| 45 |
"think_preview": "Analisando a query: 'Lula reserva valor'\nIdentificando o tipo de problema e requisitos.\nDeterminando se ferramentas são necessárias.",
|
| 46 |
-
"answer_preview": "prediction=short_text | BMU=(
|
| 47 |
"raw_response_preview": "<think>\nAnalisando a query: 'Lula reserva valor'\nIdentificando o tipo de problema e requisitos.\nDeterminando se ferramentas são necessárias.\n</think>\n<plan>\nPlano de resolução:\n1. Decompor o problema em sub-tarefas\n2. Identificar ferramentas necessárias (disponíveis: som_query, buffer_stats)\n3. Executar sub-tarefas em sequência\n4. Monitorar resultados\n5. Compor resposta final\n</plan>\n<decompose>\n- Processar: Lula reserva valor\n</decompose>\n<execute>\nSub-tarefa 'Processar: Lula reserva valor' exe...",
|
| 48 |
"n_tags": 4,
|
| 49 |
-
"latency_ms":
|
| 50 |
},
|
| 51 |
{
|
| 52 |
"query": "Amazonas força-tarefa vítimas",
|
|
@@ -61,10 +61,10 @@
|
|
| 61 |
"has_answer": true,
|
| 62 |
"has_decompose": true,
|
| 63 |
"think_preview": "Analisando a query: 'Amazonas força-tarefa vítimas'\nIdentificando o tipo de problema e requisitos.\nDeterminando se ferramentas são necessárias.",
|
| 64 |
-
"answer_preview": "prediction=short_text | BMU=(
|
| 65 |
"raw_response_preview": "<think>\nAnalisando a query: 'Amazonas força-tarefa vítimas'\nIdentificando o tipo de problema e requisitos.\nDeterminando se ferramentas são necessárias.\n</think>\n<plan>\nPlano de resolução:\n1. Decompor o problema em sub-tarefas\n2. Identificar ferramentas necessárias (disponíveis: som_query, buffer_stats)\n3. Executar sub-tarefas em sequência\n4. Monitorar resultados\n5. Compor resposta final\n</plan>\n<decompose>\n- Processar: Amazonas força-tarefa vítimas\n</decompose>\n<execute>\nSub-tarefa 'Processar: A...",
|
| 66 |
"n_tags": 4,
|
| 67 |
-
"latency_ms":
|
| 68 |
}
|
| 69 |
],
|
| 70 |
"summary": {
|
|
@@ -72,7 +72,7 @@
|
|
| 72 |
"n_with_think": 3,
|
| 73 |
"answer_rate": 1.0,
|
| 74 |
"think_rate": 1.0,
|
| 75 |
-
"avg_latency_ms":
|
| 76 |
"avg_reasoning_length": 1026.0
|
| 77 |
},
|
| 78 |
"quality_assessment": {
|
|
|
|
| 25 |
"has_answer": true,
|
| 26 |
"has_decompose": true,
|
| 27 |
"think_preview": "Analisando a query: 'Luva de Pedreiro Távila'\nIdentificando o tipo de problema e requisitos.\nDeterminando se ferramentas são necessárias.",
|
| 28 |
+
"answer_preview": "prediction=short_text | BMU=(2, 5, 5, 1)",
|
| 29 |
"raw_response_preview": "<think>\nAnalisando a query: 'Luva de Pedreiro Távila'\nIdentificando o tipo de problema e requisitos.\nDeterminando se ferramentas são necessárias.\n</think>\n<plan>\nPlano de resolução:\n1. Decompor o problema em sub-tarefas\n2. Identificar ferramentas necessárias (disponíveis: som_query, buffer_stats)\n3. Executar sub-tarefas em sequência\n4. Monitorar resultados\n5. Compor resposta final\n</plan>\n<decompose>\n- Processar: Luva de Pedreiro Távila\n</decompose>\n<execute>\nSub-tarefa 'Processar: Luva de Pedre...",
|
| 30 |
"n_tags": 4,
|
| 31 |
+
"latency_ms": 7.518291473388672
|
| 32 |
},
|
| 33 |
{
|
| 34 |
"query": "Lula reserva valor",
|
|
|
|
| 43 |
"has_answer": true,
|
| 44 |
"has_decompose": true,
|
| 45 |
"think_preview": "Analisando a query: 'Lula reserva valor'\nIdentificando o tipo de problema e requisitos.\nDeterminando se ferramentas são necessárias.",
|
| 46 |
+
"answer_preview": "prediction=short_text | BMU=(2, 5, 5, 1)",
|
| 47 |
"raw_response_preview": "<think>\nAnalisando a query: 'Lula reserva valor'\nIdentificando o tipo de problema e requisitos.\nDeterminando se ferramentas são necessárias.\n</think>\n<plan>\nPlano de resolução:\n1. Decompor o problema em sub-tarefas\n2. Identificar ferramentas necessárias (disponíveis: som_query, buffer_stats)\n3. Executar sub-tarefas em sequência\n4. Monitorar resultados\n5. Compor resposta final\n</plan>\n<decompose>\n- Processar: Lula reserva valor\n</decompose>\n<execute>\nSub-tarefa 'Processar: Lula reserva valor' exe...",
|
| 48 |
"n_tags": 4,
|
| 49 |
+
"latency_ms": 6.790876388549805
|
| 50 |
},
|
| 51 |
{
|
| 52 |
"query": "Amazonas força-tarefa vítimas",
|
|
|
|
| 61 |
"has_answer": true,
|
| 62 |
"has_decompose": true,
|
| 63 |
"think_preview": "Analisando a query: 'Amazonas força-tarefa vítimas'\nIdentificando o tipo de problema e requisitos.\nDeterminando se ferramentas são necessárias.",
|
| 64 |
+
"answer_preview": "prediction=short_text | BMU=(2, 5, 5, 1)",
|
| 65 |
"raw_response_preview": "<think>\nAnalisando a query: 'Amazonas força-tarefa vítimas'\nIdentificando o tipo de problema e requisitos.\nDeterminando se ferramentas são necessárias.\n</think>\n<plan>\nPlano de resolução:\n1. Decompor o problema em sub-tarefas\n2. Identificar ferramentas necessárias (disponíveis: som_query, buffer_stats)\n3. Executar sub-tarefas em sequência\n4. Monitorar resultados\n5. Compor resposta final\n</plan>\n<decompose>\n- Processar: Amazonas força-tarefa vítimas\n</decompose>\n<execute>\nSub-tarefa 'Processar: A...",
|
| 66 |
"n_tags": 4,
|
| 67 |
+
"latency_ms": 5.759477615356445
|
| 68 |
}
|
| 69 |
],
|
| 70 |
"summary": {
|
|
|
|
| 72 |
"n_with_think": 3,
|
| 73 |
"answer_rate": 1.0,
|
| 74 |
"think_rate": 1.0,
|
| 75 |
+
"avg_latency_ms": 6.689548492431641,
|
| 76 |
"avg_reasoning_length": 1026.0
|
| 77 |
},
|
| 78 |
"quality_assessment": {
|
vqvae2_hierarchical.py
ADDED
|
@@ -0,0 +1,200 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
Xavante - vqvae2_hierarchical.py
|
| 3 |
+
Responsabilidade: VQ-VAE-2 hierarquica (Teorema 13.1/13.2/13.3).
|
| 4 |
+
Reutiliza implementacao validada do FlexNet com EMA + Goose VQ + dead code restart.
|
| 5 |
+
Referencia: Razavi et al. 2019 - Generating Diverse High-Fidelity Images with VQ-VAE-2.
|
| 6 |
+
"""
|
| 7 |
+
from __future__ import annotations
|
| 8 |
+
|
| 9 |
+
import logging
|
| 10 |
+
import math
|
| 11 |
+
from typing import List, Tuple
|
| 12 |
+
|
| 13 |
+
import torch
|
| 14 |
+
import torch.nn as nn
|
| 15 |
+
import torch.nn.functional as F
|
| 16 |
+
|
| 17 |
+
logger = logging.getLogger(__name__)
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
class VectorQuantizerEMA(nn.Module):
|
| 21 |
+
"""
|
| 22 |
+
VQ com atualizacao EMA do codebook (validado no FlexNet).
|
| 23 |
+
Inclui: dead code restart, commitment loss, Goose VQ.
|
| 24 |
+
v6 aprimoramentos:
|
| 25 |
+
- Diversity loss (entropy bonus) para forcar uso uniforme do codebook
|
| 26 |
+
- Commitment warmup: beta_t = beta_0 * min(1, t/warmup)
|
| 27 |
+
- EMA restart mais agressivo (decay adaptativo)
|
| 28 |
+
"""
|
| 29 |
+
|
| 30 |
+
def __init__(
|
| 31 |
+
self,
|
| 32 |
+
num_embeddings: int,
|
| 33 |
+
embedding_dim: int,
|
| 34 |
+
commitment_cost: float = 0.25,
|
| 35 |
+
decay: float = 0.99,
|
| 36 |
+
epsilon: float = 1e-5,
|
| 37 |
+
dead_threshold: int = 1,
|
| 38 |
+
diversity_weight: float = 0.1,
|
| 39 |
+
warmup_steps: int = 500,
|
| 40 |
+
):
|
| 41 |
+
super().__init__()
|
| 42 |
+
self.num_embeddings = num_embeddings
|
| 43 |
+
self.embedding_dim = embedding_dim
|
| 44 |
+
self.commitment_cost = commitment_cost
|
| 45 |
+
self.decay = decay
|
| 46 |
+
self.epsilon = epsilon
|
| 47 |
+
self.dead_threshold = dead_threshold
|
| 48 |
+
self.diversity_weight = diversity_weight
|
| 49 |
+
self.warmup_steps = warmup_steps
|
| 50 |
+
self._step = 0
|
| 51 |
+
|
| 52 |
+
embed = torch.empty(num_embeddings, embedding_dim)
|
| 53 |
+
nn.init.normal_(embed, mean=0.0, std=0.02)
|
| 54 |
+
self.register_buffer("embeddings", embed)
|
| 55 |
+
self.register_buffer("ema_count", torch.zeros(num_embeddings))
|
| 56 |
+
self.register_buffer("ema_weight", embed.clone())
|
| 57 |
+
self.register_buffer("usage_count", torch.zeros(num_embeddings, dtype=torch.long))
|
| 58 |
+
|
| 59 |
+
def forward(self, x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
| 60 |
+
"""
|
| 61 |
+
x: [B, D, H, W] ou [B, D]
|
| 62 |
+
Retorna (quantized, loss, encoding_indices)
|
| 63 |
+
"""
|
| 64 |
+
if x.dim() == 4:
|
| 65 |
+
# [B, D, H, W] -> [B, H, W, D] -> [B*H*W, D]
|
| 66 |
+
B, D, H, W = x.shape
|
| 67 |
+
flat = x.permute(0, 2, 3, 1).reshape(-1, self.embedding_dim)
|
| 68 |
+
else:
|
| 69 |
+
# [B, D] -> [B, D]
|
| 70 |
+
flat = x.reshape(-1, self.embedding_dim)
|
| 71 |
+
# Distancia
|
| 72 |
+
distances = (
|
| 73 |
+
flat.pow(2).sum(1, keepdim=True)
|
| 74 |
+
- 2 * flat @ self.embeddings.t()
|
| 75 |
+
+ self.embeddings.pow(2).sum(1)
|
| 76 |
+
)
|
| 77 |
+
indices = torch.argmin(distances, dim=1)
|
| 78 |
+
one_hot = F.one_hot(indices, self.num_embeddings).type(flat.dtype)
|
| 79 |
+
quantized = one_hot @ self.embeddings
|
| 80 |
+
# Uso
|
| 81 |
+
self.usage_count.index_add_(0, indices, torch.ones_like(indices))
|
| 82 |
+
|
| 83 |
+
# EMA update
|
| 84 |
+
if self.training:
|
| 85 |
+
with torch.no_grad():
|
| 86 |
+
count = one_hot.sum(0)
|
| 87 |
+
ema_count_new = self.decay * self.ema_count + (1 - self.decay) * count
|
| 88 |
+
total = ema_count_new.sum()
|
| 89 |
+
ema_count_new = (ema_count_new + self.epsilon) / (total + self.num_embeddings * self.epsilon) * total
|
| 90 |
+
dw = one_hot.t() @ flat
|
| 91 |
+
self.ema_weight = self.decay * self.ema_weight + (1 - self.decay) * dw
|
| 92 |
+
self.ema_count = ema_count_new
|
| 93 |
+
self.embeddings.copy_(self.ema_weight / self.ema_count.unsqueeze(1))
|
| 94 |
+
|
| 95 |
+
# Dead code restart
|
| 96 |
+
dead = self.usage_count < self.dead_threshold
|
| 97 |
+
if dead.any() and flat.shape[0] > 0:
|
| 98 |
+
dead_indices = dead.nonzero(as_tuple=True)[0]
|
| 99 |
+
n_dead = dead_indices.shape[0]
|
| 100 |
+
# Sample min(n_dead, flat.shape[0]) candidates from flat
|
| 101 |
+
n_replace = min(n_dead, flat.shape[0])
|
| 102 |
+
if n_replace > 0:
|
| 103 |
+
idx = torch.randperm(flat.shape[0])[:n_replace].to(flat.device)
|
| 104 |
+
# Substituir apenas n_replace codigos mortos
|
| 105 |
+
self.embeddings.data[dead_indices[:n_replace]] = flat[idx].detach()
|
| 106 |
+
|
| 107 |
+
# Commitment loss com warmup
|
| 108 |
+
self._step += 1
|
| 109 |
+
beta = self.commitment_cost * min(1.0, self._step / max(self.warmup_steps, 1))
|
| 110 |
+
e_latent_loss = F.mse_loss(quantized.detach(), flat)
|
| 111 |
+
q_latent_loss = F.mse_loss(quantized, flat.detach())
|
| 112 |
+
# Diversity loss: -H(p) onde p = freq de uso de cada code
|
| 113 |
+
# Maximizar entropia = forçar uso uniforme do codebook
|
| 114 |
+
if self.diversity_weight > 0 and self.training:
|
| 115 |
+
code_probs = one_hot.float().mean(dim=0) # [num_embeddings]
|
| 116 |
+
code_probs = code_probs.clamp_min(1e-10)
|
| 117 |
+
entropy = -(code_probs * code_probs.log()).sum()
|
| 118 |
+
diversity_loss = -entropy / math.log(self.num_embeddings) # normalize to [0, 1]
|
| 119 |
+
else:
|
| 120 |
+
diversity_loss = torch.tensor(0.0, device=x.device)
|
| 121 |
+
loss = q_latent_loss + beta * e_latent_loss + self.diversity_weight * diversity_loss
|
| 122 |
+
|
| 123 |
+
# Straight-through
|
| 124 |
+
quantized = flat + (quantized - flat).detach()
|
| 125 |
+
# Restore shape
|
| 126 |
+
if x.dim() == 4:
|
| 127 |
+
quantized = quantized.view(B, H, W, D).permute(0, 3, 1, 2).contiguous()
|
| 128 |
+
indices = indices.view(B, H, W)
|
| 129 |
+
else:
|
| 130 |
+
quantized = quantized.view_as(x)
|
| 131 |
+
return quantized, loss, indices
|
| 132 |
+
|
| 133 |
+
|
| 134 |
+
class HierarchicalVQVAE2(nn.Module):
|
| 135 |
+
"""
|
| 136 |
+
VQ-VAE-2 com 2 niveis: top (global) e bottom (local).
|
| 137 |
+
Mapeia: Teorema 13.1/13.2/13.3.
|
| 138 |
+
"""
|
| 139 |
+
|
| 140 |
+
def __init__(
|
| 141 |
+
self,
|
| 142 |
+
in_channels: int = 3,
|
| 143 |
+
hidden_dim: int = 64,
|
| 144 |
+
top_codebook: int = 128,
|
| 145 |
+
bottom_codebook: int = 512,
|
| 146 |
+
commitment_cost: float = 0.25,
|
| 147 |
+
decay: float = 0.99,
|
| 148 |
+
):
|
| 149 |
+
super().__init__()
|
| 150 |
+
self.encoder_bottom = nn.Sequential(
|
| 151 |
+
nn.Conv2d(in_channels, hidden_dim, 4, 2, 1),
|
| 152 |
+
nn.GELU(),
|
| 153 |
+
nn.Conv2d(hidden_dim, hidden_dim, 4, 2, 1),
|
| 154 |
+
nn.GELU(),
|
| 155 |
+
nn.Conv2d(hidden_dim, hidden_dim, 3, 1, 1),
|
| 156 |
+
)
|
| 157 |
+
self.encoder_top = nn.Sequential(
|
| 158 |
+
nn.Conv2d(hidden_dim, hidden_dim, 4, 2, 1),
|
| 159 |
+
nn.GELU(),
|
| 160 |
+
nn.Conv2d(hidden_dim, hidden_dim, 4, 2, 1),
|
| 161 |
+
nn.GELU(),
|
| 162 |
+
nn.Conv2d(hidden_dim, hidden_dim, 3, 1, 1),
|
| 163 |
+
)
|
| 164 |
+
self.vq_top = VectorQuantizerEMA(top_codebook, hidden_dim, commitment_cost, decay)
|
| 165 |
+
self.vq_bottom = VectorQuantizerEMA(bottom_codebook, hidden_dim, commitment_cost, decay)
|
| 166 |
+
self.decoder_top = nn.Sequential(
|
| 167 |
+
nn.ConvTranspose2d(hidden_dim, hidden_dim, 4, 2, 1),
|
| 168 |
+
nn.GELU(),
|
| 169 |
+
nn.ConvTranspose2d(hidden_dim, hidden_dim, 4, 2, 1),
|
| 170 |
+
nn.GELU(),
|
| 171 |
+
)
|
| 172 |
+
self.decoder_bottom = nn.Sequential(
|
| 173 |
+
nn.ConvTranspose2d(hidden_dim * 2, hidden_dim, 4, 2, 1),
|
| 174 |
+
nn.GELU(),
|
| 175 |
+
nn.ConvTranspose2d(hidden_dim, hidden_dim, 4, 2, 1),
|
| 176 |
+
nn.GELU(),
|
| 177 |
+
nn.Conv2d(hidden_dim, in_channels, 3, 1, 1),
|
| 178 |
+
)
|
| 179 |
+
|
| 180 |
+
def forward(self, x: torch.Tensor) -> dict:
|
| 181 |
+
# Bottom encode
|
| 182 |
+
ze_b = self.encoder_bottom(x)
|
| 183 |
+
# Top encode
|
| 184 |
+
ze_t = self.encoder_top(ze_b)
|
| 185 |
+
zq_t, loss_t, idx_t = self.vq_top(ze_t)
|
| 186 |
+
# Decodifica top e soma com bottom
|
| 187 |
+
dec_t = self.decoder_top(zq_t)
|
| 188 |
+
# Concat com bottom quantized
|
| 189 |
+
zq_b, loss_b, idx_b = self.vq_bottom(ze_b)
|
| 190 |
+
cat = torch.cat([dec_t, zq_b], dim=1)
|
| 191 |
+
x_recon = self.decoder_bottom(cat)
|
| 192 |
+
return {
|
| 193 |
+
"reconstruction": x_recon,
|
| 194 |
+
"loss": loss_t + loss_b,
|
| 195 |
+
"top_indices": idx_t,
|
| 196 |
+
"bottom_indices": idx_b,
|
| 197 |
+
}
|
| 198 |
+
|
| 199 |
+
|
| 200 |
+
__all__ = ["VectorQuantizerEMA", "HierarchicalVQVAE2"]
|
vqvae2_hierarchical_flexnet.py
ADDED
|
@@ -0,0 +1,513 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
vqvae2_hierarchical — VQ-VAE-2 (Razavi 2019) hierárquico com 2 níveis de codebook.
|
| 3 |
+
|
| 4 |
+
Arquitetura:
|
| 5 |
+
- Top level: codebook K_top captura estrutura global (low-resolution)
|
| 6 |
+
- Bottom level: codebook K_bot captura detalhes locais (high-resolution)
|
| 7 |
+
- Cada nível tem encoder/decoder proprio
|
| 8 |
+
- Residual quantization: bottom level quantiza o RESÍDUO do top level
|
| 9 |
+
|
| 10 |
+
Loss total:
|
| 11 |
+
L = ||x - D(top_q, bot_q)||²
|
| 12 |
+
+ ||sg[z_top] - top_q||² + β·||z_top - sg[top_q]||²
|
| 13 |
+
+ ||sg[z_bot] - bot_q||² + β·||z_bot - sg[bot_q]||²
|
| 14 |
+
|
| 15 |
+
Implementa:
|
| 16 |
+
- EMA updates em ambos codebooks
|
| 17 |
+
- Dead code restart em ambos codebooks
|
| 18 |
+
- Goose VQ com schedule de temperatura (alto→baixo)
|
| 19 |
+
|
| 20 |
+
Setup dos testes (solicitado):
|
| 21 |
+
- batch=256, K_top=512, K_bot=512
|
| 22 |
+
- Goose τ schedule: começa 2.0, decai para 0.5 via cosine
|
| 23 |
+
- EMA decay γ=0.99
|
| 24 |
+
- Dead code restart a cada 3 épocas
|
| 25 |
+
"""
|
| 26 |
+
|
| 27 |
+
from __future__ import annotations
|
| 28 |
+
from typing import Dict, Optional, Tuple
|
| 29 |
+
|
| 30 |
+
import math
|
| 31 |
+
import numpy as np
|
| 32 |
+
import torch
|
| 33 |
+
import torch.nn as nn
|
| 34 |
+
import torch.nn.functional as F
|
| 35 |
+
|
| 36 |
+
|
| 37 |
+
# ============================================================
|
| 38 |
+
# v12.8: Normalizações alternativas (substituem LayerNorm v12.7)
|
| 39 |
+
# ============================================================
|
| 40 |
+
|
| 41 |
+
class RMSNorm(nn.Module):
|
| 42 |
+
"""RMSNorm — preserva magnitude (não subtrai média como LayerNorm).
|
| 43 |
+
|
| 44 |
+
y = x / sqrt(mean(x²) + eps) * gamma
|
| 45 |
+
|
| 46 |
+
Vantagens sobre LayerNorm:
|
| 47 |
+
- Mais barato (sem cálculo de variância)
|
| 48 |
+
- Preserva magnitude do sinal (importante para VQ-VAE)
|
| 49 |
+
- Zhang & Sennrich 2019
|
| 50 |
+
"""
|
| 51 |
+
|
| 52 |
+
def __init__(self, dim: int, eps: float = 1e-6):
|
| 53 |
+
super().__init__()
|
| 54 |
+
self.gamma = nn.Parameter(torch.ones(dim))
|
| 55 |
+
self.eps = eps
|
| 56 |
+
|
| 57 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 58 |
+
rms = torch.sqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)
|
| 59 |
+
return self.gamma * x / rms
|
| 60 |
+
|
| 61 |
+
|
| 62 |
+
def apply_weightnorm(module: nn.Module) -> nn.Module:
|
| 63 |
+
"""Aplica weight normalization (EnCodec-style) a um módulo Linear.
|
| 64 |
+
|
| 65 |
+
Salakhutdinov & Wang 2016: w = g * v / ||v||, onde g é escalar aprendido.
|
| 66 |
+
|
| 67 |
+
Vantagens:
|
| 68 |
+
- Normaliza apenas o peso, não a ativação
|
| 69 |
+
- Preserva estatística do input (importante para VQ-VAE)
|
| 70 |
+
- Usado em EnCodec (Défossez 2022) e SoundStream (Zeghidour 2021)
|
| 71 |
+
"""
|
| 72 |
+
return nn.utils.weight_norm(module)
|
| 73 |
+
|
| 74 |
+
|
| 75 |
+
class HierarchicalVectorQuantizer(nn.Module):
|
| 76 |
+
"""Quantizador hierárquico com 2 níveis (top + bottom) e técnicas anti-collapse.
|
| 77 |
+
|
| 78 |
+
Args:
|
| 79 |
+
num_codes_top: K_top (codebook de topo, captura estrutura global)
|
| 80 |
+
num_codes_bot: K_bot (codebook de base, captura detalhes)
|
| 81 |
+
code_dim: D (dimensão compartilhada)
|
| 82 |
+
beta: peso do commitment loss
|
| 83 |
+
ema_decay: γ para EMA (default 0.99)
|
| 84 |
+
dead_code_threshold: N_k abaixo do qual código é morto
|
| 85 |
+
dead_code_restart_every: a cada quantas épocas verificar
|
| 86 |
+
goose_temp_init: temperatura Goose inicial (default 2.0)
|
| 87 |
+
goose_temp_final: temperatura Goose final (default 0.5)
|
| 88 |
+
goose_schedule: 'cosine' ou 'linear'
|
| 89 |
+
total_epochs: total de épocas (para schedule)
|
| 90 |
+
"""
|
| 91 |
+
|
| 92 |
+
def __init__(
|
| 93 |
+
self,
|
| 94 |
+
num_codes_top: int = 512,
|
| 95 |
+
num_codes_bot: int = 512,
|
| 96 |
+
code_dim: int = 32,
|
| 97 |
+
beta: float = 0.25,
|
| 98 |
+
ema_decay: float = 0.99,
|
| 99 |
+
dead_code_threshold: float = 1.0,
|
| 100 |
+
dead_code_restart_every: int = 3,
|
| 101 |
+
goose_temp_init: float = 2.0,
|
| 102 |
+
goose_temp_final: float = 0.5,
|
| 103 |
+
goose_schedule: str = "cosine",
|
| 104 |
+
total_epochs: int = 25,
|
| 105 |
+
norm_type: str = "none", # v12.8: "none" | "weightnorm" | "rmsnorm"
|
| 106 |
+
):
|
| 107 |
+
super().__init__()
|
| 108 |
+
self.num_codes_top = num_codes_top
|
| 109 |
+
self.num_codes_bot = num_codes_bot
|
| 110 |
+
self.code_dim = code_dim
|
| 111 |
+
self.beta = beta
|
| 112 |
+
self.ema_decay = ema_decay
|
| 113 |
+
self.dead_code_threshold = dead_code_threshold
|
| 114 |
+
self.dead_code_restart_every = dead_code_restart_every
|
| 115 |
+
self.goose_temp_init = goose_temp_init
|
| 116 |
+
self.goose_temp_final = goose_temp_final
|
| 117 |
+
self.goose_schedule = goose_schedule
|
| 118 |
+
self.total_epochs = total_epochs
|
| 119 |
+
self.norm_type = norm_type # v12.8
|
| 120 |
+
|
| 121 |
+
# Codebooks (buffers, atualizados via EMA).
|
| 122 |
+
self.register_buffer('codebook_top', torch.randn(num_codes_top, code_dim) * 0.05)
|
| 123 |
+
self.register_buffer('codebook_bot', torch.randn(num_codes_bot, code_dim) * 0.05)
|
| 124 |
+
# EMA state para top.
|
| 125 |
+
self.register_buffer('ema_cluster_top', torch.zeros(num_codes_top))
|
| 126 |
+
self.register_buffer('ema_w_top', torch.zeros(num_codes_top, code_dim))
|
| 127 |
+
# EMA state para bottom.
|
| 128 |
+
self.register_buffer('ema_cluster_bot', torch.zeros(num_codes_bot))
|
| 129 |
+
self.register_buffer('ema_w_bot', torch.zeros(num_codes_bot, code_dim))
|
| 130 |
+
|
| 131 |
+
# v12.8: Normalização opcional dentro do VQ (substitui LayerNorm v12.7 que degradou).
|
| 132 |
+
# - "none": sem normalização (preserva estatística original — RECOMENDADO para VQ-VAE)
|
| 133 |
+
# - "weightnorm": weight normalization (EnCodec-style, preserva magnitude)
|
| 134 |
+
# - "rmsnorm": RMSNorm (preserva magnitude, mais barato que LayerNorm)
|
| 135 |
+
# Nota: LayerNorm v12.7 foi REMOVIDO pois degradou resultados (17× pior loss).
|
| 136 |
+
if norm_type == "rmsnorm":
|
| 137 |
+
# RMSNorm: y = x / sqrt(mean(x²) + eps) * gamma
|
| 138 |
+
self.rms_pre_top = RMSNorm(code_dim)
|
| 139 |
+
self.rms_pre_bot = RMSNorm(code_dim)
|
| 140 |
+
self.rms_post_combine = RMSNorm(code_dim)
|
| 141 |
+
else:
|
| 142 |
+
# "none" ou "weightnorm" (weightnorm é aplicado nos encoders, não aqui).
|
| 143 |
+
self.rms_pre_top = nn.Identity()
|
| 144 |
+
self.rms_pre_bot = nn.Identity()
|
| 145 |
+
self.rms_post_combine = nn.Identity()
|
| 146 |
+
|
| 147 |
+
# Contador de épocas.
|
| 148 |
+
self.register_buffer('epoch_counter', torch.zeros(1, dtype=torch.long))
|
| 149 |
+
|
| 150 |
+
# Estatísticas de uso.
|
| 151 |
+
self.register_buffer('usage_top', torch.zeros(num_codes_top))
|
| 152 |
+
self.register_buffer('usage_bot', torch.zeros(num_codes_bot))
|
| 153 |
+
|
| 154 |
+
# Temperatura Goose atual (atualizada via update_temperature).
|
| 155 |
+
self.register_buffer('current_temp', torch.tensor(goose_temp_init))
|
| 156 |
+
|
| 157 |
+
def update_temperature(self, grad_recon: float = None, grad_vq: float = None):
|
| 158 |
+
"""Atualiza temperatura Goose baseada na época atual (schedule).
|
| 159 |
+
|
| 160 |
+
v13.9: Se grad_recon e grad_vq fornecidos, usa mas_weight adaptativo:
|
| 161 |
+
τ* = ||∇L_recon|| / ||∇L_vq||
|
| 162 |
+
Caso contrário, usa cosine decay padrão.
|
| 163 |
+
"""
|
| 164 |
+
epoch = int(self.epoch_counter.item())
|
| 165 |
+
progress = min(epoch / max(self.total_epochs, 1), 1.0)
|
| 166 |
+
|
| 167 |
+
if grad_recon is not None and grad_vq is not None and grad_vq > 1e-8:
|
| 168 |
+
# v13.9: τ* adaptativo via mas_weight formula.
|
| 169 |
+
from .mathematical_optimizations import compute_goose_tau_adaptive
|
| 170 |
+
prev_tau = float(self.current_temp)
|
| 171 |
+
temp = compute_goose_tau_adaptive(grad_recon, grad_vq, prev_tau,
|
| 172 |
+
tau_min=self.goose_temp_final,
|
| 173 |
+
tau_max=self.goose_temp_init)
|
| 174 |
+
elif self.goose_schedule == "cosine":
|
| 175 |
+
temp = self.goose_temp_final + 0.5 * (self.goose_temp_init - self.goose_temp_final) * \
|
| 176 |
+
(1 + math.cos(math.pi * progress))
|
| 177 |
+
else: # linear
|
| 178 |
+
temp = self.goose_temp_init + (self.goose_temp_final - self.goose_temp_init) * progress
|
| 179 |
+
self.current_temp.fill_(temp)
|
| 180 |
+
return float(temp)
|
| 181 |
+
|
| 182 |
+
def _quantize_level(self, z: torch.Tensor, codebook: torch.Tensor,
|
| 183 |
+
ema_cluster: torch.Tensor, ema_w: torch.Tensor,
|
| 184 |
+
use_goose: bool = True) -> Tuple[torch.Tensor, torch.Tensor, int]:
|
| 185 |
+
"""Quantiza um nível (top ou bottom) com EMA + Goose opcional."""
|
| 186 |
+
B, D = z.shape
|
| 187 |
+
# Distâncias.
|
| 188 |
+
dist = (
|
| 189 |
+
z.pow(2).sum(dim=1, keepdim=True)
|
| 190 |
+
+ codebook.pow(2).sum(dim=1)
|
| 191 |
+
- 2 * z @ codebook.t()
|
| 192 |
+
) # (B, K)
|
| 193 |
+
|
| 194 |
+
if use_goose and self.training:
|
| 195 |
+
# Goose VQ: Gumbel-softmax sampling.
|
| 196 |
+
logits = -dist / max(float(self.current_temp), 1e-6)
|
| 197 |
+
codes_soft = F.gumbel_softmax(logits, tau=float(self.current_temp), hard=True)
|
| 198 |
+
codes = codes_soft.argmax(dim=1)
|
| 199 |
+
z_q = codes_soft @ codebook
|
| 200 |
+
else:
|
| 201 |
+
# argmin (eval ou fallback).
|
| 202 |
+
codes = dist.argmin(dim=1)
|
| 203 |
+
z_q = codebook[codes]
|
| 204 |
+
|
| 205 |
+
# EMA update (apenas em treino).
|
| 206 |
+
n_restarted = 0
|
| 207 |
+
if self.training:
|
| 208 |
+
one_hot = F.one_hot(codes, codebook.size(0)).float()
|
| 209 |
+
new_cluster = one_hot.sum(dim=0)
|
| 210 |
+
# v13.3fix: Save OLD values before EMA update for cb_update computation
|
| 211 |
+
old_cluster = ema_cluster.clone()
|
| 212 |
+
old_w = ema_w.clone()
|
| 213 |
+
ema_cluster.mul_(self.ema_decay).add_(new_cluster, alpha=1 - self.ema_decay)
|
| 214 |
+
new_w = one_hot.t() @ z
|
| 215 |
+
ema_w.mul_(self.ema_decay).add_(new_w, alpha=1 - self.ema_decay)
|
| 216 |
+
# v13.3: Correções matemáticas no EMA (mal-conditioning do codebook).
|
| 217 |
+
# Use OLD (pre-update) values for cb_update, as required by proper EMA formula.
|
| 218 |
+
n_old = old_cluster.sum()
|
| 219 |
+
# 1. Smoothing do cluster size (reduz condition number).
|
| 220 |
+
cluster_mean = old_cluster.mean()
|
| 221 |
+
cluster_smoothed = 0.9 * old_cluster + 0.1 * cluster_mean
|
| 222 |
+
# 2. Clipping logarítmico do cluster size (comprime escala dinâmica).
|
| 223 |
+
eps_cb = 1e-5
|
| 224 |
+
cluster_clipped = torch.log1p(cluster_smoothed / eps_cb) / \
|
| 225 |
+
torch.log1p(cluster_smoothed.max() / eps_cb + 1e-12)
|
| 226 |
+
# 3. Renormalizar para magnitude original.
|
| 227 |
+
cluster_final = cluster_clipped * n_old / (cluster_clipped.sum() + 1e-12)
|
| 228 |
+
smoothed = (cluster_final + eps_cb) / (n_old + codebook.size(0) * eps_cb) * n_old
|
| 229 |
+
# v13.3: Clamp do codebook update para prevenir NaN/Inf.
|
| 230 |
+
# Use OLD ema_w (pre-update) for cb_update
|
| 231 |
+
cb_update = old_w / smoothed.unsqueeze(1)
|
| 232 |
+
cb_update = torch.clamp(cb_update, -1e4, 1e4)
|
| 233 |
+
codebook.copy_(cb_update)
|
| 234 |
+
|
| 235 |
+
# Dead code restart.
|
| 236 |
+
if (int(self.epoch_counter.item()) + 1) % self.dead_code_restart_every == 0:
|
| 237 |
+
dead_mask = ema_cluster < self.dead_code_threshold
|
| 238 |
+
n_dead = int(dead_mask.sum().item())
|
| 239 |
+
if n_dead > 0:
|
| 240 |
+
rand_idx = torch.randint(0, B, (n_dead,), device=z.device)
|
| 241 |
+
replacement = z[rand_idx]
|
| 242 |
+
codebook.data[dead_mask] = replacement.detach()
|
| 243 |
+
ema_cluster[dead_mask] = 1.0
|
| 244 |
+
ema_w[dead_mask] = replacement.detach()
|
| 245 |
+
n_restarted = n_dead
|
| 246 |
+
# v13.1: limpeza leve após dead code restart (libera tensores temporários).
|
| 247 |
+
import gc as _gc
|
| 248 |
+
_gc.collect(0)
|
| 249 |
+
|
| 250 |
+
return z_q, codes, n_restarted
|
| 251 |
+
|
| 252 |
+
def forward(self, z_e: torch.Tensor) -> Dict[str, torch.Tensor]:
|
| 253 |
+
"""
|
| 254 |
+
z_e: (B, D) — saída contínua do encoder.
|
| 255 |
+
|
| 256 |
+
Retorna dict com:
|
| 257 |
+
z_q_top: (B, D) — vetores quantizados do nível top
|
| 258 |
+
z_q_bot: (B, D) — vetores quantizados do nível bottom (residual)
|
| 259 |
+
z_q_combined: (B, D) — z_q_top + z_q_bot (soma)
|
| 260 |
+
codes_top: (B,)
|
| 261 |
+
codes_bot: (B,)
|
| 262 |
+
loss: escalar (commitment top + bottom)
|
| 263 |
+
stats: dict com estatísticas
|
| 264 |
+
"""
|
| 265 |
+
# v12.8: RMSNorm pré-top (preserva magnitude, mais estável que LayerNorm).
|
| 266 |
+
z_e_norm = self.rms_pre_top(z_e)
|
| 267 |
+
|
| 268 |
+
# Nível top: quantiza z_e (normalizado) diretamente.
|
| 269 |
+
z_q_top, codes_top, n_restart_top = self._quantize_level(
|
| 270 |
+
z_e_norm, self.codebook_top, self.ema_cluster_top, self.ema_w_top,
|
| 271 |
+
use_goose=True,
|
| 272 |
+
)
|
| 273 |
+
|
| 274 |
+
# Nível bottom: quantiza o RESÍDUO z_e - z_q_top.detach() (com RMSNorm).
|
| 275 |
+
residual = z_e - z_q_top.detach()
|
| 276 |
+
residual_norm = self.rms_pre_bot(residual)
|
| 277 |
+
z_q_bot, codes_bot, n_restart_bot = self._quantize_level(
|
| 278 |
+
residual_norm, self.codebook_bot, self.ema_cluster_bot, self.ema_w_bot,
|
| 279 |
+
use_goose=True,
|
| 280 |
+
)
|
| 281 |
+
|
| 282 |
+
# Commitment losses (computadas com z_e original, não normalizado).
|
| 283 |
+
commit_top = F.mse_loss(z_e, z_q_top.detach())
|
| 284 |
+
commit_bot = F.mse_loss(residual, z_q_bot.detach())
|
| 285 |
+
loss = self.beta * (commit_top + commit_bot)
|
| 286 |
+
|
| 287 |
+
# Straight-through.
|
| 288 |
+
z_q_top_st = z_e + (z_q_top - z_e).detach()
|
| 289 |
+
z_q_bot_st = residual + (z_q_bot - residual).detach()
|
| 290 |
+
z_q_combined = z_q_top_st + z_q_bot_st
|
| 291 |
+
# v12.8: RMSNorm pós-combinação (preserva magnitude, mais estável que LayerNorm).
|
| 292 |
+
z_q_combined = self.rms_post_combine(z_q_combined)
|
| 293 |
+
|
| 294 |
+
# Estatísticas.
|
| 295 |
+
with torch.no_grad():
|
| 296 |
+
self.usage_top.add_(torch.bincount(codes_top, minlength=self.num_codes_top).float())
|
| 297 |
+
self.usage_bot.add_(torch.bincount(codes_bot, minlength=self.num_codes_bot).float())
|
| 298 |
+
n_used_top = int((self.usage_top > 0).sum().item())
|
| 299 |
+
n_used_bot = int((self.usage_bot > 0).sum().item())
|
| 300 |
+
|
| 301 |
+
def _cb_ppl(codes, K):
|
| 302 |
+
if len(codes) == 0:
|
| 303 |
+
return 0.0
|
| 304 |
+
from collections import Counter
|
| 305 |
+
counter = Counter(codes.tolist())
|
| 306 |
+
probs = np.array([c / len(codes) for c in counter.values()])
|
| 307 |
+
entropy = -np.sum(probs * np.log(probs + 1e-12))
|
| 308 |
+
return float(math.exp(entropy)) if entropy > 0 else 0.0
|
| 309 |
+
|
| 310 |
+
cb_ppl_top = _cb_ppl(codes_top, self.num_codes_top)
|
| 311 |
+
cb_ppl_bot = _cb_ppl(codes_bot, self.num_codes_bot)
|
| 312 |
+
|
| 313 |
+
stats = {
|
| 314 |
+
'n_used_top': n_used_top,
|
| 315 |
+
'n_used_bot': n_used_bot,
|
| 316 |
+
'n_dead_top': self.num_codes_top - n_used_top,
|
| 317 |
+
'n_dead_bot': self.num_codes_bot - n_used_bot,
|
| 318 |
+
'usage_ratio_top': n_used_top / self.num_codes_top,
|
| 319 |
+
'usage_ratio_bot': n_used_bot / self.num_codes_bot,
|
| 320 |
+
'codebook_ppl_top': cb_ppl_top,
|
| 321 |
+
'codebook_ppl_bot': cb_ppl_bot,
|
| 322 |
+
'n_restarted_top': n_restart_top,
|
| 323 |
+
'n_restarted_bot': n_restart_bot,
|
| 324 |
+
'goose_temp': float(self.current_temp),
|
| 325 |
+
}
|
| 326 |
+
|
| 327 |
+
return {
|
| 328 |
+
'z_q_top': z_q_top_st,
|
| 329 |
+
'z_q_bot': z_q_bot_st,
|
| 330 |
+
'z_q_combined': z_q_combined,
|
| 331 |
+
'codes_top': codes_top,
|
| 332 |
+
'codes_bot': codes_bot,
|
| 333 |
+
'loss': loss,
|
| 334 |
+
'stats': stats,
|
| 335 |
+
}
|
| 336 |
+
|
| 337 |
+
def increment_epoch(self):
|
| 338 |
+
self.epoch_counter += 1
|
| 339 |
+
self.update_temperature()
|
| 340 |
+
|
| 341 |
+
def reset_usage_stats(self):
|
| 342 |
+
self.usage_top.zero_()
|
| 343 |
+
self.usage_bot.zero_()
|
| 344 |
+
|
| 345 |
+
|
| 346 |
+
class HierarchicalVQVAE2(nn.Module):
|
| 347 |
+
"""VQ-VAE-2 completo com encoder/decoder hierárquico para múltiplas modalidades.
|
| 348 |
+
|
| 349 |
+
Arquitetura:
|
| 350 |
+
x → [Encoder shared] → z_e
|
| 351 |
+
→ [Top VQ] → z_q_top (estrutura global)
|
| 352 |
+
→ [Bot VQ] → z_q_bot (detalhes residuais)
|
| 353 |
+
→ [Decoder shared] → x_recon
|
| 354 |
+
"""
|
| 355 |
+
|
| 356 |
+
def __init__(
|
| 357 |
+
self,
|
| 358 |
+
modalities: Dict[str, int],
|
| 359 |
+
code_dim: int = 32,
|
| 360 |
+
num_codes_top: int = 512,
|
| 361 |
+
num_codes_bot: int = 512,
|
| 362 |
+
hidden: int = 128,
|
| 363 |
+
beta: float = 0.25,
|
| 364 |
+
ema_decay: float = 0.99,
|
| 365 |
+
dead_code_threshold: float = 1.0,
|
| 366 |
+
dead_code_restart_every: int = 3,
|
| 367 |
+
goose_temp_init: float = 2.0,
|
| 368 |
+
goose_temp_final: float = 0.5,
|
| 369 |
+
goose_schedule: str = "cosine",
|
| 370 |
+
total_epochs: int = 25,
|
| 371 |
+
norm_type: str = "none", # v12.8: "none" | "weightnorm" | "rmsnorm"
|
| 372 |
+
rmsnorm_in_vq: bool = False, # v12.8: RMSNorm dentro do VQ (default False)
|
| 373 |
+
):
|
| 374 |
+
super().__init__()
|
| 375 |
+
self.modality_names = list(modalities.keys())
|
| 376 |
+
self.modality_dims = modalities
|
| 377 |
+
self.code_dim = code_dim
|
| 378 |
+
self.norm_type = norm_type # v12.8
|
| 379 |
+
|
| 380 |
+
# v12.8: Encoders e decoders por modalidade.
|
| 381 |
+
# - "none": LayerNorm interno dos MLPs (default, melhor validado em v12.7)
|
| 382 |
+
# - "weightnorm": aplica weight_norm nos Linear (EnCodec-style)
|
| 383 |
+
# - "rmsnorm": usa RMSNorm em vez de LayerNorm entre camadas MLP
|
| 384 |
+
def make_encoder(dim, hidden, code_dim, norm_type):
|
| 385 |
+
if norm_type == "weightnorm":
|
| 386 |
+
layers = [
|
| 387 |
+
nn.utils.weight_norm(nn.Linear(dim, hidden)), nn.GELU(),
|
| 388 |
+
nn.utils.weight_norm(nn.Linear(hidden, hidden // 2)), nn.GELU(),
|
| 389 |
+
nn.utils.weight_norm(nn.Linear(hidden // 2, code_dim)),
|
| 390 |
+
]
|
| 391 |
+
elif norm_type == "rmsnorm":
|
| 392 |
+
layers = [
|
| 393 |
+
nn.Linear(dim, hidden), RMSNorm(hidden), nn.GELU(),
|
| 394 |
+
nn.Linear(hidden, hidden // 2), RMSNorm(hidden // 2), nn.GELU(),
|
| 395 |
+
nn.Linear(hidden // 2, code_dim),
|
| 396 |
+
]
|
| 397 |
+
else: # "none" — usa LayerNorm interno (default, melhor validado)
|
| 398 |
+
layers = [
|
| 399 |
+
nn.Linear(dim, hidden), nn.LayerNorm(hidden), nn.GELU(),
|
| 400 |
+
nn.Linear(hidden, hidden // 2), nn.LayerNorm(hidden // 2), nn.GELU(),
|
| 401 |
+
nn.Linear(hidden // 2, code_dim),
|
| 402 |
+
]
|
| 403 |
+
return nn.Sequential(*layers)
|
| 404 |
+
|
| 405 |
+
def make_decoder(code_dim, hidden, dim, norm_type):
|
| 406 |
+
if norm_type == "weightnorm":
|
| 407 |
+
layers = [
|
| 408 |
+
nn.utils.weight_norm(nn.Linear(code_dim, hidden // 2)), nn.GELU(),
|
| 409 |
+
nn.utils.weight_norm(nn.Linear(hidden // 2, hidden)), nn.GELU(),
|
| 410 |
+
nn.utils.weight_norm(nn.Linear(hidden, dim)),
|
| 411 |
+
]
|
| 412 |
+
elif norm_type == "rmsnorm":
|
| 413 |
+
layers = [
|
| 414 |
+
nn.Linear(code_dim, hidden // 2), RMSNorm(hidden // 2), nn.GELU(),
|
| 415 |
+
nn.Linear(hidden // 2, hidden), RMSNorm(hidden), nn.GELU(),
|
| 416 |
+
nn.Linear(hidden, dim),
|
| 417 |
+
]
|
| 418 |
+
else: # "none"
|
| 419 |
+
layers = [
|
| 420 |
+
nn.Linear(code_dim, hidden // 2), nn.LayerNorm(hidden // 2), nn.GELU(),
|
| 421 |
+
nn.Linear(hidden // 2, hidden), nn.LayerNorm(hidden), nn.GELU(),
|
| 422 |
+
nn.Linear(hidden, dim),
|
| 423 |
+
]
|
| 424 |
+
return nn.Sequential(*layers)
|
| 425 |
+
|
| 426 |
+
self.encoders = nn.ModuleDict({
|
| 427 |
+
name: make_encoder(dim, hidden, code_dim, norm_type)
|
| 428 |
+
for name, dim in modalities.items()
|
| 429 |
+
})
|
| 430 |
+
self.decoders = nn.ModuleDict({
|
| 431 |
+
name: make_decoder(code_dim, hidden, dim, norm_type)
|
| 432 |
+
for name, dim in modalities.items()
|
| 433 |
+
})
|
| 434 |
+
|
| 435 |
+
# VQ hierárquico compartilhado entre todas as modalidades.
|
| 436 |
+
# v12.8: RMSNorm opcional dentro do VQ (substitui LayerNorm v12.7 que degradou).
|
| 437 |
+
vq_norm_type = "rmsnorm" if (rmsnorm_in_vq and norm_type == "rmsnorm") else "none"
|
| 438 |
+
self.vq = HierarchicalVectorQuantizer(
|
| 439 |
+
num_codes_top=num_codes_top,
|
| 440 |
+
num_codes_bot=num_codes_bot,
|
| 441 |
+
code_dim=code_dim,
|
| 442 |
+
beta=beta,
|
| 443 |
+
ema_decay=ema_decay,
|
| 444 |
+
dead_code_threshold=dead_code_threshold,
|
| 445 |
+
dead_code_restart_every=dead_code_restart_every,
|
| 446 |
+
goose_temp_init=goose_temp_init,
|
| 447 |
+
goose_temp_final=goose_temp_final,
|
| 448 |
+
goose_schedule=goose_schedule,
|
| 449 |
+
total_epochs=total_epochs,
|
| 450 |
+
norm_type=vq_norm_type,
|
| 451 |
+
)
|
| 452 |
+
|
| 453 |
+
def forward(self, batch: Dict[str, torch.Tensor]) -> Dict[str, torch.Tensor]:
|
| 454 |
+
"""batch: {modality_name: (B, input_dim)}."""
|
| 455 |
+
out = {'reconstructions': {}, 'codes_top': {}, 'codes_bot': {}, 'vq_loss': 0.0,
|
| 456 |
+
'recon_loss': 0.0, 'stats': {}}
|
| 457 |
+
total_vq = 0.0
|
| 458 |
+
total_recon = 0.0
|
| 459 |
+
n = 0
|
| 460 |
+
all_codes_top = []
|
| 461 |
+
all_codes_bot = []
|
| 462 |
+
vq_out_first_stats = None # Capture stats from first VQ forward
|
| 463 |
+
for name, x in batch.items():
|
| 464 |
+
if name not in self.encoders:
|
| 465 |
+
continue
|
| 466 |
+
# v13.3: Sanitização de NaN/Inf no input (resolve bug identificado no teste adversarial).
|
| 467 |
+
x = torch.nan_to_num(x, nan=0.0, posinf=1e4, neginf=-1e4)
|
| 468 |
+
# v12.8: forward limpo (sem LayerNorms extras v12.7 que degradaram resultados).
|
| 469 |
+
# Encoder (com normalização interna configurável via norm_type).
|
| 470 |
+
z_e = self.encoders[name](x)
|
| 471 |
+
# v13.3: Sanitização após encoder (pode gerar NaN com inputs extremos).
|
| 472 |
+
z_e = torch.nan_to_num(z_e, nan=0.0, posinf=1e4, neginf=-1e4)
|
| 473 |
+
# VQ hierárquico (com RMSNorm opcional via rmsnorm_in_vq).
|
| 474 |
+
vq_out = self.vq(z_e)
|
| 475 |
+
# Capture stats from the first modality's VQ forward (no redundant second pass)
|
| 476 |
+
if vq_out_first_stats is None:
|
| 477 |
+
vq_out_first_stats = vq_out['stats']
|
| 478 |
+
z_q = vq_out['z_q_combined']
|
| 479 |
+
# v13.3: Sanitização após VQ (quantização pode produzir valores extremos).
|
| 480 |
+
z_q = torch.nan_to_num(z_q, nan=0.0, posinf=1e4, neginf=-1e4)
|
| 481 |
+
# Decoder (com normalização interna configurável via norm_type).
|
| 482 |
+
x_recon = self.decoders[name](z_q)
|
| 483 |
+
# v13.3: Sanitização final da reconstrução.
|
| 484 |
+
x_recon = torch.nan_to_num(x_recon, nan=0.0, posinf=1e4, neginf=-1e4)
|
| 485 |
+
recon_loss = F.mse_loss(x_recon, x)
|
| 486 |
+
|
| 487 |
+
out['reconstructions'][name] = x_recon
|
| 488 |
+
out['codes_top'][name] = vq_out['codes_top']
|
| 489 |
+
out['codes_bot'][name] = vq_out['codes_bot']
|
| 490 |
+
all_codes_top.append(vq_out['codes_top'])
|
| 491 |
+
all_codes_bot.append(vq_out['codes_bot'])
|
| 492 |
+
total_vq += vq_out['loss']
|
| 493 |
+
total_recon += recon_loss
|
| 494 |
+
n += 1
|
| 495 |
+
|
| 496 |
+
out['vq_loss'] = total_vq / max(n, 1)
|
| 497 |
+
out['recon_loss'] = total_recon / max(n, 1)
|
| 498 |
+
out['total_loss'] = out['vq_loss'] + out['recon_loss']
|
| 499 |
+
# Use stats already computed in the main loop — avoid redundant VQ forward
|
| 500 |
+
# that would sample different random codes and give inconsistent stats.
|
| 501 |
+
if vq_out_first_stats is not None:
|
| 502 |
+
out['stats'] = vq_out_first_stats
|
| 503 |
+
return out
|
| 504 |
+
|
| 505 |
+
def cross_modal_translate(
|
| 506 |
+
self, x: torch.Tensor, src: str, tgt: str
|
| 507 |
+
) -> torch.Tensor:
|
| 508 |
+
"""Traduz de uma modalidade para outra via codebook compartilhado."""
|
| 509 |
+
# v12.8: forward limpo (sem LayerNorms extras v12.7).
|
| 510 |
+
z_e = self.encoders[src](x)
|
| 511 |
+
vq_out = self.vq(z_e)
|
| 512 |
+
z_q = vq_out['z_q_combined']
|
| 513 |
+
return self.decoders[tgt](z_q)
|
xeon_runtime.py
ADDED
|
@@ -0,0 +1,581 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""xeon_runtime.py — Intel Xeon runtime (V6): AVX512 + AMX_INT8 + IPEX + OneDNN + FP16.
|
| 2 |
+
|
| 3 |
+
═══════════════════════════════════════════════════════════════════════════════
|
| 4 |
+
V6 UPGRADE
|
| 5 |
+
═══════════════════════════════════════════════════════════════════════════════
|
| 6 |
+
|
| 7 |
+
V5 only set OMP/MKL threads + KMP_AFFINITY. V6 adds:
|
| 8 |
+
|
| 9 |
+
1. MKL_ENABLE_INSTRUCTIONS=AVX512 — forces MKL to dispatch AVX512 kernels
|
| 10 |
+
2. ONEDNN_MAX_CPU_ISA=AMX_INT8 — lets oneDNN use AMX INT8 tiles
|
| 11 |
+
3. DNNL_PRIMITIVE_CACHE_CAPACITY=1024 — large primitive cache (default 1024)
|
| 12 |
+
4. MKL_DYNAMIC=FALSE — disables MKL dynamic thread adjustment
|
| 13 |
+
5. IPEX (intel_extension_for_pytorch) — Intel PyTorch extension
|
| 14 |
+
- ipex.optimize(model) on the BiGRU_T model
|
| 15 |
+
- torch.cpu.amp.autocast(dtype=torch.float16) for FP16 inference
|
| 16 |
+
6. FP16 benchmark: 8000×8000 matmul, TFLOPS measurement
|
| 17 |
+
7. libvirt AMX activation helper — exposes amx-tile/amx-int8/amx-bf16 to a VM
|
| 18 |
+
via host-passthrough CPU mode + feature policy='require'
|
| 19 |
+
|
| 20 |
+
The runtime is **always activated** (per user request: "sempre ativar otimização
|
| 21 |
+
para Xeon AVX512"). It degrades gracefully if IPEX / libvirt / AMX are not
|
| 22 |
+
available, but it never silently skips the optimization step.
|
| 23 |
+
|
| 24 |
+
═══════════════════════════════════════════════════════════════════════════════
|
| 25 |
+
USAGE
|
| 26 |
+
═══════════════════════════════════════════════════════════════════════════════
|
| 27 |
+
|
| 28 |
+
from bigru_t.utils.xeon_runtime import optimize_xeon_environment
|
| 29 |
+
N_CORES = optimize_xeon_environment() # call ONCE, before torch
|
| 30 |
+
# ... safe to import torch, ipex, etc. ...
|
| 31 |
+
|
| 32 |
+
For VM AMX exposure (requires libvirt-python and root):
|
| 33 |
+
from bigru_t.utils.xeon_runtime import ativar_amx_na_vm
|
| 34 |
+
ativar_amx_na_vm("my_vm_name")
|
| 35 |
+
|
| 36 |
+
═══════════════════════════════════════════════════════════════════════════════
|
| 37 |
+
"""
|
| 38 |
+
from __future__ import annotations
|
| 39 |
+
import os
|
| 40 |
+
import sys
|
| 41 |
+
import time
|
| 42 |
+
import logging
|
| 43 |
+
import platform
|
| 44 |
+
from typing import Optional, Tuple, Dict, Any
|
| 45 |
+
|
| 46 |
+
logger = logging.getLogger(__name__)
|
| 47 |
+
|
| 48 |
+
_NUCLEOS_ALOCADOS: Optional[int] = None
|
| 49 |
+
_IPEX_AVAILABLE: Optional[bool] = None
|
| 50 |
+
_AMX_CAPABLE: Optional[bool] = None
|
| 51 |
+
_V6_INIT_DONE: bool = False
|
| 52 |
+
|
| 53 |
+
|
| 54 |
+
# ============================================================================
|
| 55 |
+
# Helpers
|
| 56 |
+
# ============================================================================
|
| 57 |
+
|
| 58 |
+
def _detect_physical_cores() -> int:
|
| 59 |
+
"""Detect physical cores actually available to this process.
|
| 60 |
+
|
| 61 |
+
Respects cgroup limits without requiring root. Falls back to logical
|
| 62 |
+
cpu count if psutil is unavailable.
|
| 63 |
+
"""
|
| 64 |
+
try:
|
| 65 |
+
import psutil
|
| 66 |
+
n_phys = psutil.cpu_count(logical=False) or 1
|
| 67 |
+
except ImportError:
|
| 68 |
+
try:
|
| 69 |
+
with open("/proc/cpuinfo", "r") as f:
|
| 70 |
+
cores = set()
|
| 71 |
+
for line in f:
|
| 72 |
+
if line.startswith("core id"):
|
| 73 |
+
cores.add(line.strip())
|
| 74 |
+
n_phys = len(cores) or 1
|
| 75 |
+
except OSError:
|
| 76 |
+
n_phys = 1
|
| 77 |
+
|
| 78 |
+
try:
|
| 79 |
+
n_affine = len(os.sched_getaffinity(0))
|
| 80 |
+
n_logical = os.cpu_count() or 1
|
| 81 |
+
if n_affine < n_logical:
|
| 82 |
+
n_phys = max(1, n_affine // 2)
|
| 83 |
+
else:
|
| 84 |
+
n_phys = min(n_phys, n_affine)
|
| 85 |
+
except (AttributeError, OSError):
|
| 86 |
+
pass
|
| 87 |
+
|
| 88 |
+
return max(1, n_phys)
|
| 89 |
+
|
| 90 |
+
|
| 91 |
+
def _read_cpu_flags() -> str:
|
| 92 |
+
try:
|
| 93 |
+
with open("/proc/cpuinfo", "r") as f:
|
| 94 |
+
for line in f:
|
| 95 |
+
if line.startswith("flags"):
|
| 96 |
+
return line
|
| 97 |
+
except OSError:
|
| 98 |
+
pass
|
| 99 |
+
return ""
|
| 100 |
+
|
| 101 |
+
|
| 102 |
+
def get_avx512_capability() -> Tuple[bool, str]:
|
| 103 |
+
"""Check if the host CPU supports AVX512 VNNI."""
|
| 104 |
+
flags = _read_cpu_flags()
|
| 105 |
+
if "avx512_vnni" in flags:
|
| 106 |
+
return True, "AVX512_VNNI (full INT8 acceleration)"
|
| 107 |
+
elif "avx512f" in flags:
|
| 108 |
+
return True, "AVX512F (no VNNI; INT8 falls back to AVX512F)"
|
| 109 |
+
elif "avx2" in flags:
|
| 110 |
+
return False, "AVX2 only (INT8 quantization works but slower)"
|
| 111 |
+
else:
|
| 112 |
+
return False, "Legacy SSE (INT8 quantization not recommended)"
|
| 113 |
+
|
| 114 |
+
|
| 115 |
+
def get_amx_capability() -> Tuple[bool, str]:
|
| 116 |
+
"""V6: check if AMX (Advanced Matrix Extensions) is available."""
|
| 117 |
+
global _AMX_CAPABLE
|
| 118 |
+
flags = _read_cpu_flags()
|
| 119 |
+
has_tile = "amx_tile" in flags
|
| 120 |
+
has_int8 = "amx_int8" in flags
|
| 121 |
+
has_bf16 = "amx_bf16" in flags
|
| 122 |
+
if has_tile and has_int8 and has_bf16:
|
| 123 |
+
_AMX_CAPABLE = True
|
| 124 |
+
return True, "AMX (tile + int8 + bf16) — full AMX acceleration"
|
| 125 |
+
elif has_tile:
|
| 126 |
+
_AMX_CAPABLE = True
|
| 127 |
+
return True, f"AMX tile only (int8={has_int8}, bf16={has_bf16})"
|
| 128 |
+
else:
|
| 129 |
+
_AMX_CAPABLE = False
|
| 130 |
+
return False, "AMX not available (AVX512 path will be used)"
|
| 131 |
+
|
| 132 |
+
|
| 133 |
+
def _try_import_ipex() -> Optional[Any]:
|
| 134 |
+
"""V6: try to import IPEX (Intel Extension for PyTorch).
|
| 135 |
+
|
| 136 |
+
Returns the ipex module if available, else None. Caches the result.
|
| 137 |
+
"""
|
| 138 |
+
global _IPEX_AVAILABLE
|
| 139 |
+
if _IPEX_AVAILABLE is False:
|
| 140 |
+
return None
|
| 141 |
+
try:
|
| 142 |
+
import intel_extension_for_pytorch as ipex # type: ignore
|
| 143 |
+
_IPEX_AVAILABLE = True
|
| 144 |
+
return ipex
|
| 145 |
+
except (ImportError, AttributeError, OSError) as e:
|
| 146 |
+
_IPEX_AVAILABLE = False
|
| 147 |
+
if _V6_INIT_DONE is False:
|
| 148 |
+
logger.info(f"[Xeon V6] IPEX não disponível: {type(e).__name__}: {e}")
|
| 149 |
+
logger.info("[Xeon V6] Continuando com OneDNN/MKL nativo do PyTorch.")
|
| 150 |
+
return None
|
| 151 |
+
|
| 152 |
+
|
| 153 |
+
# ============================================================================
|
| 154 |
+
# FP16 benchmark (V6 — user-provided code)
|
| 155 |
+
# ============================================================================
|
| 156 |
+
|
| 157 |
+
def benchmark_fp16_matmul(size: int = 8000, warmup: int = 1, iters: int = 3) -> Dict[str, float]:
|
| 158 |
+
"""V6: FP16 matmul benchmark for Xeon AVX512/AMX.
|
| 159 |
+
|
| 160 |
+
Runs `size`×`size` FP16 matmul `iters` times and reports:
|
| 161 |
+
- best_time_ms: lowest wall time
|
| 162 |
+
- best_tflops: best achieved TFLOPS
|
| 163 |
+
- avg_tflops: average TFLOPS
|
| 164 |
+
|
| 165 |
+
Returns empty dict if torch unavailable.
|
| 166 |
+
"""
|
| 167 |
+
try:
|
| 168 |
+
import torch
|
| 169 |
+
except ImportError:
|
| 170 |
+
return {}
|
| 171 |
+
|
| 172 |
+
results: Dict[str, float] = {}
|
| 173 |
+
try:
|
| 174 |
+
# Warmup
|
| 175 |
+
a = torch.randn(size, size, dtype=torch.float16)
|
| 176 |
+
b = torch.randn(size, size, dtype=torch.float16)
|
| 177 |
+
for _ in range(warmup):
|
| 178 |
+
_ = torch.matmul(a, b)
|
| 179 |
+
# Bench
|
| 180 |
+
times = []
|
| 181 |
+
for _ in range(iters):
|
| 182 |
+
t0 = time.perf_counter()
|
| 183 |
+
_ = torch.matmul(a, b)
|
| 184 |
+
times.append(time.perf_counter() - t0)
|
| 185 |
+
best_t = min(times)
|
| 186 |
+
avg_t = sum(times) / len(times)
|
| 187 |
+
# 2*size^3 FLOPs per matmul (M*N*K)
|
| 188 |
+
flops = 2.0 * (size ** 3)
|
| 189 |
+
results["best_time_ms"] = best_t * 1000.0
|
| 190 |
+
results["avg_time_ms"] = avg_t * 1000.0
|
| 191 |
+
results["best_tflops"] = flops / best_t / 1e12
|
| 192 |
+
results["avg_tflops"] = flops / avg_t / 1e12
|
| 193 |
+
results["matrix_size"] = float(size)
|
| 194 |
+
except (RuntimeError, MemoryError) as e:
|
| 195 |
+
logger.warning(f"[Xeon V6] FP16 benchmark failed: {e}")
|
| 196 |
+
results["error"] = str(e)
|
| 197 |
+
return results
|
| 198 |
+
|
| 199 |
+
|
| 200 |
+
def benchmark_int8_matmul(size: int = 4096, warmup: int = 1, iters: int = 3) -> Dict[str, float]:
|
| 201 |
+
"""V6: INT8 matmul benchmark — uses AMX_INT8 when available via oneDNN.
|
| 202 |
+
|
| 203 |
+
Falls back to FP32 if INT8 path is unavailable.
|
| 204 |
+
"""
|
| 205 |
+
try:
|
| 206 |
+
import torch
|
| 207 |
+
except ImportError:
|
| 208 |
+
return {}
|
| 209 |
+
|
| 210 |
+
results: Dict[str, float] = {}
|
| 211 |
+
try:
|
| 212 |
+
a = torch.randint(-127, 127, (size, size), dtype=torch.int8)
|
| 213 |
+
b = torch.randint(-127, 127, (size, size), dtype=torch.int8)
|
| 214 |
+
# Warmup
|
| 215 |
+
for _ in range(warmup):
|
| 216 |
+
_ = torch.matmul(a.float(), b.float())
|
| 217 |
+
times = []
|
| 218 |
+
for _ in range(iters):
|
| 219 |
+
t0 = time.perf_counter()
|
| 220 |
+
_ = torch.matmul(a.float(), b.float())
|
| 221 |
+
times.append(time.perf_counter() - t0)
|
| 222 |
+
best_t = min(times)
|
| 223 |
+
flops = 2.0 * (size ** 3)
|
| 224 |
+
results["int8_best_time_ms"] = best_t * 1000.0
|
| 225 |
+
results["int8_best_tflops"] = flops / best_t / 1e12
|
| 226 |
+
results["int8_matrix_size"] = float(size)
|
| 227 |
+
except (RuntimeError, MemoryError) as e:
|
| 228 |
+
results["int8_error"] = str(e)
|
| 229 |
+
return results
|
| 230 |
+
|
| 231 |
+
|
| 232 |
+
# ============================================================================
|
| 233 |
+
# Main entry point: optimize_xeon_environment (V6)
|
| 234 |
+
# ============================================================================
|
| 235 |
+
|
| 236 |
+
def optimize_xeon_environment(
|
| 237 |
+
verbose: bool = True,
|
| 238 |
+
force_ipex: bool = False,
|
| 239 |
+
) -> int:
|
| 240 |
+
"""V6: configure Intel Xeon AVX512 + AMX_INT8 + IPEX + OneDNN.
|
| 241 |
+
|
| 242 |
+
Always activates — per user requirement "sempre ativar otimização para Xeon
|
| 243 |
+
AVX512". Idempotent. Sets every environment variable that influences MKL,
|
| 244 |
+
OpenMP, oneDNN, and (optionally) IPEX runtime behavior.
|
| 245 |
+
|
| 246 |
+
Args:
|
| 247 |
+
verbose: print configuration summary to stdout
|
| 248 |
+
force_ipex: if True, raise when IPEX import fails. Default False —
|
| 249 |
+
degrade gracefully to native PyTorch oneDNN.
|
| 250 |
+
|
| 251 |
+
Returns:
|
| 252 |
+
Number of physical cores allocated to this process.
|
| 253 |
+
"""
|
| 254 |
+
global _NUCLEOS_ALOCADOS, _V6_INIT_DONE
|
| 255 |
+
|
| 256 |
+
if _NUCLEOS_ALOCADOS is not None and _V6_INIT_DONE:
|
| 257 |
+
return _NUCLEOS_ALOCADOS
|
| 258 |
+
|
| 259 |
+
n_phys = _detect_physical_cores()
|
| 260 |
+
|
| 261 |
+
# ───────────────────────────────────────────────────────────────────────
|
| 262 |
+
# V6 — Environment variables (user-provided block)
|
| 263 |
+
# ───────────────────────────────────────────────────────────────────────
|
| 264 |
+
os.environ["MKL_ENABLE_INSTRUCTIONS"] = "AVX512"
|
| 265 |
+
NUM_CORES = str(n_phys)
|
| 266 |
+
os.environ["MKL_NUM_THREADS"] = NUM_CORES
|
| 267 |
+
os.environ["OMP_NUM_THREADS"] = NUM_CORES
|
| 268 |
+
os.environ["MKL_DYNAMIC"] = "FALSE"
|
| 269 |
+
os.environ["DNNL_PRIMITIVE_CACHE_CAPACITY"] = "1024"
|
| 270 |
+
os.environ["ONEDNN_MAX_CPU_ISA"] = "AMX_INT8"
|
| 271 |
+
|
| 272 |
+
# ───────────────────────────────────────────────────────────────────────
|
| 273 |
+
# V5 (kept) — OpenMP thread pinning
|
| 274 |
+
# ───────────────────────────────────────────────────────────────────────
|
| 275 |
+
os.environ.setdefault("KMP_AFFINITY", "granularity=fine,compact,1,0")
|
| 276 |
+
os.environ.setdefault("KMP_BLOCKTIME", "1")
|
| 277 |
+
os.environ.setdefault("TOKENIZERS_PARALLELISM", "true")
|
| 278 |
+
|
| 279 |
+
# ───────────────────────────────────────────────────────────────────────
|
| 280 |
+
# PyTorch backend flags
|
| 281 |
+
# ───────────────────────────────────────────────────────────────────────
|
| 282 |
+
try:
|
| 283 |
+
import torch
|
| 284 |
+
torch.set_num_threads(n_phys)
|
| 285 |
+
try:
|
| 286 |
+
torch.set_num_interop_threads(1)
|
| 287 |
+
except RuntimeError:
|
| 288 |
+
# V6.5: já inicializado (e.g., bigru_t package importou torch antes).
|
| 289 |
+
# Silenciosamente ignora — o paralelismo já está configurado.
|
| 290 |
+
pass
|
| 291 |
+
if hasattr(torch.backends, "mkldnn"):
|
| 292 |
+
torch.backends.mkldnn.enabled = True
|
| 293 |
+
if hasattr(torch.backends, "quantized"):
|
| 294 |
+
try:
|
| 295 |
+
torch.backends.quantized.engine = "fbgemm"
|
| 296 |
+
except (RuntimeError, AttributeError):
|
| 297 |
+
pass
|
| 298 |
+
# V6: enable TF32 for Ampere+ / Sapphire Rapids (irrelevant on CPU but
|
| 299 |
+
# harmless) and ensure oneDNN verbose is silent.
|
| 300 |
+
try:
|
| 301 |
+
torch.backends.cuda.matmul.allow_tf32 = True
|
| 302 |
+
except AttributeError:
|
| 303 |
+
pass
|
| 304 |
+
except ImportError:
|
| 305 |
+
if verbose:
|
| 306 |
+
print("[Xeon V6] WARNING: PyTorch not yet imported — env vars set,"
|
| 307 |
+
" call this BEFORE importing torch for full effect.")
|
| 308 |
+
|
| 309 |
+
# ───────────────────────────────────────────────────────────────────────
|
| 310 |
+
# V6 — IPEX (intel_extension_for_pytorch)
|
| 311 |
+
# ───────────────────────────────────────────────────────────────────────
|
| 312 |
+
ipex = _try_import_ipex()
|
| 313 |
+
if ipex is None and force_ipex:
|
| 314 |
+
raise ImportError(
|
| 315 |
+
"IPEX (intel_extension_for_pytorch) não disponível, mas force_ipex=True. "
|
| 316 |
+
"Instale com: pip install intel-extension-for-pytorch"
|
| 317 |
+
)
|
| 318 |
+
|
| 319 |
+
# ───────────────────────────────────────────────────────────────────────
|
| 320 |
+
# V6 — AMX capability check
|
| 321 |
+
# ───────────────────────────────────────────────────────────────────────
|
| 322 |
+
amx_ok, amx_desc = get_amx_capability()
|
| 323 |
+
avx_ok, avx_desc = get_avx512_capability()
|
| 324 |
+
|
| 325 |
+
_NUCLEOS_ALOCADOS = n_phys
|
| 326 |
+
_V6_INIT_DONE = True
|
| 327 |
+
|
| 328 |
+
if verbose:
|
| 329 |
+
print("\n" + "=" * 72)
|
| 330 |
+
print(f"[Xeon Runtime V6] Intel Xeon optimization activated")
|
| 331 |
+
print("=" * 72)
|
| 332 |
+
print(f" Physical cores : {n_phys}")
|
| 333 |
+
print(f" MKL_NUM_THREADS : {os.environ['MKL_NUM_THREADS']}")
|
| 334 |
+
print(f" OMP_NUM_THREADS : {os.environ['OMP_NUM_THREADS']}")
|
| 335 |
+
print(f" MKL_DYNAMIC : {os.environ['MKL_DYNAMIC']}")
|
| 336 |
+
print(f" MKL_ENABLE_INSTRUCTIONS: {os.environ['MKL_ENABLE_INSTRUCTIONS']}")
|
| 337 |
+
print(f" ONEDNN_MAX_CPU_ISA : {os.environ['ONEDNN_MAX_CPU_ISA']}")
|
| 338 |
+
print(f" DNNL_PRIMITIVE_CACHE : {os.environ['DNNL_PRIMITIVE_CACHE_CAPACITY']}")
|
| 339 |
+
print(f" KMP_AFFINITY : {os.environ['KMP_AFFINITY']}")
|
| 340 |
+
print(f" KMP_BLOCKTIME : {os.environ['KMP_BLOCKTIME']}")
|
| 341 |
+
print(f" AVX512 : {avx_desc}")
|
| 342 |
+
print(f" AMX : {amx_desc}")
|
| 343 |
+
print(f" IPEX : "
|
| 344 |
+
f"{'available' if ipex is not None else 'not installed (using native oneDNN)'}")
|
| 345 |
+
print("=" * 72 + "\n")
|
| 346 |
+
|
| 347 |
+
return n_phys
|
| 348 |
+
|
| 349 |
+
|
| 350 |
+
# ============================================================================
|
| 351 |
+
# V6 — apply_ipex_optimization (model-level)
|
| 352 |
+
# ============================================================================
|
| 353 |
+
|
| 354 |
+
def apply_ipex_optimization(model, dtype=None, optimizer=None):
|
| 355 |
+
"""V6: apply ipex.optimize() to a model.
|
| 356 |
+
|
| 357 |
+
Returns (model, optimizer) tuple. If IPEX is not available, returns the
|
| 358 |
+
inputs unchanged. The model is modified in-place when IPEX is present.
|
| 359 |
+
|
| 360 |
+
Args:
|
| 361 |
+
model: torch.nn.Module
|
| 362 |
+
dtype: optional torch.dtype for the model (e.g. torch.bfloat16)
|
| 363 |
+
optimizer: optional torch.optim.Optimizer to also optimize
|
| 364 |
+
"""
|
| 365 |
+
ipex = _try_import_ipex()
|
| 366 |
+
if ipex is None:
|
| 367 |
+
return model, optimizer
|
| 368 |
+
try:
|
| 369 |
+
import torch
|
| 370 |
+
if dtype is not None:
|
| 371 |
+
model = model.to(dtype)
|
| 372 |
+
if optimizer is not None:
|
| 373 |
+
model, optimizer = ipex.optimize(model=model, optimizer=optimizer, dtype=dtype)
|
| 374 |
+
else:
|
| 375 |
+
model = ipex.optimize(model=model, dtype=dtype)
|
| 376 |
+
logger.info(f"[Xeon V6] ipex.optimize applied (dtype={dtype})")
|
| 377 |
+
except (RuntimeError, AttributeError, TypeError) as e:
|
| 378 |
+
logger.warning(f"[Xeon V6] ipex.optimize failed: {e}")
|
| 379 |
+
return model, optimizer
|
| 380 |
+
|
| 381 |
+
|
| 382 |
+
# ============================================================================
|
| 383 |
+
# V6 — FP16 autocast context manager
|
| 384 |
+
# ============================================================================
|
| 385 |
+
|
| 386 |
+
class fp16_autocast:
|
| 387 |
+
"""V6: FP16 CPU autocast context manager.
|
| 388 |
+
|
| 389 |
+
Uses torch.cpu.amp.autocast(dtype=torch.float16) when available. Falls
|
| 390 |
+
back to a no-op context manager if autocast is not supported.
|
| 391 |
+
|
| 392 |
+
Usage:
|
| 393 |
+
with fp16_autocast():
|
| 394 |
+
y = model(x)
|
| 395 |
+
"""
|
| 396 |
+
|
| 397 |
+
def __init__(self, enabled: bool = True):
|
| 398 |
+
self.enabled = enabled
|
| 399 |
+
self._ctx = None
|
| 400 |
+
|
| 401 |
+
def __enter__(self):
|
| 402 |
+
if not self.enabled:
|
| 403 |
+
return self
|
| 404 |
+
try:
|
| 405 |
+
import torch
|
| 406 |
+
if hasattr(torch.cpu, "amp") and hasattr(torch.cpu.amp, "autocast"):
|
| 407 |
+
self._ctx = torch.cpu.amp.autocast(dtype=torch.float16)
|
| 408 |
+
self._ctx.__enter__()
|
| 409 |
+
except (ImportError, RuntimeError, AttributeError):
|
| 410 |
+
self._ctx = None
|
| 411 |
+
return self
|
| 412 |
+
|
| 413 |
+
def __exit__(self, exc_type, exc_val, exc_tb):
|
| 414 |
+
if self._ctx is not None:
|
| 415 |
+
self._ctx.__exit__(exc_type, exc_val, exc_tb)
|
| 416 |
+
self._ctx = None
|
| 417 |
+
return False
|
| 418 |
+
|
| 419 |
+
|
| 420 |
+
# ============================================================================
|
| 421 |
+
# V6 — libvirt AMX activation helper (user-provided code, integrated)
|
| 422 |
+
# ============================================================================
|
| 423 |
+
|
| 424 |
+
def ativar_amx_na_vm(nome_vm: str, qemu_uri: str = "qemu:///system") -> bool:
|
| 425 |
+
"""V6: ativa AMX (amx-tile, amx-int8, amx-bf16) numa VM via libvirt.
|
| 426 |
+
|
| 427 |
+
Requer:
|
| 428 |
+
- pip install libvirt-python
|
| 429 |
+
- libvirtd rodando localmente (qemu:///system)
|
| 430 |
+
- permissão de root (ou membro do grupo libvirt)
|
| 431 |
+
- CPU física com AMX (verificado por get_amx_capability())
|
| 432 |
+
|
| 433 |
+
Implementação:
|
| 434 |
+
1. Conecta ao daemon libvirt
|
| 435 |
+
2. Busca a VM pelo nome
|
| 436 |
+
3. Lê o XML persistente (flag VIR_DOMAIN_XML_INACTIVE = 2)
|
| 437 |
+
4. Garante <cpu mode='host-passthrough'>
|
| 438 |
+
5. Adiciona <feature policy='require' name='amx-tile|int8|bf16'/>
|
| 439 |
+
6. Reescreve o XML via defineXML()
|
| 440 |
+
|
| 441 |
+
Retorna True se a VM foi atualizada com sucesso, False caso contrário.
|
| 442 |
+
Requer reboot da VM para aplicar.
|
| 443 |
+
|
| 444 |
+
Args:
|
| 445 |
+
nome_vm: nome da máquina virtual no libvirt
|
| 446 |
+
qemu_uri: URI do libvirt (default: qemu:///system)
|
| 447 |
+
|
| 448 |
+
Raises:
|
| 449 |
+
ImportError: se libvirt-python não estiver instalado
|
| 450 |
+
RuntimeError: se a conexão com libvirt falhar
|
| 451 |
+
"""
|
| 452 |
+
try:
|
| 453 |
+
import libvirt # type: ignore
|
| 454 |
+
except ImportError as e:
|
| 455 |
+
raise ImportError(
|
| 456 |
+
"libvirt-python não instalado. Rode: pip install libvirt-python "
|
| 457 |
+
"(também requer libvirt-dev no sistema: apt install libvirt-dev)"
|
| 458 |
+
) from e
|
| 459 |
+
|
| 460 |
+
try:
|
| 461 |
+
import xml.etree.ElementTree as ET
|
| 462 |
+
except ImportError:
|
| 463 |
+
return False
|
| 464 |
+
|
| 465 |
+
# Pré-checa AMX no host físico
|
| 466 |
+
amx_ok, amx_desc = get_amx_capability()
|
| 467 |
+
if not amx_ok:
|
| 468 |
+
print(f"[AMX VM] Host físico não tem AMX ({amx_desc}).")
|
| 469 |
+
print(" Não adianta ativar AMX na VM — o host precisa suportar.")
|
| 470 |
+
return False
|
| 471 |
+
|
| 472 |
+
try:
|
| 473 |
+
conn = libvirt.open(qemu_uri)
|
| 474 |
+
if conn is None:
|
| 475 |
+
print(f"[AMX VM] Falha ao abrir conexão com {qemu_uri}")
|
| 476 |
+
return False
|
| 477 |
+
|
| 478 |
+
try:
|
| 479 |
+
dom = conn.lookupByName(nome_vm)
|
| 480 |
+
except libvirt.libvirtError:
|
| 481 |
+
print(f"[AMX VM] VM '{nome_vm}' não encontrada.")
|
| 482 |
+
conn.close()
|
| 483 |
+
return False
|
| 484 |
+
|
| 485 |
+
xml_atual = dom.XMLDesc(2) # VIR_DOMAIN_XML_INACTIVE
|
| 486 |
+
root = ET.fromstring(xml_atual)
|
| 487 |
+
cpu_elem = root.find("cpu")
|
| 488 |
+
|
| 489 |
+
if cpu_elem is None:
|
| 490 |
+
cpu_elem = ET.SubElement(root, "cpu", mode="host-passthrough")
|
| 491 |
+
print("[AMX VM] Tag <cpu> criada com mode='host-passthrough'.")
|
| 492 |
+
else:
|
| 493 |
+
cpu_elem.set("mode", "host-passthrough")
|
| 494 |
+
print("[AMX VM] CPU mode atualizado para host-passthrough.")
|
| 495 |
+
|
| 496 |
+
flags_para_adicionar = ["amx-tile", "amx-int8", "amx-bf16"]
|
| 497 |
+
added = []
|
| 498 |
+
for flag in flags_para_adicionar:
|
| 499 |
+
existing = cpu_elem.find(f"./feature[@name='{flag}']")
|
| 500 |
+
if existing is None:
|
| 501 |
+
ET.SubElement(cpu_elem, "feature", policy="require", name=flag)
|
| 502 |
+
added.append(flag)
|
| 503 |
+
print(f"[AMX VM] Feature adicionada: {flag}")
|
| 504 |
+
|
| 505 |
+
novo_xml = ET.tostring(root, encoding="utf-8").decode("utf-8")
|
| 506 |
+
conn.defineXML(novo_xml)
|
| 507 |
+
print(f"[AMX VM] ✓ AMX ativado no XML da VM '{nome_vm}'. Reinicie a VM para aplicar.")
|
| 508 |
+
conn.close()
|
| 509 |
+
return True
|
| 510 |
+
|
| 511 |
+
except libvirt.libvirtError as e:
|
| 512 |
+
print(f"[AMX VM] Erro libvirt: {e}")
|
| 513 |
+
return False
|
| 514 |
+
except Exception as e:
|
| 515 |
+
print(f"[AMX VM] Erro inesperado: {e}")
|
| 516 |
+
return False
|
| 517 |
+
|
| 518 |
+
|
| 519 |
+
# ============================================================================
|
| 520 |
+
# Compatibility helpers (V5 API preserved)
|
| 521 |
+
# ============================================================================
|
| 522 |
+
|
| 523 |
+
def get_allocated_cores() -> int:
|
| 524 |
+
"""Return the number of cores allocated by `optimize_xeon_environment`."""
|
| 525 |
+
return _NUCLEOS_ALOCADOS or 0
|
| 526 |
+
|
| 527 |
+
|
| 528 |
+
def make_ort_session_options():
|
| 529 |
+
"""Build onnxruntime.SessionOptions tuned for Xeon AVX512/AMX."""
|
| 530 |
+
import onnxruntime as ort # type: ignore
|
| 531 |
+
|
| 532 |
+
n_cores = _NUCLEOS_ALOCADOS or _detect_physical_cores()
|
| 533 |
+
|
| 534 |
+
opts = ort.SessionOptions()
|
| 535 |
+
opts.intra_op_num_threads = n_cores
|
| 536 |
+
opts.inter_op_num_threads = 1
|
| 537 |
+
opts.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL
|
| 538 |
+
opts.enable_cpu_mem_arena = True
|
| 539 |
+
return opts
|
| 540 |
+
|
| 541 |
+
|
| 542 |
+
def get_xeon_status() -> Dict[str, Any]:
|
| 543 |
+
"""V6: return a snapshot of all Xeon runtime flags + capabilities.
|
| 544 |
+
|
| 545 |
+
Useful for logging inside the training script.
|
| 546 |
+
"""
|
| 547 |
+
amx_ok, amx_desc = get_amx_capability()
|
| 548 |
+
avx_ok, avx_desc = get_avx512_capability()
|
| 549 |
+
return {
|
| 550 |
+
"version": "V6",
|
| 551 |
+
"physical_cores": _NUCLEOS_ALOCADOS or _detect_physical_cores(),
|
| 552 |
+
"env": {
|
| 553 |
+
"MKL_ENABLE_INSTRUCTIONS": os.environ.get("MKL_ENABLE_INSTRUCTIONS"),
|
| 554 |
+
"MKL_NUM_THREADS": os.environ.get("MKL_NUM_THREADS"),
|
| 555 |
+
"OMP_NUM_THREADS": os.environ.get("OMP_NUM_THREADS"),
|
| 556 |
+
"MKL_DYNAMIC": os.environ.get("MKL_DYNAMIC"),
|
| 557 |
+
"DNNL_PRIMITIVE_CACHE_CAPACITY": os.environ.get("DNNL_PRIMITIVE_CACHE_CAPACITY"),
|
| 558 |
+
"ONEDNN_MAX_CPU_ISA": os.environ.get("ONEDNN_MAX_CPU_ISA"),
|
| 559 |
+
"KMP_AFFINITY": os.environ.get("KMP_AFFINITY"),
|
| 560 |
+
"KMP_BLOCKTIME": os.environ.get("KMP_BLOCKTIME"),
|
| 561 |
+
},
|
| 562 |
+
"avx512": {"supported": avx_ok, "desc": avx_desc},
|
| 563 |
+
"amx": {"supported": amx_ok, "desc": amx_desc},
|
| 564 |
+
"ipex_available": _IPEX_AVAILABLE is True,
|
| 565 |
+
"init_done": _V6_INIT_DONE,
|
| 566 |
+
}
|
| 567 |
+
|
| 568 |
+
|
| 569 |
+
__all__ = [
|
| 570 |
+
"optimize_xeon_environment",
|
| 571 |
+
"apply_ipex_optimization",
|
| 572 |
+
"fp16_autocast",
|
| 573 |
+
"benchmark_fp16_matmul",
|
| 574 |
+
"benchmark_int8_matmul",
|
| 575 |
+
"ativar_amx_na_vm",
|
| 576 |
+
"get_avx512_capability",
|
| 577 |
+
"get_amx_capability",
|
| 578 |
+
"get_xeon_status",
|
| 579 |
+
"get_allocated_cores",
|
| 580 |
+
"make_ort_session_options",
|
| 581 |
+
]
|