PowerMachine commited on
Commit
0cee009
·
verified ·
1 Parent(s): a22ed76

V3.2: corrigir 6 bugs + meta_T_moved (T_tracking_loss) + re-teste com Xeon IPEX 2.8 + AMX

Browse files
src/bigru_t/model/u8cell_t.py CHANGED
@@ -12,11 +12,33 @@ Reaproveita a arquitetura "4 camadas BiGRU cooperativas" do v13.9.2, mas:
12
  - TransformerUnit funde as 8 saídas (vs. Confidence Gate no v13.9.2)
13
  - Sem Hypothesis Network interna (a hipótese é global, Lema 3)
14
  - Sem Trust Network interna (o orquestrador global faz o papel, Lema 4)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
15
  """
16
  from __future__ import annotations
17
 
18
  import torch
19
  import torch.nn as nn
 
 
20
 
21
  from .bigru4 import BiGRU4
22
  from .transformer_unit import TransformerUnit
@@ -33,8 +55,11 @@ class u8cell_T(nn.Module):
33
  d_ff: dimensão interna do FFN do TransformerUnit
34
  output_dim: dimensão de saída do u8cell_T
35
 
36
- Forward:
37
  x: (batch, T, input_dim) → h_out: (batch, output_dim)
 
 
 
38
  """
39
 
