PowerMachine commited on
Commit
e30765f
·
verified ·
1 Parent(s): 1f3d188

V6.5-V2-dynamic: upload batch (scripts + state + model) [16 files]

Browse files
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": 10680,
7
  "n_errors": 0,
8
- "last_norm_in": 111.9294204711914,
9
- "last_norm_out": 116.00782012939453,
10
  "last_attn_activated": true,
11
- "last_attn_diff_norm": 116.50959014892578,
12
  "n_heads": 8,
13
  "logic_functional": true
14
  },
15
  "active": true,
16
  "logic_functional": true,
17
- "n_calls": 10680,
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-08T23:52:08.862775",
4
  "config": {
5
  "som_grid": [
6
  6,
@@ -9,9 +9,9 @@
9
  4
10
  ],
11
  "n_neurons": 864,
12
- "hidden_dim": 1024,
13
- "vocab_size": 16384,
14
- "n_hypotheses": 16,
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": 91.0796530006337,
58
- "avg_time_ms": 91.46852800040506,
59
- "best_tflops": 1.4053632812930177,
60
- "avg_tflops": 1.3993884322641899,
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": 722.1084952354431,
90
  "storage_critical_stopped": false
91
  },
92
  "fase_2_punicão_summary": {
93
  "total_samples": 2000,
94
  "meta_atingida": true,
95
- "elapsed_s": 507.18996810913086,
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": 10680,
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": 11.263132095336914,
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": 220.174132,
 
 
 
128
  "reason": "end_of_training_v2",
129
  "step": 700,
130
  "total_samples": 10000,
131
  "n_tensors": 7,
132
- "n_buffer_tail": 128
133
  },
134
  "final_v2_metrics": {
135
  "version": "V2-dynamic",
136
- "n_hypotheses": 16,
137
- "n_hypotheses_active": 16,
138
- "max_n_hypotheses": 32,
139
  "n_trials": 6,
140
  "hyp_train_steps": 30,
141
  "hyp_lr": 0.0001,
142
  "delta_scale": 0.0010000000474974513,
143
- "n_generators": 32,
144
  "punishment_count": 0,
145
  "success_count": 0,
146
  "training_ready": false,
147
  "classifier_trained": true,
148
- "ewc_reference_set": true,
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": 32,
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": 16,
176
  "n_trials": 6,
177
  "hyp_train_steps": 30
178
  },
179
  "after": {
180
- "n_hypotheses": 16,
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": 16,
200
  "n_trials": 6,
201
  "hyp_train_steps": 30
202
  },
203
  "after": {
204
- "n_hypotheses": 16,
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": 16,
224
  "n_trials": 6,
225
  "hyp_train_steps": 30
226
  },
227
  "after": {
228
- "n_hypotheses": 16,
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": 16,
248
  "n_trials": 6,
249
  "hyp_train_steps": 30
250
  },
251
  "after": {
252
- "n_hypotheses": 16,
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": 16,
272
  "n_trials": 6,
273
  "hyp_train_steps": 30
274
  },
275
  "after": {
276
- "n_hypotheses": 16,
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=(0, 0, 0, 0)",
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": 9.467601776123047
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=(0, 0, 0, 0)",
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": 11.906862258911133
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=(0, 0, 0, 0)",
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": 12.414932250976562
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": 11.263132095336914,
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
+ ]