40
  def __init__(
@@ -50,10 +75,11 @@ class u8cell_T(nn.Module):
50
  super().__init__()
51
  self.input_dim = input_dim
52
  self.output_dim = output_dim
 
53
 
54
  # 8 BiGRU4 paralelas (cada uma com 4 camadas BiGRU)
55
  self.bigru_list = nn.ModuleList(
56
- [BiGRU4(input_dim, hidden_dim_bigru, dropout=dropout) for _ in range(8)]
57
  )
58
  bigru_out_dim = 2 * hidden_dim_bigru # bidirecional
59
 
@@ -66,12 +92,23 @@ class u8cell_T(nn.Module):
66
  # Projeção final para output_dim
67
  self.proj_out = nn.Linear(d_transformer, output_dim)
68
 
69
- def forward(self, x: torch.Tensor) -> torch.Tensor:
70
- """x: (batch, T, input_dim) → h_out: (batch, output_dim)"""
71
- representations = []
 
 
 
 
 
 
 
 
 
 
72
  for bigru in self.bigru_list:
73
  v_i = bigru(x) # (batch, bigru_out_dim)
74
  v_i_proj = self.proj_in(v_i) # (batch, d_transformer)
 
75
  representations.append(v_i_proj)
76
 
77
  # Empilha as 8 representações: (batch, 8, d_transformer)
@@ -82,4 +119,91 @@ class u8cell_T(nn.Module):
82
 
83
  # Projeção final → (batch, output_dim)
84
  h_out = self.proj_out(h_unit)
85
- return h_out
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
12
  - TransformerUnit funde as 8 saídas (vs. Confidence Gate no v13.9.2)
13
  - Sem Hypothesis Network interna (a hipótese é global, Lema 3)
14
  - Sem Trust Network interna (o orquestrador global faz o papel, Lema 4)
15
+
16
+ V4 — Aprimoramento da detecção de ativação BiGRU:
17
+ Adiciona `return_aux=True` no forward para retornar métricas de
18
+ paralelismo real e diversificação de dados entre as 8 BiGRU4:
19
+ - `bigru_representations`: lista de (batch, d_transformer) — saídas das 8 BiGRU4
20
+ - `cosine_diversity`: 1 - média(cos_sim entre pares) ∈ [0, 2]
21
+ Alto (próximo de 1) = módulos produzindo representações diversas
22
+ Baixo (próximo de 0) = módulos colapsados (mesma saída)
23
+ - `gate_norm_ratio`: std(||v_i||) / mean(||v_i||) — coef. variação das normas
24
+ Alto = módulos com magnitudes diferentes (especialização de escala)
25
+ Baixo = módulos com mesma magnitude (sem diferenciação)
26
+ - `activation_variance`: variância inter-BiGRU por sample (média sobre batch)
27
+ Alto = diversificação de conteúdo
28
+ Baixo = redundância
29
+ - `parallelism_score`: geom_mean(||v_i||) / mean(||v_i||) — score de paralelismo
30
+ Próximo de 1 = todas as 8 BiGRU contribuem ativamente
31
+ Próximo de 0 = apenas algumas dominam
32
+
33
+ Essas métricas permitem monitorar se os 8 BiGRUs estão realmente trabalhando
34
+ em paralelo com dados diversificados (Lema 1) ou se colapsaram.
35
  """
36
  from __future__ import annotations
37
 
38
  import torch
39
  import torch.nn as nn
40
+ import torch.nn.functional as F
41
+ from typing import Optional
42
 
43
  from .bigru4 import BiGRU4
44
  from .transformer_unit import TransformerUnit
 
55
  d_ff: dimensão interna do FFN do TransformerUnit
56
  output_dim: dimensão de saída do u8cell_T
57
 
58
+ Forward (default):
59
  x: (batch, T, input_dim) → h_out: (batch, output_dim)
60
+
61
+ Forward (return_aux=True):
62
+ → (h_out, aux_dict) onde aux_dict contém métricas BiGRU (V4)
63
  """
64
 
65
  def __init__(
 
75
  super().__init__()
76
  self.input_dim = input_dim
77
  self.output_dim = output_dim
78
+ self.n_bigru = 8 # Número de BiGRU4 paralelas (constante arquitetural)
79
 
80
  # 8 BiGRU4 paralelas (cada uma com 4 camadas BiGRU)
81
  self.bigru_list = nn.ModuleList(
82
+ [BiGRU4(input_dim, hidden_dim_bigru, dropout=dropout) for _ in range(self.n_bigru)]
83
  )
84
  bigru_out_dim = 2 * hidden_dim_bigru # bidirecional
85
 
 
92
  # Projeção final para output_dim
93
  self.proj_out = nn.Linear(d_transformer, output_dim)
94
 
95
+ def forward(
96
+ self,
97
+ x: torch.Tensor,
98
+ return_aux: bool = False,
99
+ ) -> torch.Tensor:
100
+ """x: (batch, T, input_dim) → h_out: (batch, output_dim)
101
+
102
+ Se return_aux=True, retorna (h_out, aux_dict) com métricas de
103
+ paralelismo e diversificação das 8 BiGRU4 (V4).
104
+ """
105
+ representations = [] # cada (batch, d_transformer) — saídas projetadas
106
+ bigru_raw = [] # cada (batch, 2*hidden_dim) — saídas brutas das BiGRU4
107
+
108
  for bigru in self.bigru_list:
109
  v_i = bigru(x) # (batch, bigru_out_dim)
110
  v_i_proj = self.proj_in(v_i) # (batch, d_transformer)
111
+ bigru_raw.append(v_i)
112
  representations.append(v_i_proj)
113
 
114
  # Empilha as 8 representações: (batch, 8, d_transformer)
 
119
 
120
  # Projeção final → (batch, output_dim)
121
  h_out = self.proj_out(h_unit)
122
+
123
+ if not return_aux:
124
+ return h_out
125
+
126
+ # =====================================================================
127
+ # V4: Métricas de paralelismo e diversificação das 8 BiGRU4
128
+ # =====================================================================
129
+ aux = self._compute_bigru_diversity_metrics(representations, bigru_raw)
130
+ return h_out, aux
131
+
132
+ def _compute_bigru_diversity_metrics(
133
+ self,
134
+ representations: list[torch.Tensor],
135
+ bigru_raw: list[torch.Tensor],
136
+ ) -> dict:
137
+ """Computa métricas de paralelismo/diversificação das 8 BiGRU4.
138
+
139
+ Args:
140
+ representations: lista de 8 tensores (batch, d_transformer) — saídas projetadas
141
+ bigru_raw: lista de 8 tensores (batch, 2*hidden_dim) — saídas brutas
142
+
143
+ Returns:
144
+ dict com:
145
+ - cosine_diversity: 1 - mean(cos_sim entre pares) ∈ [0, 2]
146
+ - gate_norm_ratio: std(||v_i||) / mean(||v_i||)
147
+ - activation_variance: variância inter-BiGRU (média sobre batch)
148
+ - parallelism_score: geom_mean(||v_i||) / mean(||v_i||)
149
+ - per_bigru_norms: lista de 8 normas L2 (uma por BiGRU4)
150
+ - bigru_representations: stacked (batch, 8, d_transformer) — para análise externa
151
+ """
152
+ # Stack para (8, batch, d_transformer) → (8, batch)
153
+ norms = torch.stack(
154
+ [v.norm(p=2, dim=-1) for v in representations], dim=0
155
+ ) # (8, batch)
156
+ per_bigru_norms = norms.mean(dim=1) # (8,) — norma média de cada BiGRU4
157
+
158
+ # Stack das representações projetadas: (batch, 8, d_transformer)
159
+ stacked_reps = torch.stack(representations, dim=1)
160
+
161
+ # 1. Cosine diversity: 1 - mean(cos_sim entre pares i<j)
162
+ # representations[i]: (batch, d_transformer)
163
+ # cosine sim(i, j) médio sobre batch
164
+ n = len(representations)
165
+ cos_sims = []
166
+ for i in range(n):
167
+ for j in range(i + 1, n):
168
+ # (batch,) — cosine sim por sample, depois média sobre batch
169
+ cs = F.cosine_similarity(
170
+ representations[i], representations[j], dim=-1
171
+ ) # (batch,)
172
+ cos_sims.append(cs.mean())
173
+ if cos_sims:
174
+ mean_cos = torch.stack(cos_sims).mean()
175
+ cosine_diversity = 1.0 - mean_cos
176
+ else:
177
+ cosine_diversity = torch.tensor(0.0, device=representations[0].device)
178
+
179
+ # 2. Gate norm ratio: coef. de variação das normas L2 das 8 BiGRU4
180
+ # Médio sobre batch: usamos per_bigru_norms (já é média sobre batch)
181
+ mean_norm = per_bigru_norms.mean()
182
+ std_norm = per_bigru_norms.std()
183
+ gate_norm_ratio = (std_norm / (mean_norm + 1e-8)).clamp(min=0.0)
184
+
185
+ # 3. Activation variance: variância inter-BiGRU por sample, média sobre batch
186
+ # stacked_raw: (batch, 8, bigru_out_dim)
187
+ stacked_raw = torch.stack(bigru_raw, dim=1) # (batch, 8, bigru_out_dim)
188
+ # variância sobre dim=1 (inter-BiGRU), depois média sobre batch e features
189
+ activation_variance = stacked_raw.var(dim=1).mean()
190
+
191
+ # 4. Parallelism score: geom_mean(||v_i||) / mean(||v_i||) — score ∈ (0, 1]
192
+ # geom_mean próximo de mean → 1.0 (todas contribuem)
193
+ # geom_mean muito menor que mean → algumas BiGRUs dominam
194
+ log_norms = torch.log(per_bigru_norms + 1e-8)
195
+ log_geom_mean = log_norms.mean()
196
+ geom_mean = torch.exp(log_geom_mean)
197
+ parallelism_score = (geom_mean / (mean_norm + 1e-8)).clamp(min=0.0, max=1.0)
198
+
199
+ return {
200
+ "cosine_diversity": float(cosine_diversity.item()),
201
+ "gate_norm_ratio": float(gate_norm_ratio.item()),
202
+ "activation_variance": float(activation_variance.item()),
203
+ "parallelism_score": float(parallelism_score.item()),
204
+ "per_bigru_norms": [float(x.item()) for x in per_bigru_norms],
205
+ # V4: NÃO retornar bigru_representations (consome RAM nos K módulos)
206
+ # A diversidade INTER-módulos é computada no UnifiedModel usando
207
+ # h_k (saída do u8cell_T), não as representações internas.
208
+ # "bigru_representations": stacked_reps.detach(), # comentado para economizar RAM
209
+ }
src/bigru_t/model/unified_model.py CHANGED
@@ -106,10 +106,11 @@ class UnifiedModelConfig:
106
  num_layers_train: int = 2
107
 
108
  # HypT
 
109
  hypT_dim: int = 128
110
  nhead_hyp: int = 4
111
  d_ff_hyp: int = 256
112
- num_layers_hyp: int = 2
113
 
114
  # Quantization (Lema 3)
115
  num_bits: int = 8
@@ -310,12 +311,23 @@ class UnifiedModel(nn.Module):
310
  alpha, entropy_reg = self.module_selector(temperature=temperature)
311
 
312
  # Calcula saídas de todos os módulos e pondera por alpha
 
313
  H_list = []
 
 
314
  for k, cell in enumerate(self.u8cells):
315
- h_k = cell(x) # (batch, output_dim_u8cell)
 
 
 
 
 
 
316
  h_k_weighted = h_k * alpha[k]
317
  H_list.append(h_k_weighted)
318
  H = torch.stack(H_list, dim=1) # (batch, max_modules, output_dim_u8cell)
 
 
319
 
320
  # Orquestrador: produz representação agregada
321
  o = self.orq_cell(H) # (batch, d_cache)
@@ -352,15 +364,103 @@ class UnifiedModel(nn.Module):
352
  delta = torch.zeros_like(y_hat_main)
353
 
354
  if return_aux:
 
 
 
 
355
  aux = {
356
  "entropy_reg": entropy_reg,
357
  "alpha": alpha.detach(),
358
  "o": o.detach(),
359
  "circular_info": circular_info,
 
360
  }
361
  return y_hat_main, delta, aux
362
  return y_hat_main, delta
363
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
364
  def get_active_modules(self, threshold: float = 0.01) -> int:
365
  """Retorna nº efetivo de módulos ativos (para auto-configuração)."""
366
  return self.module_selector.get_active_modules(threshold)
 
106
  num_layers_train: int = 2
107
 
108
  # HypT
109
+ # V4: num_layers_hyp aumentado de 2 → 8 (mais capacidade para correção delta)
110
  hypT_dim: int = 128
111
  nhead_hyp: int = 4
112
  d_ff_hyp: int = 256
113
+ num_layers_hyp: int = 8
114
 
115
  # Quantization (Lema 3)
116
  num_bits: int = 8
 
311
  alpha, entropy_reg = self.module_selector(temperature=temperature)
312
 
313
  # Calcula saídas de todos os módulos e pondera por alpha
314
+ # V4: coleta métricas de diversificação BiGRU (paralelismo real)
315
  H_list = []
316
+ H_unweighted = [] # V4: saídas h_k SEM ponderar, para diversidade INTER-módulos
317
+ bigru_metrics_per_module = [] # lista de dicts (1 por u8cell_T)
318
  for k, cell in enumerate(self.u8cells):
319
+ if return_aux:
320
+ # V4: forward com return_aux para capturar diversidade BiGRU
321
+ h_k, bigru_aux = cell(x, return_aux=True)
322
+ bigru_metrics_per_module.append(bigru_aux)
323
+ else:
324
+ h_k = cell(x) # (batch, output_dim_u8cell)
325
+ H_unweighted.append(h_k)
326
  h_k_weighted = h_k * alpha[k]
327
  H_list.append(h_k_weighted)
328
  H = torch.stack(H_list, dim=1) # (batch, max_modules, output_dim_u8cell)
329
+ # V4: stack das saídas não-ponderadas para diversidade INTER-módulos
330
+ H_unweighted_stack = torch.stack(H_unweighted, dim=1) # (batch, K, output_dim_u8cell)
331
 
332
  # Orquestrador: produz representação agregada
333
  o = self.orq_cell(H) # (batch, d_cache)
 
364
  delta = torch.zeros_like(y_hat_main)
365
 
366
  if return_aux:
367
+ # V4: agrega métricas BiGRU de todos os K u8cell_T módulos
368
+ bigru_diversity_summary = self._aggregate_bigru_diversity(
369
+ bigru_metrics_per_module, alpha, H_unweighted_stack
370
+ )
371
  aux = {
372
  "entropy_reg": entropy_reg,
373
  "alpha": alpha.detach(),
374
  "o": o.detach(),
375
  "circular_info": circular_info,
376
+ "bigru_diversity": bigru_diversity_summary, # V4
377
  }
378
  return y_hat_main, delta, aux
379
  return y_hat_main, delta
380
 
381
+ def _aggregate_bigru_diversity(
382
+ self,
383
+ metrics_per_module: list[dict],
384
+ alpha: torch.Tensor,
385
+ H_unweighted: torch.Tensor,
386
+ ) -> dict:
387
+ """Agrega métricas BiGRU de todos os K u8cell_T módulos.
388
+
389
+ Args:
390
+ metrics_per_module: lista de K dicts (cada um = saída de u8cell_T._compute_bigru_diversity_metrics)
391
+ alpha: (K,) — pesos do ModuleSelector
392
+ H_unweighted: (batch, K, output_dim_u8cell) — saídas h_k SEM ponderar
393
+
394
+ Returns:
395
+ dict com médias agregadas sobre os K módulos (ponderadas por alpha):
396
+ - mean_cosine_diversity: diversidade média inter-BiGRU (ponderada por alpha)
397
+ - mean_gate_norm_ratio: variação média das normas
398
+ - mean_activation_variance: variância média
399
+ - mean_parallelism_score: score médio de paralelismo
400
+ - module_level_cosine_diversity: 1 - mean(cos_sim entre h_k dos K módulos)
401
+ (diversidade INTER-módulos, não intra-módulo)
402
+ - per_module_summary: lista de K dicts com métricas essenciais
403
+ """
404
+ if not metrics_per_module:
405
+ return {
406
+ "mean_cosine_diversity": 0.0,
407
+ "mean_gate_norm_ratio": 0.0,
408
+ "mean_activation_variance": 0.0,
409
+ "mean_parallelism_score": 0.0,
410
+ "module_level_cosine_diversity": 0.0,
411
+ "per_module_summary": [],
412
+ }
413
+
414
+ # Ponderar por alpha_k (módulos mais ativos pesam mais)
415
+ alpha_np = alpha.detach().cpu().numpy()
416
+ alpha_norm = alpha_np / (alpha_np.sum() + 1e-8)
417
+
418
+ cd_list = [m["cosine_diversity"] for m in metrics_per_module]
419
+ gnr_list = [m["gate_norm_ratio"] for m in metrics_per_module]
420
+ av_list = [m["activation_variance"] for m in metrics_per_module]
421
+ ps_list = [m["parallelism_score"] for m in metrics_per_module]
422
+
423
+ mean_cd = sum(w * v for w, v in zip(alpha_norm, cd_list))
424
+ mean_gnr = sum(w * v for w, v in zip(alpha_norm, gnr_list))
425
+ mean_av = sum(w * v for w, v in zip(alpha_norm, av_list))
426
+ mean_ps = sum(w * v for w, v in zip(alpha_norm, ps_list))
427
+
428
+ # V4: Diversidade INTER-módulos usando H_unweighted (batch, K, output_dim_u8cell)
429
+ # Média sobre batch → (K, output_dim_u8cell), depois cosine sim entre pares
430
+ if H_unweighted is not None and H_unweighted.size(1) >= 2:
431
+ module_reps = H_unweighted.detach().mean(dim=0) # (K, output_dim_u8cell)
432
+ K = module_reps.size(0)
433
+ inter_cos_sims = []
434
+ for i in range(K):
435
+ for j in range(i + 1, K):
436
+ cs = F.cosine_similarity(
437
+ module_reps[i].unsqueeze(0),
438
+ module_reps[j].unsqueeze(0),
439
+ dim=-1,
440
+ )
441
+ inter_cos_sims.append(float(cs.item()))
442
+ mean_inter_cos = sum(inter_cos_sims) / len(inter_cos_sims)
443
+ module_level_cd = 1.0 - mean_inter_cos
444
+ else:
445
+ module_level_cd = 0.0
446
+
447
+ return {
448
+ "mean_cosine_diversity": float(mean_cd),
449
+ "mean_gate_norm_ratio": float(mean_gnr),
450
+ "mean_activation_variance": float(mean_av),
451
+ "mean_parallelism_score": float(mean_ps),
452
+ "module_level_cosine_diversity": float(module_level_cd),
453
+ "per_module_summary": [
454
+ {
455
+ "cosine_diversity": cd_list[i],
456
+ "gate_norm_ratio": gnr_list[i],
457
+ "parallelism_score": ps_list[i],
458
+ "per_bigru_norms": metrics_per_module[i]["per_bigru_norms"],
459
+ }
460
+ for i in range(len(metrics_per_module))
461
+ ],
462
+ }
463
+
464
  def get_active_modules(self, threshold: float = 0.01) -> int:
465
  """Retorna nº efetivo de módulos ativos (para auto-configuração)."""
466
  return self.module_selector.get_active_modules(threshold)
src/bigru_t/training/__init__.py CHANGED
@@ -3,10 +3,12 @@ from .gradient_surgery import apply_gradient_surgery, orthogonalize_gradient
3
  from .meta_configurator import MetaConfigurator
4
  from .kill_switch import KillSwitch, KillSwitchState
5
  from .trainer import BiGRU_T_Trainer, TrainerConfig
 
6
 
7
  __all__ = [
8
  "apply_gradient_surgery", "orthogonalize_gradient",
9
  "MetaConfigurator",
10
  "KillSwitch", "KillSwitchState",
11
  "BiGRU_T_Trainer", "TrainerConfig",
 
12
  ]
 
3
  from .meta_configurator import MetaConfigurator
4
  from .kill_switch import KillSwitch, KillSwitchState
5
  from .trainer import BiGRU_T_Trainer, TrainerConfig
6
+ from .triplet_prototype_loss import TripletPrototypeLoss # V4 — compatibilizado de gru-ring-v13-9-2
7
 
8
  __all__ = [
9
  "apply_gradient_surgery", "orthogonalize_gradient",
10
  "MetaConfigurator",
11
  "KillSwitch", "KillSwitchState",
12
  "BiGRU_T_Trainer", "TrainerConfig",
13
+ "TripletPrototypeLoss", # V4
14
  ]
src/bigru_t/training/triplet_prototype_loss.py ADDED
@@ -0,0 +1,125 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """triplet_prototype_loss.py — Triplet loss com protótipos (vetorizado V4).
2
+
3
+ Origem: compatibilizado matematicamente de gru-ring-v13-9-2/flexnet/triplet_prototype_loss.py
4
+ (V11.23 original do PowerMachine/gru-ring-v13-9-2 no HF).
5
+
6
+ Melhorias V4 (compatibilização matemática para BiGRU_T_version):
7
+ 1. **Vetorização total**: o loop O(N*C) original é substituído por operações
8
+ matriciais. Para N=64 samples e C=8 protótipos, isso reduz de 512 iterações
9
+ para 1 única matriz de distâncias (N, C).
10
+ 2. **Detached prototypes**: protótipos são tratados como referências (não
11
+ recebem gradiente) — isso evita o colapso trivial onde todos os protótipos
12
+ convergem para o mesmo ponto. O gradiente flui apenas para os embeddings.
13
+ 3. **Integração com ModuleSelector (Lema 1)**: os protótipos podem ser as
14
+ próprias saídas agregadas dos u8cell_T módulos (1 protótipo por módulo),
15
+ forçando diversificação INTER-módulos via loss triplet.
16
+ 4. **Hard negative mining**: seleção vetorial do protótipo negativo mais
17
+ próximo (menor distância) entre classes diferentes, sem loop Python.
18
+ 5. **Compatibilidade matemática**: a fórmula mantém a semântica do original:
19
+ loss = mean(max(d_pos - d_neg + margin, 0))
20
+ onde d_pos = ||emb_i - proto[label_i]||_2 e d_neg = min_{c≠label_i} ||emb_i - proto[c]||_2
21
+
22
+ Uso típico no BiGRU_T_version:
23
+ - embeddings: (B, D) — saída agregada do OrqCell (d_cache)
24
+ - labels: (B,) — índices do módulo ativo (argmax do ModuleSelector) ou classes reais
25
+ - prototypes: (K, D) — uma protótipo por u8cell_T (média móvel das ativações)
26
+ """
27
+ from __future__ import annotations
28
+
29
+ import torch
30
+ import torch.nn as nn
31
+ import torch.nn.functional as F
32
+
33
+
34
+ class TripletPrototypeLoss(nn.Module):
35
+ """Triplet loss com protótipos (vetorizado V4).
36
+
37
+ Args:
38
+ margin: margem do triplet (default 0.3)
39
+ lambda_triplet: peso da loss (default 0.1)
40
+ prototypes_detach: se True (default), protótipos não recebem gradiente
41
+
42
+ Forward:
43
+ embeddings: (N, D)
44
+ labels: (N,) — índices inteiros em [0, C)
45
+ prototypes: (C, D)
46
+
47
+ Returns:
48
+ escalar (loss triplet média sobre N samples)
49
+ """
50
+
51
+ def __init__(self, margin: float = 0.3, lambda_triplet: float = 0.1,
52
+ prototypes_detach: bool = True):
53
+ super().__init__()
54
+ self.margin = margin
55
+ self.lambda_triplet = lambda_triplet
56
+ self.prototypes_detach = prototypes_detach
57
+
58
+ def forward(
59
+ self,
60
+ embeddings: torch.Tensor,
61
+ labels: torch.Tensor,
62
+ prototypes: torch.Tensor,
63
+ ) -> torch.Tensor:
64
+ """Computa triplet loss vetorizada.
65
+
66
+ Args:
67
+ embeddings: (N, D)
68
+ labels: (N,) — índices em [0, C)
69
+ prototypes: (C, D)
70
+
71
+ Returns:
72
+ scalar tensor (loss média)
73
+ """
74
+ N, D = embeddings.shape
75
+ C = prototypes.shape[0]
76
+
77
+ if N == 0 or C == 0:
78
+ return torch.tensor(0.0, device=embeddings.device, requires_grad=True)
79
+
80
+ # Detach protótipos (não receber gradiente — referência fixa por step)
81
+ if self.prototypes_detach:
82
+ prototypes = prototypes.detach()
83
+
84
+ # =====================================================================
85
+ # 1. Matriz de distâncias (N, C) — pairwise L2 entre embeddings e protótipos
86
+ # =====================================================================
87
+ # ||a - b||² = ||a||² + ||b||² - 2 * a·b
88
+ emb_sq = (embeddings ** 2).sum(dim=-1, keepdim=True) # (N, 1)
89
+ proto_sq = (prototypes ** 2).sum(dim=-1, keepdim=True).t() # (1, C)
90
+ dot = embeddings @ prototypes.t() # (N, C)
91
+ dist_sq = emb_sq + proto_sq - 2.0 * dot
92
+ dist_sq = dist_sq.clamp(min=0.0) # estabilidade numérica
93
+ dist = torch.sqrt(dist_sq + 1e-8) # (N, C)
94
+
95
+ # =====================================================================
96
+ # 2. Distância positiva: d_pos[i] = dist[i, labels[i]]
97
+ # =====================================================================
98
+ # Usar gather para selecionar a coluna correta por linha
99
+ labels_clamped = labels.clamp(min=0, max=C - 1).long() # (N,)
100
+ d_pos = dist.gather(1, labels_clamped.unsqueeze(1)).squeeze(1) # (N,)
101
+
102
+ # =====================================================================
103
+ # 3. Hard negative: d_neg[i] = min_{c != labels[i]} dist[i, c]
104
+ # =====================================================================
105
+ # Máscara (N, C) — True onde c != labels[i]
106
+ mask_neg = torch.ones_like(dist, dtype=torch.bool)
107
+ mask_neg.scatter_(1, labels_clamped.unsqueeze(1), False)
108
+
109
+ # Onde mask_neg é False, setar para +inf para que o min não pegue a classe positiva
110
+ dist_neg = dist.masked_fill(~mask_neg, float('inf'))
111
+ d_neg = dist_neg.min(dim=1).values # (N,)
112
+
113
+ # =====================================================================
114
+ # 4. Triplet loss: max(d_pos - d_neg + margin, 0) — média sobre N
115
+ # =====================================================================
116
+ # Se algum sample não tem negative válido (C=1), d_neg = inf → triplet = 0
117
+ triplet = F.relu(d_pos - d_neg + self.margin)
118
+ # Filtrar samples sem negative válido (d_neg = inf → triplet = nan se subtrair)
119
+ valid_mask = torch.isfinite(d_neg)
120
+ if valid_mask.any():
121
+ loss = triplet[valid_mask].mean()
122
+ else:
123
+ loss = torch.tensor(0.0, device=embeddings.device, requires_grad=True)
124
+
125
+ return self.lambda_triplet * loss
v32_upload_report.json ADDED
@@ -0,0 +1,48 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "version": "V3.2_upload",
3
+ "timestamp": 1786013028.9061728,
4
+ "repo_id": "PowerMachine/BiGRU_T_version",
5
+ "repo_type": "model",
6
+ "folder_path": "/home/z/my-project/BiGRU_T_version",
7
+ "allow_patterns": [
8
+ "*.py",
9
+ "*.json",
10
+ "*.md",
11
+ "*.txt",
12
+ "tokenizer.json"
13
+ ],
14
+ "ignore_patterns": [
15
+ "__pycache__/*",
16
+ "*.pyc",
17
+ "*.pyo",
18
+ "_test_*",
19
+ "*.safetensors",
20
+ "*.bin",
21
+ "*.pt",
22
+ "*.pth",
23
+ ".git",
24
+ ".git/*",
25
+ "*/.git",
26
+ "**/.git/**",
27
+ ".cache/huggingface",
28
+ ".cache/huggingface/*",
29
+ "*/.cache/huggingface",
30
+ "**/.cache/huggingface/**"
31
+ ],
32
+ "local_files_count": 64,
33
+ "local_files_total_bytes": 788009,
34
+ "remote_files_before": 63,
35
+ "remote_files_after": 67,
36
+ "new_files_count": 4,
37
+ "overwritten_count": 64,
38
+ "critical_files_all_present": true,
39
+ "elapsed_s": 2.654574849000028,
40
+ "commit_info": "https://huggingface.co/PowerMachine/BiGRU_T_version/commit/a22ed76d93580e2cc99e084b60e54080bb45cc9c",
41
+ "env": {
42
+ "HF_HUB_ENABLE_HF_TRANSFER": "1",
43
+ "HF_HUB_ETAG_TIMEOUT": "60",
44
+ "OMP_NUM_THREADS": "2",
45
+ "TOKENIZERS_PARALLELISM": "true",
46
+ "PYTHONUNBUFFERED": "1"
47
+ }
48
+ }
v4_report.json ADDED
@@ -0,0 +1,122 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "version": "V4",
3
+ "timestamp": 1786016531.119913,
4
+ "config": {
5
+ "n_multimodal": 10,
6
+ "n_text": 10,
7
+ "epochs": 2,
8
+ "batch_size": 4,
9
+ "max_seq_len": 16,
10
+ "num_layers_hyp_sweep": [
11
+ 2,
12
+ 8
13
+ ],
14
+ "triplet_weight": 0.05
15
+ },
16
+ "xeon_env": {
17
+ "OMP_NUM_THREADS": "2",
18
+ "MKL_NUM_THREADS": "2",
19
+ "ONEDNN_MAX_CPU_ISA": "AMX_INT8",
20
+ "IPEX": true,
21
+ "CPU_ISA": "AMX",
22
+ "torch_threads": 2
23
+ },
24
+ "results": {
25
+ "num_layers_hyp_2": {
26
+ "n_steps": 6,
27
+ "n_epochs": 2,
28
+ "final_train_loss": 1.492246150970459,
29
+ "final_val_loss": 5.687218344211578,
30
+ "final_perplexity": 295.0716901317322,
31
+ "final_overfitting_gap": -1.3897555629412333,
32
+ "final_eval_benchmark": 0.375,
33
+ "avg_tokens_per_sec": 30.952014369483933,
34
+ "avg_tflops_per_cpu": 0.00010286999861443428,
35
+ "avg_time_per_step_ms": 2092.294035499966,
36
+ "avg_grad_norm": 19.630864143371582,
37
+ "avg_cpu_pct": 120.28333333333335,
38
+ "n_alerts": 2,
39
+ "alerts": [
40
+ {
41
+ "step": 5,
42
+ "type": "loss_spike",
43
+ "delta_pct": 0.5995766717862026,
44
+ "from": 8.140497207641602,
45
+ "to": 3.2596449851989746
46
+ },
47
+ {
48
+ "step": 6,
49
+ "type": "loss_spike",
50
+ "delta_pct": 0.5422059280239778,
51
+ "from": 3.2596449851989746,
52
+ "to": 1.492246150970459
53
+ }
54
+ ],
55
+ "alert_types": [
56
+ "loss_spike"
57
+ ],
58
+ "elapsed_s": 18.622390889000144,
59
+ "n_params": 8862751,
60
+ "hyp_T_params": 1001034,
61
+ "num_layers_hyp": 2,
62
+ "final_bigru_diversity": {
63
+ "mean_cosine_diversity": 0.6990604996681213,
64
+ "mean_parallelism_score": 0.9938665628433228,
65
+ "mean_gate_norm_ratio": 0.11281532049179077,
66
+ "module_level_cosine_diversity": 0.8235123827388244,
67
+ "per_module_parallelism": [
68
+ 0.9952576756477356,
69
+ 0.9962905645370483,
70
+ 0.9829346537590027,
71
+ 0.9921810626983643,
72
+ 0.9968430399894714,
73
+ 0.9968783855438232,
74
+ 0.9929060935974121,
75
+ 0.9943419694900513
76
+ ]
77
+ }
78
+ },
79
+ "num_layers_hyp_8": {
80
+ "n_steps": 6,
81
+ "n_epochs": 2,
82
+ "final_train_loss": 4.706010818481445,
83
+ "final_val_loss": 5.576698219776153,
84
+ "final_perplexity": 264.19784088397313,
85
+ "final_overfitting_gap": -1.148307160536448,
86
+ "final_eval_benchmark": 0.25,
87
+ "avg_tokens_per_sec": 31.42855771366708,
88
+ "avg_tflops_per_cpu": 0.0001138220299728001,
89
+ "avg_time_per_step_ms": 2074.0577321666933,
90
+ "avg_grad_norm": 15.662032763163248,
91
+ "avg_cpu_pct": 120.39999999999999,
92
+ "n_alerts": 0,
93
+ "alerts": [],
94
+ "alert_types": [],
95
+ "elapsed_s": 18.50429308999992,
96
+ "n_params": 9657631,
97
+ "hyp_T_params": 1795914,
98
+ "num_layers_hyp": 8,
99
+ "final_bigru_diversity": {
100
+ "mean_cosine_diversity": 0.6962708830833435,
101
+ "mean_parallelism_score": 0.9929798245429993,
102
+ "mean_gate_norm_ratio": 0.124363474547863,
103
+ "module_level_cosine_diversity": 0.7712274699338845,
104
+ "per_module_parallelism": [
105
+ 0.9903037548065186,
106
+ 0.9907274842262268,
107
+ 0.9923153519630432,
108
+ 0.9945146441459656,
109
+ 0.9895800948143005,
110
+ 0.9969634413719177,
111
+ 0.9970486164093018,
112
+ 0.997769296169281
113
+ ]
114
+ }
115
+ }
116
+ },
117
+ "analysis": {
118
+ "best_num_layers_hyp": 8,
119
+ "best_val_loss": 5.576698219776153,
120
+ "best_eval_acc": 0.375
121
+ }
122
+ }