V3.2: corrigir 6 bugs + meta_T_moved (T_tracking_loss) + re-teste com Xeon IPEX 2.8 + AMX
Browse files- bug_hunt_v3_report.json +150 -0
- model_final/config.json +15 -34
- src/bigru_t/__init__.py +4 -2
- src/bigru_t/data/streaming_datasets.py +2 -1
- src/bigru_t/model/module_selector.py +22 -6
- src/bigru_t/model/unified_model.py +54 -1
- src/bigru_t/training/hypothesis_monitor.py +373 -0
- src/bigru_t/training/meta_configurator.py +40 -3
- src/bigru_t/training/trainer.py +226 -14
- training_report.json +60 -61
- v32_test_report.json +74 -0
- v3_report.json +119 -0
bug_hunt_v3_report.json
ADDED
|
@@ -0,0 +1,150 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"total": 31,
|
| 3 |
+
"passed": 31,
|
| 4 |
+
"failed": 0,
|
| 5 |
+
"results": {
|
| 6 |
+
"system_ram": {
|
| 7 |
+
"ok": true,
|
| 8 |
+
"detail": "3284MB free / 4042MB total"
|
| 9 |
+
},
|
| 10 |
+
"system_disk": {
|
| 11 |
+
"ok": true,
|
| 12 |
+
"detail": "7.82GB free"
|
| 13 |
+
},
|
| 14 |
+
"xeon_ipex": {
|
| 15 |
+
"ok": true,
|
| 16 |
+
"detail": "IPEX=OK"
|
| 17 |
+
},
|
| 18 |
+
"xeon_isa": {
|
| 19 |
+
"ok": true,
|
| 20 |
+
"detail": "ISA=AMX"
|
| 21 |
+
},
|
| 22 |
+
"xeon_omp": {
|
| 23 |
+
"ok": true,
|
| 24 |
+
"detail": "OMP_NUM_THREADS=2"
|
| 25 |
+
},
|
| 26 |
+
"xeon_mkl": {
|
| 27 |
+
"ok": true,
|
| 28 |
+
"detail": "MKL_NUM_THREADS=2"
|
| 29 |
+
},
|
| 30 |
+
"xeonednn": {
|
| 31 |
+
"ok": true,
|
| 32 |
+
"detail": "ONEDNN_MAX_CPU_ISA=AMX_INT8"
|
| 33 |
+
},
|
| 34 |
+
"xeon_torch_threads": {
|
| 35 |
+
"ok": true,
|
| 36 |
+
"detail": "torch_threads=2"
|
| 37 |
+
},
|
| 38 |
+
"imports_all": {
|
| 39 |
+
"ok": true,
|
| 40 |
+
"detail": "40+ s\u00edmbolos importados OK"
|
| 41 |
+
},
|
| 42 |
+
"model_params": {
|
| 43 |
+
"ok": true,
|
| 44 |
+
"detail": "13,033,867 (13.03M)"
|
| 45 |
+
},
|
| 46 |
+
"model_forward_shape": {
|
| 47 |
+
"ok": true,
|
| 48 |
+
"detail": "y_hat=(2, 16384)"
|
| 49 |
+
},
|
| 50 |
+
"model_forward_no_nan": {
|
| 51 |
+
"ok": true,
|
| 52 |
+
"detail": "y_hat e delta sem NaN/Inf"
|
| 53 |
+
},
|
| 54 |
+
"model_crw_active": {
|
| 55 |
+
"ok": true,
|
| 56 |
+
"detail": "CRW: cycles=3, w2=0.005203, q\u0304=0.9958"
|
| 57 |
+
},
|
| 58 |
+
"model_crw_contraction": {
|
| 59 |
+
"ok": true,
|
| 60 |
+
"detail": "contraction_valid=True"
|
| 61 |
+
},
|
| 62 |
+
"model_memory": {
|
| 63 |
+
"ok": true,
|
| 64 |
+
"detail": "RSS ap\u00f3s modelo: 620.0MB"
|
| 65 |
+
},
|
| 66 |
+
"backward_grad_flow": {
|
| 67 |
+
"ok": true,
|
| 68 |
+
"detail": "2228/2256 params com grad"
|
| 69 |
+
},
|
| 70 |
+
"backward_zero_grad": {
|
| 71 |
+
"ok": true,
|
| 72 |
+
"detail": "0 params com grad \u2248 0"
|
| 73 |
+
},
|
| 74 |
+
"backward_grad_norm": {
|
| 75 |
+
"ok": true,
|
| 76 |
+
"detail": "||grad||=20.0964"
|
| 77 |
+
},
|
| 78 |
+
"gs_grad_assigned": {
|
| 79 |
+
"ok": true,
|
| 80 |
+
"detail": "1161 params com .grad atribu\u00eddo"
|
| 81 |
+
},
|
| 82 |
+
"gs_no_nan": {
|
| 83 |
+
"ok": true,
|
| 84 |
+
"detail": "0 params com grad NaN"
|
| 85 |
+
},
|
| 86 |
+
"meta_runs": {
|
| 87 |
+
"ok": true,
|
| 88 |
+
"detail": "val_loss=9.9373, sharpness=385.9175"
|
| 89 |
+
},
|
| 90 |
+
"meta_T_moved": {
|
| 91 |
+
"ok": true,
|
| 92 |
+
"detail": "T: 1.000000 \u2192 1.010050"
|
| 93 |
+
},
|
| 94 |
+
"meta_tau_moved": {
|
| 95 |
+
"ok": true,
|
| 96 |
+
"detail": "\u03c4: 1.000000 \u2192 1.010050"
|
| 97 |
+
},
|
| 98 |
+
"dpo_loss_runs": {
|
| 99 |
+
"ok": true,
|
| 100 |
+
"detail": "loss=0.6931, \u03b2=0.2750"
|
| 101 |
+
},
|
| 102 |
+
"dpo_loss_grad": {
|
| 103 |
+
"ok": true,
|
| 104 |
+
"detail": "gradientes flu\u00edram"
|
| 105 |
+
},
|
| 106 |
+
"dpo_step_runs": {
|
| 107 |
+
"ok": true,
|
| 108 |
+
"detail": "dpo_loss=0.6735, acc=1.00"
|
| 109 |
+
},
|
| 110 |
+
"monitor_runs": {
|
| 111 |
+
"ok": true,
|
| 112 |
+
"detail": "steps=5, punishment=2"
|
| 113 |
+
},
|
| 114 |
+
"monitor_bug_detection": {
|
| 115 |
+
"ok": true,
|
| 116 |
+
"detail": "has_bug=False, bugs=0"
|
| 117 |
+
},
|
| 118 |
+
"multimodal_encoders": {
|
| 119 |
+
"ok": true,
|
| 120 |
+
"detail": "5/5 OK: {'text': torch.Size([2, 16, 128]), 'image': (2, 128), 'audio': (2, 128), 'video': (2, 128), 'router_fusion': ((1, 128), (2, 128))}"
|
| 121 |
+
},
|
| 122 |
+
"ks_normal": {
|
| 123 |
+
"ok": true,
|
| 124 |
+
"detail": "step 1: reason=None, RAM=24.7%"
|
| 125 |
+
},
|
| 126 |
+
"ks_summary": {
|
| 127 |
+
"ok": true,
|
| 128 |
+
"detail": "keys=['total_steps', 'min_loss', 'max_ram_pct', 'min_disk_free_gb', 'max_ppl', 'killed', 'kill_reason']"
|
| 129 |
+
}
|
| 130 |
+
},
|
| 131 |
+
"model_info": {
|
| 132 |
+
"params": {
|
| 133 |
+
"total": 13033867,
|
| 134 |
+
"total_M": 13.033867,
|
| 135 |
+
"trainable": 13033867,
|
| 136 |
+
"trainable_M": 13.033867
|
| 137 |
+
},
|
| 138 |
+
"param_breakdown": {
|
| 139 |
+
"token_embedding": 2097152,
|
| 140 |
+
"module_selector": 8,
|
| 141 |
+
"u8cells": 5937664,
|
| 142 |
+
"orq_cell": 142848,
|
| 143 |
+
"train_T": 2395008,
|
| 144 |
+
"hyp_T": 2395008,
|
| 145 |
+
"circular_reasoning": 66179
|
| 146 |
+
},
|
| 147 |
+
"rss_mb": 620.03125
|
| 148 |
+
},
|
| 149 |
+
"timestamp": 1786012789.1891277
|
| 150 |
+
}
|
model_final/config.json
CHANGED
|
@@ -1,6 +1,6 @@
|
|
| 1 |
{
|
| 2 |
"model_type": "bigru_t_version",
|
| 3 |
-
"arch": "BiGRU_T_version
|
| 4 |
"vocab_size": 16384,
|
| 5 |
"d_model": 128,
|
| 6 |
"max_seq_len": 32,
|
|
@@ -27,41 +27,22 @@
|
|
| 27 |
"num_layers_hyp": 2,
|
| 28 |
"num_bits": 8,
|
| 29 |
"dropout": 0.1,
|
| 30 |
-
"total_params":
|
| 31 |
-
"total_params_M":
|
| 32 |
"training": {
|
| 33 |
-
"version": "
|
| 34 |
-
"
|
| 35 |
-
"
|
| 36 |
-
"
|
| 37 |
-
"
|
| 38 |
-
"
|
| 39 |
-
"
|
| 40 |
-
"
|
| 41 |
-
"
|
| 42 |
-
"
|
| 43 |
-
"optimizer": "HamiltonianWassersteinOptimizer",
|
| 44 |
-
"hw_features": [
|
| 45 |
-
"AdamW",
|
| 46 |
-
"W2_adaptive",
|
| 47 |
-
"repulsion_topological",
|
| 48 |
-
"LR_cyclic_Van_der_Pol"
|
| 49 |
-
],
|
| 50 |
-
"meta_configurator": "Lema 4 active (T=1.01, tau=1.0)",
|
| 51 |
-
"gradient_surgery": "Lema 2 active",
|
| 52 |
-
"w8a8_quantization": "Lema 3 active",
|
| 53 |
-
"module_selector": "Lema 1 active (8 modules, softmax+entropy)",
|
| 54 |
-
"multimodal_test": "5/5 encoders OK",
|
| 55 |
-
"circular_reasoning": "CircularReasoningWasserstein available",
|
| 56 |
-
"dpo": "Standalone dpo_loss available (beta adaptive)",
|
| 57 |
-
"inference": "BiGRUTGenerator (greedy + top-k sampling)",
|
| 58 |
-
"tokenizer": "BBPE parallel Map-Reduce (ProcessPoolExecutor)",
|
| 59 |
-
"memory_cleanup": "aggressive_cleanup + TimeBudget (Xavante style)",
|
| 60 |
"datasets": [
|
| 61 |
-
"
|
| 62 |
],
|
| 63 |
-
"
|
| 64 |
-
"killed": false,
|
| 65 |
-
"note": "Epoch 1 fully completed + saved. Epoch 2 reached step 5 (loss 5.69) before process ended."
|
| 66 |
}
|
| 67 |
}
|
|
|
|
| 1 |
{
|
| 2 |
"model_type": "bigru_t_version",
|
| 3 |
+
"arch": "BiGRU_T_version (4 lemas)",
|
| 4 |
"vocab_size": 16384,
|
| 5 |
"d_model": 128,
|
| 6 |
"max_seq_len": 32,
|
|
|
|
| 27 |
"num_layers_hyp": 2,
|
| 28 |
"num_bits": 8,
|
| 29 |
"dropout": 0.1,
|
| 30 |
+
"total_params": 13033867,
|
| 31 |
+
"total_params_M": 13.033867,
|
| 32 |
"training": {
|
| 33 |
+
"version": "BiGRU_T_version",
|
| 34 |
+
"epochs": 5,
|
| 35 |
+
"global_step": 25,
|
| 36 |
+
"best_loss": 2.1027424335479736,
|
| 37 |
+
"final_loss": 4.172516345977783,
|
| 38 |
+
"final_avg_loss": 3.9418057203292847,
|
| 39 |
+
"final_perplexity": 51.511532793090105,
|
| 40 |
+
"killed": false,
|
| 41 |
+
"kill_reason": null,
|
| 42 |
+
"elapsed_s": 427.32008838653564,
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 43 |
"datasets": [
|
| 44 |
+
"nvidia/OpenMathInstruct-2"
|
| 45 |
],
|
| 46 |
+
"max_samples_per_dataset": 12
|
|
|
|
|
|
|
| 47 |
}
|
| 48 |
}
|
src/bigru_t/__init__.py
CHANGED
|
@@ -26,7 +26,8 @@ from .training.gradient_surgery import apply_gradient_surgery, orthogonalize_gra
|
|
| 26 |
from .training.meta_configurator import MetaConfigurator
|
| 27 |
from .training.kill_switch import KillSwitch, KillSwitchState
|
| 28 |
from .training.trainer import BiGRU_T_Trainer, TrainerConfig
|
| 29 |
-
from .training.dpo import dpo_loss, compute_dynamic_beta, compute_sequence_logps
|
|
|
|
| 30 |
|
| 31 |
from .reasoning.circular_reasoning_wasserstein import CircularReasoningWasserstein
|
| 32 |
|
|
@@ -47,7 +48,8 @@ __all__ = [
|
|
| 47 |
"MetaConfigurator",
|
| 48 |
"KillSwitch", "KillSwitchState",
|
| 49 |
"BiGRU_T_Trainer", "TrainerConfig",
|
| 50 |
-
"dpo_loss", "compute_dynamic_beta", "compute_sequence_logps",
|
|
|
|
| 51 |
"CircularReasoningWasserstein",
|
| 52 |
"BiGRUTGenerator",
|
| 53 |
"HamiltonianWassersteinOptimizer",
|
|
|
|
| 26 |
from .training.meta_configurator import MetaConfigurator
|
| 27 |
from .training.kill_switch import KillSwitch, KillSwitchState
|
| 28 |
from .training.trainer import BiGRU_T_Trainer, TrainerConfig
|
| 29 |
+
from .training.dpo import dpo_loss, compute_dynamic_beta, compute_sequence_logps, dpo_step
|
| 30 |
+
from .training.hypothesis_monitor import HypothesisMonitor, HypothesisMonitorReport
|
| 31 |
|
| 32 |
from .reasoning.circular_reasoning_wasserstein import CircularReasoningWasserstein
|
| 33 |
|
|
|
|
| 48 |
"MetaConfigurator",
|
| 49 |
"KillSwitch", "KillSwitchState",
|
| 50 |
"BiGRU_T_Trainer", "TrainerConfig",
|
| 51 |
+
"dpo_loss", "compute_dynamic_beta", "compute_sequence_logps", "dpo_step",
|
| 52 |
+
"HypothesisMonitor", "HypothesisMonitorReport",
|
| 53 |
"CircularReasoningWasserstein",
|
| 54 |
"BiGRUTGenerator",
|
| 55 |
"HamiltonianWassersteinOptimizer",
|
src/bigru_t/data/streaming_datasets.py
CHANGED
|
@@ -94,7 +94,8 @@ DATASET_FORMATS: Dict[str, Dict[str, Any]] = {
|
|
| 94 |
"nvidia/OpenMathInstruct-2": {
|
| 95 |
"type": "math",
|
| 96 |
"text_fields": ["problem", "question", "input", "prompt"],
|
| 97 |
-
|
|
|
|
| 98 |
"split": "train",
|
| 99 |
"config": None,
|
| 100 |
},
|
|
|
|
| 94 |
"nvidia/OpenMathInstruct-2": {
|
| 95 |
"type": "math",
|
| 96 |
"text_fields": ["problem", "question", "input", "prompt"],
|
| 97 |
+
# V3 FIX: campos reais do dataset são generated_solution e expected_answer
|
| 98 |
+
"label_fields": ["generated_solution", "expected_answer", "solution", "answer", "output", "response"],
|
| 99 |
"split": "train",
|
| 100 |
"config": None,
|
| 101 |
},
|
src/bigru_t/model/module_selector.py
CHANGED
|
@@ -48,18 +48,34 @@ class ModuleSelector(nn.Module):
|
|
| 48 |
self.max_modules = max_modules
|
| 49 |
self.lambda_ent = lambda_ent
|
| 50 |
|
| 51 |
-
#
|
| 52 |
-
|
| 53 |
-
|
| 54 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 55 |
"""Retorna (alpha, entropy_reg).
|
| 56 |
|
| 57 |
alpha: (K,) — pesos de seleção dos módulos (soma = 1)
|
| 58 |
entropy_reg: scalar — penalização de entropia (a somar à loss)
|
|
|
|
|
|
|
|
|
|
|
|
|
| 59 |
"""
|
| 60 |
-
#
|
|
|
|
|
|
|
|
|
|
| 61 |
# T alta → distribuição uniforme; T baixa → argmax (especialização)
|
| 62 |
-
|
|
|
|
| 63 |
|
| 64 |
# Regularização de entropia (minimizar H(alpha) → especialização)
|
| 65 |
# H(alpha) = -sum(alpha * log(alpha + eps))
|
|
|
|
| 48 |
self.max_modules = max_modules
|
| 49 |
self.lambda_ent = lambda_ent
|
| 50 |
|
| 51 |
+
# V3 BUGFIX: init com pequeno ruído aleatório (em vez de zeros) para
|
| 52 |
+
# quebrar a simetria. Com zeros, softmax(0/T) = uniforme, e o gradiente
|
| 53 |
+
# w.r.t. module_logits é o mesmo para todos os módulos → nenhum aprende
|
| 54 |
+
# a se diferenciar. Com ruído, cada módulo tem um leve viés inicial
|
| 55 |
+
# que o otimizador pode amplificar.
|
| 56 |
+
# Scale = 1/sqrt(K) para manter variance razoável.
|
| 57 |
+
init_scale = 1.0 / (max_modules ** 0.5)
|
| 58 |
+
self.module_logits = nn.Parameter(
|
| 59 |
+
torch.randn(max_modules) * init_scale
|
| 60 |
+
)
|
| 61 |
+
|
| 62 |
+
def forward(self, temperature=1.0) -> tuple[torch.Tensor, torch.Tensor]:
|
| 63 |
"""Retorna (alpha, entropy_reg).
|
| 64 |
|
| 65 |
alpha: (K,) — pesos de seleção dos módulos (soma = 1)
|
| 66 |
entropy_reg: scalar — penalização de entropia (a somar à loss)
|
| 67 |
+
|
| 68 |
+
V3 BUGFIX: temperature pode ser tensor (do MetaConfigurator) ou float.
|
| 69 |
+
Usamos torch.clamp (diferenciável) em vez de max() (que pode quebrar
|
| 70 |
+
o grafo computacional quando temperature é tensor 0-dim).
|
| 71 |
"""
|
| 72 |
+
# Garante que temperature é um tensor (preserva grafo se for tensor)
|
| 73 |
+
if not isinstance(temperature, torch.Tensor):
|
| 74 |
+
temperature = torch.tensor(float(temperature), device=self.module_logits.device)
|
| 75 |
+
# Softmax com temperatura (Lema 1) — torch.clamp é diferenciável
|
| 76 |
# T alta → distribuição uniforme; T baixa → argmax (especialização)
|
| 77 |
+
T_clamped = torch.clamp(temperature, min=1e-6)
|
| 78 |
+
alpha = F.softmax(self.module_logits / T_clamped, dim=0) # (K,)
|
| 79 |
|
| 80 |
# Regularização de entropia (minimizar H(alpha) → especialização)
|
| 81 |
# H(alpha) = -sum(alpha * log(alpha + eps))
|
src/bigru_t/model/unified_model.py
CHANGED
|
@@ -53,6 +53,7 @@ from .train_t import TrainT
|
|
| 53 |
from .hyp_t import HypT
|
| 54 |
from .module_selector import ModuleSelector
|
| 55 |
from ..quantization.quantized_linear import QuantizedLinear, apply_w8a8
|
|
|
|
| 56 |
|
| 57 |
|
| 58 |
@dataclass
|
|
@@ -113,6 +114,12 @@ class UnifiedModelConfig:
|
|
| 113 |
# Quantization (Lema 3)
|
| 114 |
num_bits: int = 8
|
| 115 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 116 |
# Misc
|
| 117 |
dropout: float = 0.1
|
| 118 |
|
|
@@ -243,6 +250,20 @@ class UnifiedModel(nn.Module):
|
|
| 243 |
# Init = 1.0 (sem hipótese até a loss cair abaixo de 1.0 — raro no início)
|
| 244 |
self.register_buffer("tau", torch.tensor(float("inf")))
|
| 245 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 246 |
# Aplicar W8A8 em todas as camadas lineares (exceto Embedding, que é nn.Embedding)
|
| 247 |
# Nota: nn.GRU não é afetada (não é Linear)
|
| 248 |
self.apply(apply_w8a8)
|
|
@@ -254,6 +275,7 @@ class UnifiedModel(nn.Module):
|
|
| 254 |
use_hypothesis: bool = False,
|
| 255 |
stop_grad_hyp: bool = True,
|
| 256 |
return_aux: bool = False,
|
|
|
|
| 257 |
) -> tuple[torch.Tensor, torch.Tensor]:
|
| 258 |
"""Forward do UnifiedModel.
|
| 259 |
|
|
@@ -262,11 +284,19 @@ class UnifiedModel(nn.Module):
|
|
| 262 |
temperature: temperatura do softmax do Lema 1
|
| 263 |
use_hypothesis: se True, calcula delta (cabeça HypT ativa)
|
| 264 |
stop_grad_hyp: se True, detach a entrada da HypT (Lema 3)
|
| 265 |
-
return_aux: se True, retorna também entropy_reg e
|
|
|
|
|
|
|
| 266 |
|
| 267 |
Returns:
|
| 268 |
(y_hat, delta) por padrão.
|
| 269 |
(y_hat, delta, aux_dict) se return_aux=True.
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 270 |
"""
|
| 271 |
# Embedding se necessário
|
| 272 |
if x.dim() == 2:
|
|
@@ -290,6 +320,28 @@ class UnifiedModel(nn.Module):
|
|
| 290 |
# Orquestrador: produz representação agregada
|
| 291 |
o = self.orq_cell(H) # (batch, d_cache)
|
| 292 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 293 |
# Predição principal
|
| 294 |
y_hat_main = self.train_T(o) # (batch, vocab_size)
|
| 295 |
|
|
@@ -304,6 +356,7 @@ class UnifiedModel(nn.Module):
|
|
| 304 |
"entropy_reg": entropy_reg,
|
| 305 |
"alpha": alpha.detach(),
|
| 306 |
"o": o.detach(),
|
|
|
|
| 307 |
}
|
| 308 |
return y_hat_main, delta, aux
|
| 309 |
return y_hat_main, delta
|
|
|
|
| 53 |
from .hyp_t import HypT
|
| 54 |
from .module_selector import ModuleSelector
|
| 55 |
from ..quantization.quantized_linear import QuantizedLinear, apply_w8a8
|
| 56 |
+
from ..reasoning.circular_reasoning_wasserstein import CircularReasoningWasserstein
|
| 57 |
|
| 58 |
|
| 59 |
@dataclass
|
|
|
|
| 114 |
# Quantization (Lema 3)
|
| 115 |
num_bits: int = 8
|
| 116 |
|
| 117 |
+
# Circular Reasoning (Teorema 4 — contração W₂)
|
| 118 |
+
use_circular_reasoning: bool = True
|
| 119 |
+
circular_n_cycles: int = 3
|
| 120 |
+
circular_contraction_target: float = 0.9
|
| 121 |
+
circular_init_gate: float = 0.05
|
| 122 |
+
|
| 123 |
# Misc
|
| 124 |
dropout: float = 0.1
|
| 125 |
|
|
|
|
| 250 |
# Init = 1.0 (sem hipótese até a loss cair abaixo de 1.0 — raro no início)
|
| 251 |
self.register_buffer("tau", torch.tensor(float("inf")))
|
| 252 |
|
| 253 |
+
# Circular Reasoning (Teorema 4 — contração W₂ sobre a representação agregada)
|
| 254 |
+
# Aplicado APÓS OrqCell e ANTES de TrainT/HypT. Refina 'o' em n_cycles
|
| 255 |
+
# iterações de Φ(s) = LayerNorm(s + gate * RefineHead(s)).
|
| 256 |
+
self.circular_reasoning = (
|
| 257 |
+
CircularReasoningWasserstein(
|
| 258 |
+
d_model=c.d_cache,
|
| 259 |
+
n_cycles=c.circular_n_cycles,
|
| 260 |
+
contraction_target=c.circular_contraction_target,
|
| 261 |
+
init_gate=c.circular_init_gate,
|
| 262 |
+
)
|
| 263 |
+
if c.use_circular_reasoning
|
| 264 |
+
else None
|
| 265 |
+
)
|
| 266 |
+
|
| 267 |
# Aplicar W8A8 em todas as camadas lineares (exceto Embedding, que é nn.Embedding)
|
| 268 |
# Nota: nn.GRU não é afetada (não é Linear)
|
| 269 |
self.apply(apply_w8a8)
|
|
|
|
| 275 |
use_hypothesis: bool = False,
|
| 276 |
stop_grad_hyp: bool = True,
|
| 277 |
return_aux: bool = False,
|
| 278 |
+
use_circular: Optional[bool] = None,
|
| 279 |
) -> tuple[torch.Tensor, torch.Tensor]:
|
| 280 |
"""Forward do UnifiedModel.
|
| 281 |
|
|
|
|
| 284 |
temperature: temperatura do softmax do Lema 1
|
| 285 |
use_hypothesis: se True, calcula delta (cabeça HypT ativa)
|
| 286 |
stop_grad_hyp: se True, detach a entrada da HypT (Lema 3)
|
| 287 |
+
return_aux: se True, retorna também entropy_reg, alpha, e info circular
|
| 288 |
+
use_circular: override p/ ativar/desativar CircularReasoning no forward
|
| 289 |
+
(default: usa self.config.use_circular_reasoning)
|
| 290 |
|
| 291 |
Returns:
|
| 292 |
(y_hat, delta) por padrão.
|
| 293 |
(y_hat, delta, aux_dict) se return_aux=True.
|
| 294 |
+
|
| 295 |
+
Fluxo completo (V3):
|
| 296 |
+
x → embedding → [L1: module_selector α] → K×u8cell_T → OrqCell → o
|
| 297 |
+
→ [CRW: Φ^k(o) por n_cycles] → o* (contração W₂)
|
| 298 |
+
→ TrainT(o*) = y_hat
|
| 299 |
+
→ HypT(o*, stop_grad) = delta (se use_hypothesis)
|
| 300 |
"""
|
| 301 |
# Embedding se necessário
|
| 302 |
if x.dim() == 2:
|
|
|
|
| 320 |
# Orquestrador: produz representação agregada
|
| 321 |
o = self.orq_cell(H) # (batch, d_cache)
|
| 322 |
|
| 323 |
+
# Circular Reasoning (Teorema 4): refine 'o' via contração W₂
|
| 324 |
+
circular_info = None
|
| 325 |
+
do_circular = (
|
| 326 |
+
use_circular if use_circular is not None
|
| 327 |
+
else self.config.use_circular_reasoning
|
| 328 |
+
)
|
| 329 |
+
if do_circular and self.circular_reasoning is not None:
|
| 330 |
+
# CRW aceita (B, D) ou (B, T, D); passamos (B, D) — LayerNorm normaliza última dim
|
| 331 |
+
crw_out = self.circular_reasoning(o)
|
| 332 |
+
o = crw_out["output"] # representação refinada s*
|
| 333 |
+
circular_info = {
|
| 334 |
+
"converged": crw_out["converged"],
|
| 335 |
+
"w2_final": crw_out["w2_final"],
|
| 336 |
+
"n_cycles_used": crw_out["n_cycles_used"],
|
| 337 |
+
"converged_step": crw_out["converged_step"],
|
| 338 |
+
"contraction_valid": crw_out["contraction_valid"],
|
| 339 |
+
"q_empirical_avg": (
|
| 340 |
+
sum(crw_out["q_empirical"]) / max(1, len(crw_out["q_empirical"]))
|
| 341 |
+
if crw_out["q_empirical"] else 1.0
|
| 342 |
+
),
|
| 343 |
+
}
|
| 344 |
+
|
| 345 |
# Predição principal
|
| 346 |
y_hat_main = self.train_T(o) # (batch, vocab_size)
|
| 347 |
|
|
|
|
| 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
|
src/bigru_t/training/hypothesis_monitor.py
ADDED
|
@@ -0,0 +1,373 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""hypothesis_monitor.py — Detector de ativação das hipóteses (L1-L4).
|
| 2 |
+
|
| 3 |
+
Filosofia: "se as hipóteses não são ativadas durante punição, há um bug".
|
| 4 |
+
|
| 5 |
+
Cada um dos 4 lemas tem uma condição de ativação observável:
|
| 6 |
+
L1 (ModuleSelector): α = softmax(logits/T) — ativa quando há variação em α
|
| 7 |
+
(entropy < log(K) - ε). Se entropy ≈ log(K), nenhum módulo é selecionado
|
| 8 |
+
(todos com peso igual) — possível bug no temperature ou nos logits.
|
| 9 |
+
L2 (GradientSurgery): ativa quando use_hyp=True e apply_gradient_surgery
|
| 10 |
+
modifica os gradientes. Se ||g_hyp|| ≈ 0 ou projeção ortogonal ≈ 0,
|
| 11 |
+
a cirurgia é no-op (possível bug: loss_hyp não tem gradiente).
|
| 12 |
+
L3 (HypT): ativa quando loss > τ (modelo está "lutando"). Se loss > τ mas
|
| 13 |
+
use_hyp=False (não foi ativada), é bug. Se use_hyp=True mas ||delta|| ≈ 0,
|
| 14 |
+
a hipótese não contribui (possível bug: HypT morta ou stop_grad excessivo).
|
| 15 |
+
L4 (MetaConfigurator): ativa a cada meta_interval passos. Se T e τ não mudam
|
| 16 |
+
entre chamadas, o meta-otimizador está estagnado (possível bug: lr=0 ou
|
| 17 |
+
gradientes None).
|
| 18 |
+
|
| 19 |
+
O monitor mantém contadores e dispara alertas quando:
|
| 20 |
+
- Punishment_count > 0 mas activation_count = 0 (hipótese nunca ativada)
|
| 21 |
+
- Activation_count > 0 mas delta_norm_medio < epsilon (HypT morta)
|
| 22 |
+
- T ou τ não mudam em N chamadas do meta (estagnação)
|
| 23 |
+
- Entropy do L1 saturada (> 0.95 * log(K)) — sem seleção real
|
| 24 |
+
"""
|
| 25 |
+
from __future__ import annotations
|
| 26 |
+
|
| 27 |
+
import logging
|
| 28 |
+
import math
|
| 29 |
+
from collections import deque
|
| 30 |
+
from dataclasses import dataclass, field
|
| 31 |
+
from typing import Optional
|
| 32 |
+
|
| 33 |
+
import torch
|
| 34 |
+
|
| 35 |
+
logger = logging.getLogger(__name__)
|
| 36 |
+
|
| 37 |
+
|
| 38 |
+
@dataclass
|
| 39 |
+
class LemmaState:
|
| 40 |
+
"""Snapshot do estado de ativação de cada lema em um step."""
|
| 41 |
+
step: int
|
| 42 |
+
# L1
|
| 43 |
+
alpha_entropy: float = 0.0
|
| 44 |
+
alpha_max: float = 0.0
|
| 45 |
+
alpha_active_modules: int = 0
|
| 46 |
+
temperature: float = 1.0
|
| 47 |
+
l1_activated: bool = False # True se entropy < 0.95 * log(K)
|
| 48 |
+
# L2
|
| 49 |
+
l2_applied: bool = False
|
| 50 |
+
grad_main_norm: float = 0.0
|
| 51 |
+
grad_hyp_norm: float = 0.0
|
| 52 |
+
orth_proj_norm: float = 0.0
|
| 53 |
+
# L3
|
| 54 |
+
loss_main: float = 0.0
|
| 55 |
+
tau: float = float("inf")
|
| 56 |
+
use_hyp_activated: bool = False # True se use_hyp=True
|
| 57 |
+
delta_norm: float = 0.0 # ||delta||₂ — se 0, HypT morta
|
| 58 |
+
punishment_triggered: bool = False # True se loss > τ
|
| 59 |
+
# L4
|
| 60 |
+
T_new: float = 1.0
|
| 61 |
+
tau_new: float = 1.0
|
| 62 |
+
sharpness: float = 0.0
|
| 63 |
+
l4_meta_called: bool = False
|
| 64 |
+
l4_T_moved: bool = False
|
| 65 |
+
l4_tau_moved: bool = False
|
| 66 |
+
# Circular Reasoning (Teorema 4)
|
| 67 |
+
crw_converged: bool = False
|
| 68 |
+
crw_w2_final: float = 0.0
|
| 69 |
+
crw_q_avg: float = 1.0
|
| 70 |
+
crw_contraction_valid: bool = False
|
| 71 |
+
|
| 72 |
+
|
| 73 |
+
@dataclass
|
| 74 |
+
class HypothesisMonitorReport:
|
| 75 |
+
"""Relatório final do monitor."""
|
| 76 |
+
total_steps: int = 0
|
| 77 |
+
punishment_steps: int = 0 # steps com loss > τ
|
| 78 |
+
hyp_activated_steps: int = 0 # steps com use_hyp=True
|
| 79 |
+
hyp_activated_during_punishment: int = 0 # ⚠️ chave: punição E ativação
|
| 80 |
+
delta_zero_when_activated: int = 0 # ⚠️ bug: HypT morta
|
| 81 |
+
l1_entropy_saturated_steps: int = 0 # ⚠️ bug: sem seleção
|
| 82 |
+
l2_zero_grad_hyp_steps: int = 0 # ⚠️ bug: loss_hyp sem grad
|
| 83 |
+
l4_T_stagnation: int = 0 # ⚠️ bug: T não move
|
| 84 |
+
l4_tau_stagnation: int = 0 # ⚠️ bug: τ não move
|
| 85 |
+
crw_converged_steps: int = 0
|
| 86 |
+
crw_contraction_invalid_steps: int = 0 # ⚠️ bug: W₂ cresceu
|
| 87 |
+
last_T: float = 1.0
|
| 88 |
+
last_tau: float = 1.0
|
| 89 |
+
history: list = field(default_factory=list)
|
| 90 |
+
|
| 91 |
+
@property
|
| 92 |
+
def has_bug(self) -> bool:
|
| 93 |
+
"""True se algum bug provável foi detectado."""
|
| 94 |
+
# Punição ocorreu mas hipótese nunca ativada — BUG CRÍTICO
|
| 95 |
+
if self.punishment_steps > 0 and self.hyp_activated_during_punishment == 0:
|
| 96 |
+
return True
|
| 97 |
+
# Hipótese ativada mas delta sempre zero — HypT morta
|
| 98 |
+
if self.hyp_activated_steps > 0 and self.hyp_activated_steps == self.delta_zero_when_activated:
|
| 99 |
+
return True
|
| 100 |
+
# L1 saturada em mais de 80% dos steps
|
| 101 |
+
if self.total_steps > 5 and self.l1_entropy_saturated_steps > 0.8 * self.total_steps:
|
| 102 |
+
return True
|
| 103 |
+
# CRW com contração inválida em mais de 50% dos steps
|
| 104 |
+
if self.crw_contraction_invalid_steps > 0.5 * max(1, self.total_steps):
|
| 105 |
+
return True
|
| 106 |
+
return False
|
| 107 |
+
|
| 108 |
+
@property
|
| 109 |
+
def bug_summary(self) -> list[str]:
|
| 110 |
+
bugs = []
|
| 111 |
+
if self.punishment_steps > 0 and self.hyp_activated_during_punishment == 0:
|
| 112 |
+
bugs.append(
|
| 113 |
+
f"CRÍTICO: {self.punishment_steps} passos de punição (loss>τ) "
|
| 114 |
+
f"mas hipótese nunca ativada — verifique tau e use_hypothesis"
|
| 115 |
+
)
|
| 116 |
+
if self.hyp_activated_steps > 0 and self.hyp_activated_steps == self.delta_zero_when_activated:
|
| 117 |
+
bugs.append(
|
| 118 |
+
f"CRÍTICO: hipótese ativada {self.hyp_activated_steps}x mas "
|
| 119 |
+
f"||delta|| sempre zero — HypT morta ou stop_grad excessivo"
|
| 120 |
+
)
|
| 121 |
+
if self.total_steps > 5 and self.l1_entropy_saturated_steps > 0.8 * self.total_steps:
|
| 122 |
+
bugs.append(
|
| 123 |
+
f"ALERTA: L1 entropy saturada em {self.l1_entropy_saturated_steps}/{self.total_steps} "
|
| 124 |
+
f"passos — module_selector não está selecionando"
|
| 125 |
+
)
|
| 126 |
+
if self.l4_T_stagnation > 3:
|
| 127 |
+
bugs.append(
|
| 128 |
+
f"ALERTA: T não mudou em {self.l4_T_stagnation} chamadas meta — "
|
| 129 |
+
f"meta_otim estagnado"
|
| 130 |
+
)
|
| 131 |
+
if self.crw_contraction_invalid_steps > 0.5 * max(1, self.total_steps):
|
| 132 |
+
bugs.append(
|
| 133 |
+
f"ALERTA: CRW contração inválida em "
|
| 134 |
+
f"{self.crw_contraction_invalid_steps}/{self.total_steps} passos — "
|
| 135 |
+
f"W₂ cresceu em vez de contrair"
|
| 136 |
+
)
|
| 137 |
+
return bugs
|
| 138 |
+
|
| 139 |
+
|
| 140 |
+
class HypothesisMonitor:
|
| 141 |
+
"""Monitora ativação dos 4 lemas + CircularReasoning durante o treino.
|
| 142 |
+
|
| 143 |
+
Uso:
|
| 144 |
+
monitor = HypothesisMonitor(max_history=100)
|
| 145 |
+
# No loop de treino, após cada micro-batch:
|
| 146 |
+
state = monitor.record(
|
| 147 |
+
step=global_step,
|
| 148 |
+
alpha=aux["alpha"], # (K,) tensor
|
| 149 |
+
entropy_reg=aux["entropy_reg"],
|
| 150 |
+
loss_main=loss_main.item(),
|
| 151 |
+
tau=model.tau.item(),
|
| 152 |
+
use_hyp=use_hyp,
|
| 153 |
+
delta=delta,
|
| 154 |
+
temperature=meta_cfg.temperature,
|
| 155 |
+
circular_info=aux.get("circular_info"),
|
| 156 |
+
)
|
| 157 |
+
# Após cada meta step:
|
| 158 |
+
monitor.record_meta(T_new, tau_new, sharpness)
|
| 159 |
+
# Ao final:
|
| 160 |
+
report = monitor.summary()
|
| 161 |
+
if report.has_bug:
|
| 162 |
+
for bug in report.bug_summary:
|
| 163 |
+
logger.error(bug)
|
| 164 |
+
"""
|
| 165 |
+
|
| 166 |
+
def __init__(
|
| 167 |
+
self,
|
| 168 |
+
max_modules: int = 8,
|
| 169 |
+
delta_epsilon: float = 1e-4,
|
| 170 |
+
meta_stagnation_epsilon: float = 1e-4,
|
| 171 |
+
history_size: int = 200,
|
| 172 |
+
):
|
| 173 |
+
self.max_modules = max_modules
|
| 174 |
+
self.delta_epsilon = delta_epsilon
|
| 175 |
+
self.meta_stagnation_epsilon = meta_stagnation_epsilon
|
| 176 |
+
self.history_size = history_size
|
| 177 |
+
self._history: deque = deque(maxlen=history_size)
|
| 178 |
+
# Acumuladores para relatório
|
| 179 |
+
self._punishment_steps = 0
|
| 180 |
+
self._hyp_activated_steps = 0
|
| 181 |
+
self._hyp_activated_during_punishment = 0
|
| 182 |
+
self._delta_zero_when_activated = 0
|
| 183 |
+
self._l1_entropy_saturated_steps = 0
|
| 184 |
+
self._l2_zero_grad_hyp_steps = 0
|
| 185 |
+
self._crw_converged_steps = 0
|
| 186 |
+
self._crw_contraction_invalid_steps = 0
|
| 187 |
+
# Meta tracking (L4)
|
| 188 |
+
self._last_T: float = 1.0
|
| 189 |
+
self._last_tau: float = 1.0
|
| 190 |
+
self._T_stagnation_count: int = 0
|
| 191 |
+
self._tau_stagnation_count: int = 0
|
| 192 |
+
self._meta_calls: int = 0
|
| 193 |
+
# L2 tracking
|
| 194 |
+
self._last_l2_applied: bool = False
|
| 195 |
+
self._last_grad_hyp_norm: float = 0.0
|
| 196 |
+
self._last_orth_proj_norm: float = 0.0
|
| 197 |
+
|
| 198 |
+
def record(
|
| 199 |
+
self,
|
| 200 |
+
step: int,
|
| 201 |
+
alpha: Optional[torch.Tensor] = None,
|
| 202 |
+
entropy_reg: float = 0.0,
|
| 203 |
+
loss_main: float = 0.0,
|
| 204 |
+
tau: float = float("inf"),
|
| 205 |
+
use_hyp: bool = False,
|
| 206 |
+
delta: Optional[torch.Tensor] = None,
|
| 207 |
+
temperature: float = 1.0,
|
| 208 |
+
circular_info: Optional[dict] = None,
|
| 209 |
+
l2_applied: bool = False,
|
| 210 |
+
grad_main_norm: float = 0.0,
|
| 211 |
+
grad_hyp_norm: float = 0.0,
|
| 212 |
+
orth_proj_norm: float = 0.0,
|
| 213 |
+
) -> LemmaState:
|
| 214 |
+
"""Registra estado de um step. Retorna LemmaState para inspeção imediata."""
|
| 215 |
+
# L1: entropy e max(alpha)
|
| 216 |
+
if alpha is not None:
|
| 217 |
+
alpha = alpha.detach().float()
|
| 218 |
+
# alpha pode ser (K,) ou (B, K) — reduzir para (K,)
|
| 219 |
+
if alpha.dim() > 1:
|
| 220 |
+
alpha = alpha.mean(dim=0)
|
| 221 |
+
alpha_entropy = float(-(alpha * torch.log(alpha + 1e-10)).sum().item())
|
| 222 |
+
alpha_max = float(alpha.max().item())
|
| 223 |
+
alpha_active = int((alpha > 0.01).sum().item())
|
| 224 |
+
# L1 ativa se entropy < 0.95 * log(K) (seleção real está acontecendo)
|
| 225 |
+
max_entropy = math.log(self.max_modules)
|
| 226 |
+
l1_activated = alpha_entropy < 0.95 * max_entropy
|
| 227 |
+
if not l1_activated:
|
| 228 |
+
self._l1_entropy_saturated_steps += 1
|
| 229 |
+
else:
|
| 230 |
+
alpha_entropy = 0.0
|
| 231 |
+
alpha_max = 0.0
|
| 232 |
+
alpha_active = 0
|
| 233 |
+
l1_activated = False
|
| 234 |
+
|
| 235 |
+
# L3: punição e delta
|
| 236 |
+
punishment_triggered = loss_main > tau
|
| 237 |
+
if punishment_triggered:
|
| 238 |
+
self._punishment_steps += 1
|
| 239 |
+
if use_hyp:
|
| 240 |
+
self._hyp_activated_steps += 1
|
| 241 |
+
if punishment_triggered:
|
| 242 |
+
self._hyp_activated_during_punishment += 1
|
| 243 |
+
# Verifica se delta é significativo
|
| 244 |
+
if delta is not None:
|
| 245 |
+
delta_norm = float(delta.detach().float().norm().item())
|
| 246 |
+
if delta_norm < self.delta_epsilon:
|
| 247 |
+
self._delta_zero_when_activated += 1
|
| 248 |
+
else:
|
| 249 |
+
delta_norm = 0.0
|
| 250 |
+
self._delta_zero_when_activated += 1
|
| 251 |
+
else:
|
| 252 |
+
delta_norm = 0.0
|
| 253 |
+
|
| 254 |
+
# L2 tracking (será preenchido por record_l2 após backward)
|
| 255 |
+
self._last_l2_applied = l2_applied
|
| 256 |
+
self._last_grad_hyp_norm = grad_hyp_norm
|
| 257 |
+
self._last_orth_proj_norm = orth_proj_norm
|
| 258 |
+
if l2_applied and grad_hyp_norm < 1e-10:
|
| 259 |
+
self._l2_zero_grad_hyp_steps += 1
|
| 260 |
+
|
| 261 |
+
# Circular Reasoning
|
| 262 |
+
crw_converged = False
|
| 263 |
+
crw_w2_final = 0.0
|
| 264 |
+
crw_q_avg = 1.0
|
| 265 |
+
crw_contraction_valid = True
|
| 266 |
+
if circular_info is not None:
|
| 267 |
+
crw_converged = bool(circular_info.get("converged", False))
|
| 268 |
+
crw_w2_final = float(circular_info.get("w2_final", 0.0))
|
| 269 |
+
crw_q_avg = float(circular_info.get("q_empirical_avg", 1.0))
|
| 270 |
+
crw_contraction_valid = bool(circular_info.get("contraction_valid", True))
|
| 271 |
+
if crw_converged:
|
| 272 |
+
self._crw_converged_steps += 1
|
| 273 |
+
if not crw_contraction_valid:
|
| 274 |
+
self._crw_contraction_invalid_steps += 1
|
| 275 |
+
|
| 276 |
+
state = LemmaState(
|
| 277 |
+
step=step,
|
| 278 |
+
alpha_entropy=alpha_entropy,
|
| 279 |
+
alpha_max=alpha_max,
|
| 280 |
+
alpha_active_modules=alpha_active,
|
| 281 |
+
temperature=temperature,
|
| 282 |
+
l1_activated=l1_activated,
|
| 283 |
+
l2_applied=l2_applied,
|
| 284 |
+
grad_main_norm=grad_main_norm,
|
| 285 |
+
grad_hyp_norm=grad_hyp_norm,
|
| 286 |
+
orth_proj_norm=orth_proj_norm,
|
| 287 |
+
loss_main=loss_main,
|
| 288 |
+
tau=tau,
|
| 289 |
+
use_hyp_activated=use_hyp,
|
| 290 |
+
delta_norm=delta_norm,
|
| 291 |
+
punishment_triggered=punishment_triggered,
|
| 292 |
+
T_new=self._last_T,
|
| 293 |
+
tau_new=self._last_tau,
|
| 294 |
+
l4_meta_called=False,
|
| 295 |
+
crw_converged=crw_converged,
|
| 296 |
+
crw_w2_final=crw_w2_final,
|
| 297 |
+
crw_q_avg=crw_q_avg,
|
| 298 |
+
crw_contraction_valid=crw_contraction_valid,
|
| 299 |
+
)
|
| 300 |
+
self._history.append(state)
|
| 301 |
+
return state
|
| 302 |
+
|
| 303 |
+
def record_meta(
|
| 304 |
+
self,
|
| 305 |
+
T_new: float,
|
| 306 |
+
tau_new: float,
|
| 307 |
+
sharpness: float = 0.0,
|
| 308 |
+
) -> None:
|
| 309 |
+
"""Registra resultado de um passo do MetaConfigurator (L4)."""
|
| 310 |
+
self._meta_calls += 1
|
| 311 |
+
# Detecta estagnação
|
| 312 |
+
if abs(T_new - self._last_T) < self.meta_stagnation_epsilon:
|
| 313 |
+
self._T_stagnation_count += 1
|
| 314 |
+
else:
|
| 315 |
+
self._T_stagnation_count = 0
|
| 316 |
+
if abs(tau_new - self._last_tau) < self.meta_stagnation_epsilon:
|
| 317 |
+
self._tau_stagnation_count += 1
|
| 318 |
+
else:
|
| 319 |
+
self._tau_stagnation_count = 0
|
| 320 |
+
self._last_T = T_new
|
| 321 |
+
self._last_tau = tau_new
|
| 322 |
+
|
| 323 |
+
def summary(self) -> HypothesisMonitorReport:
|
| 324 |
+
"""Gera relatório final com detecção de bugs."""
|
| 325 |
+
report = HypothesisMonitorReport(
|
| 326 |
+
total_steps=len(self._history),
|
| 327 |
+
punishment_steps=self._punishment_steps,
|
| 328 |
+
hyp_activated_steps=self._hyp_activated_steps,
|
| 329 |
+
hyp_activated_during_punishment=self._hyp_activated_during_punishment,
|
| 330 |
+
delta_zero_when_activated=self._delta_zero_when_activated,
|
| 331 |
+
l1_entropy_saturated_steps=self._l1_entropy_saturated_steps,
|
| 332 |
+
l2_zero_grad_hyp_steps=self._l2_zero_grad_hyp_steps,
|
| 333 |
+
l4_T_stagnation=self._T_stagnation_count,
|
| 334 |
+
l4_tau_stagnation=self._tau_stagnation_count,
|
| 335 |
+
crw_converged_steps=self._crw_converged_steps,
|
| 336 |
+
crw_contraction_invalid_steps=self._crw_contraction_invalid_steps,
|
| 337 |
+
last_T=self._last_T,
|
| 338 |
+
last_tau=self._last_tau,
|
| 339 |
+
history=[s.__dict__ for s in list(self._history)[-20:]], # últimos 20
|
| 340 |
+
)
|
| 341 |
+
return report
|
| 342 |
+
|
| 343 |
+
def log_summary(self) -> None:
|
| 344 |
+
"""Imprime relatório no log."""
|
| 345 |
+
r = self.summary()
|
| 346 |
+
logger.info("=" * 60)
|
| 347 |
+
logger.info("HIPÓTESE MONITOR — RELATÓRIO DE ATIVAÇÃO")
|
| 348 |
+
logger.info("=" * 60)
|
| 349 |
+
logger.info(f" Total steps: {r.total_steps}")
|
| 350 |
+
logger.info(f" Punição (loss > τ): {r.punishment_steps}")
|
| 351 |
+
logger.info(f" Hipótese ativada (use_hyp=True): {r.hyp_activated_steps}")
|
| 352 |
+
logger.info(f" Hipótese ativada DURANTE punição: {r.hyp_activated_during_punishment}")
|
| 353 |
+
logger.info(f" ||delta||≈0 quando ativada: {r.delta_zero_when_activated}")
|
| 354 |
+
logger.info(f" L1 entropy saturada: {r.l1_entropy_saturated_steps}")
|
| 355 |
+
logger.info(f" L2 grad_hyp≈0: {r.l2_zero_grad_hyp_steps}")
|
| 356 |
+
logger.info(f" L4 T estagnado: {r.l4_T_stagnation}")
|
| 357 |
+
logger.info(f" L4 τ estagnado: {r.l4_tau_stagnation}")
|
| 358 |
+
logger.info(f" CRW convergiu: {r.crw_converged_steps}")
|
| 359 |
+
logger.info(f" CRW contração inválida: {r.crw_contraction_invalid_steps}")
|
| 360 |
+
logger.info(f" T final: {r.last_T:.4f}")
|
| 361 |
+
logger.info(f" τ final: {r.last_tau:.4f}")
|
| 362 |
+
if r.has_bug:
|
| 363 |
+
logger.error("=" * 60)
|
| 364 |
+
logger.error("⚠️ BUGS DETECTADOS:")
|
| 365 |
+
for bug in r.bug_summary:
|
| 366 |
+
logger.error(f" - {bug}")
|
| 367 |
+
logger.error("=" * 60)
|
| 368 |
+
else:
|
| 369 |
+
logger.info(" ✓ Nenhum bug crítico detectado pelo monitor")
|
| 370 |
+
logger.info("=" * 60)
|
| 371 |
+
|
| 372 |
+
|
| 373 |
+
__all__ = ["HypothesisMonitor", "HypothesisMonitorReport", "LemmaState"]
|
src/bigru_t/training/meta_configurator.py
CHANGED
|
@@ -96,6 +96,18 @@ class MetaConfigurator:
|
|
| 96 |
|
| 97 |
Returns:
|
| 98 |
dict com: loss_val, sharpness, meta_loss, T_new, tau_new
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 99 |
"""
|
| 100 |
if criterion is None:
|
| 101 |
criterion = nn.CrossEntropyLoss()
|
|
@@ -125,9 +137,30 @@ class MetaConfigurator:
|
|
| 125 |
(g.detach() ** 2).sum() for g in grads if g is not None
|
| 126 |
)
|
| 127 |
|
| 128 |
-
#
|
| 129 |
-
#
|
| 130 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 131 |
|
| 132 |
# Atualiza log_temperature e log_tau via gradiente de meta_loss
|
| 133 |
self.meta_optim.zero_grad()
|
|
@@ -149,4 +182,8 @@ class MetaConfigurator:
|
|
| 149 |
"tau_new": float(torch.exp(self.log_tau.detach()).item()),
|
| 150 |
"grad_log_T": grad_log_T,
|
| 151 |
"grad_log_tau": grad_log_tau,
|
|
|
|
|
|
|
|
|
|
|
|
|
| 152 |
}
|
|
|
|
| 96 |
|
| 97 |
Returns:
|
| 98 |
dict com: loss_val, sharpness, meta_loss, T_new, tau_new
|
| 99 |
+
|
| 100 |
+
V3 BUGFIX: log_tau não tinha caminho de gradiente (aparecia como None).
|
| 101 |
+
Agora adicionamos um termo de target-tracking: τ deve seguir loss_val
|
| 102 |
+
(com margem), dando a log_tau um gradiente via MSE.
|
| 103 |
+
Também corrigimos module_selector para preservar o grafo quando
|
| 104 |
+
temperature é um tensor (torch.clamp em vez de max()).
|
| 105 |
+
|
| 106 |
+
V3.2 BUGFIX (complementar): log_temperature também não movia em 1 step
|
| 107 |
+
porque o gradiente via alpha = softmax(logits/T) é ≈0 quando logits ≈0
|
| 108 |
+
(init randn*1/sqrt(K) ≈ 0). Adicionamos T_target = 1 + 0.1*(loss - 1):
|
| 109 |
+
loss alta → T alto (exploração); loss baixa → T baixo (especialização).
|
| 110 |
+
Isto dá a log_T um gradiente via MSE mesmo quando alpha é uniforme.
|
| 111 |
"""
|
| 112 |
if criterion is None:
|
| 113 |
criterion = nn.CrossEntropyLoss()
|
|
|
|
| 137 |
(g.detach() ** 2).sum() for g in grads if g is not None
|
| 138 |
)
|
| 139 |
|
| 140 |
+
# V3 FIX: meta_loss agora inclui um termo para log_tau (target-tracking).
|
| 141 |
+
# τ deve seguir (loss_val + margem) para que a hipótese seja ativada
|
| 142 |
+
# quando loss > τ (modelo está "lutando"). Sem este termo, log_tau.grad
|
| 143 |
+
# era sempre None e τ nunca se movia — BUG crítico do Lema 4.
|
| 144 |
+
tau_target = loss_val.detach() + 1.0 # margem de 1.0 acima da loss
|
| 145 |
+
tau_pred = torch.exp(self.log_tau)
|
| 146 |
+
tau_tracking_loss = 0.1 * (tau_pred - tau_target) ** 2
|
| 147 |
+
|
| 148 |
+
# V3.2 FIX: T tracking loss — análogo ao tau_tracking_loss.
|
| 149 |
+
# Sem este termo, log_T.grad era ≈0 em 1 step (alpha uniforme quando
|
| 150 |
+
# logits ≈ 0), e T nunca movia no bug_hunt. T_target baseado em loss:
|
| 151 |
+
# loss=1 → T_target=1.0 (neutro)
|
| 152 |
+
# loss=10 → T_target=1.9 (exploração: alpha mais uniforme)
|
| 153 |
+
# loss=0.1→ T_target=0.99 (especialização: alpha mais peaked)
|
| 154 |
+
# Peso menor (0.05) que tau (0.1) para evitar oscilação.
|
| 155 |
+
loss_for_target = loss_val.detach().clamp(min=0.0, max=20.0)
|
| 156 |
+
T_target = 1.0 + 0.1 * (loss_for_target - 1.0)
|
| 157 |
+
T_pred = torch.exp(self.log_temperature)
|
| 158 |
+
T_tracking_loss = 0.05 * (T_pred - T_target) ** 2
|
| 159 |
+
|
| 160 |
+
# Meta-perda: loss_val (atualiza log_T via temperatura → alpha → y_hat)
|
| 161 |
+
# + T_tracking_loss (atualiza log_T via MSE target)
|
| 162 |
+
# + tau_tracking_loss (atualiza log_tau)
|
| 163 |
+
meta_loss = loss_val + T_tracking_loss + tau_tracking_loss
|
| 164 |
|
| 165 |
# Atualiza log_temperature e log_tau via gradiente de meta_loss
|
| 166 |
self.meta_optim.zero_grad()
|
|
|
|
| 182 |
"tau_new": float(torch.exp(self.log_tau.detach()).item()),
|
| 183 |
"grad_log_T": grad_log_T,
|
| 184 |
"grad_log_tau": grad_log_tau,
|
| 185 |
+
"tau_target": float(tau_target.item()),
|
| 186 |
+
"T_target": float(T_target.item()),
|
| 187 |
+
"T_tracking_loss": float(T_tracking_loss.item()),
|
| 188 |
+
"tau_tracking_loss": float(tau_tracking_loss.item()),
|
| 189 |
}
|
src/bigru_t/training/trainer.py
CHANGED
|
@@ -43,7 +43,8 @@ from ..model.unified_model import UnifiedModel, UnifiedModelConfig, create_unifi
|
|
| 43 |
from .gradient_surgery import apply_gradient_surgery
|
| 44 |
from .meta_configurator import MetaConfigurator
|
| 45 |
from .kill_switch import KillSwitch
|
| 46 |
-
from .dpo import compute_dynamic_beta
|
|
|
|
| 47 |
from ..optim.hamiltonian_wasserstein import HamiltonianWassersteinOptimizer
|
| 48 |
from ..utils.memory_cleanup import aggressive_cleanup, production_cleanup, TimeBudget, StepTimer, get_rss_mb
|
| 49 |
|
|
@@ -52,9 +53,9 @@ logger = logging.getLogger(__name__)
|
|
| 52 |
|
| 53 |
@dataclass
|
| 54 |
class TrainerConfig:
|
| 55 |
-
"""Configuração do treino (
|
| 56 |
-
# Epochs (
|
| 57 |
-
epochs: int =
|
| 58 |
|
| 59 |
# Datasets
|
| 60 |
datasets: str = "CEIA-POSITIVO/ultrachat_br_clustred_balanced_v1,Madras1/corpus-ptbr-v2,dominguesm/restore-punctuation-ptbr-dataset"
|
|
@@ -78,11 +79,20 @@ class TrainerConfig:
|
|
| 78 |
hw_prune_every_n: int = 0 # pruning desativado por padrão (0 = off)
|
| 79 |
hw_ref_buffer_size: int = 2 # buffer de referência reduzido (era 10 → 2 para economizar RAM)
|
| 80 |
|
| 81 |
-
# DPO (
|
| 82 |
-
|
|
|
|
|
|
|
| 83 |
dpo_beta_min: float = 0.05
|
| 84 |
dpo_beta_max: float = 0.5
|
| 85 |
dpo_warmup_steps: int = 100
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 86 |
|
| 87 |
# Meta-configurator (Lema 4)
|
| 88 |
meta_interval: int = 10 # a cada N micro-batches
|
|
@@ -99,11 +109,12 @@ class TrainerConfig:
|
|
| 99 |
loss_patience: int = 30
|
| 100 |
|
| 101 |
# Time budget (estilo Xavante — streaming com timed steps)
|
| 102 |
-
|
| 103 |
-
|
|
|
|
| 104 |
|
| 105 |
# Memory cleanup agressivo (estilo Xavante)
|
| 106 |
-
cleanup_every_n_steps: int =
|
| 107 |
|
| 108 |
# Logging
|
| 109 |
log_every: int = 5
|
|
@@ -196,6 +207,26 @@ class BiGRU_T_Trainer:
|
|
| 196 |
log_dir=config.temp_dir,
|
| 197 |
)
|
| 198 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 199 |
# Diretórios
|
| 200 |
Path(config.output_dir).mkdir(parents=True, exist_ok=True)
|
| 201 |
Path(config.temp_dir).mkdir(parents=True, exist_ok=True)
|
|
@@ -206,6 +237,8 @@ class BiGRU_T_Trainer:
|
|
| 206 |
self.best_loss = float("inf")
|
| 207 |
self.losses_history: list[float] = []
|
| 208 |
self.killed_reason: Optional[str] = None
|
|
|
|
|
|
|
| 209 |
|
| 210 |
def _collate(self, samples: list) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
| 211 |
"""Tokeniza samples → (input_ids, attention_mask, labels).
|
|
@@ -352,13 +385,14 @@ class BiGRU_T_Trainer:
|
|
| 352 |
# Target = último token
|
| 353 |
target = input_ids[:, -1].clone()
|
| 354 |
|
| 355 |
-
# SINGLE forward: retorna y_hat, delta, e aux (entropy_reg, alpha)
|
| 356 |
y_hat, delta, aux = self.model(
|
| 357 |
input_ids,
|
| 358 |
temperature=self.meta_cfg.temperature,
|
| 359 |
use_hypothesis=True, # sempre computa delta (barato se stop_grad)
|
| 360 |
stop_grad_hyp=self.config.stop_grad_hyp,
|
| 361 |
return_aux=True,
|
|
|
|
| 362 |
)
|
| 363 |
|
| 364 |
# Loss principal: CrossEntropy sobre y_hat + entropy_reg (Lema 1)
|
|
@@ -385,6 +419,21 @@ class BiGRU_T_Trainer:
|
|
| 385 |
else:
|
| 386 |
loss_hyp = torch.zeros((), device=y_hat.device)
|
| 387 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 388 |
except RuntimeError as e:
|
| 389 |
err = str(e).lower()
|
| 390 |
if "out of memory" in err:
|
|
@@ -397,8 +446,8 @@ class BiGRU_T_Trainer:
|
|
| 397 |
traceback.print_exc()
|
| 398 |
break
|
| 399 |
|
| 400 |
-
# DPO beta dinâmico (se DPO ativo)
|
| 401 |
-
if
|
| 402 |
beta_t = compute_dynamic_beta(
|
| 403 |
self.global_step,
|
| 404 |
warmup_steps=cfg.dpo_warmup_steps,
|
|
@@ -410,9 +459,24 @@ class BiGRU_T_Trainer:
|
|
| 410 |
|
| 411 |
# Backward + step
|
| 412 |
self.step_timer.start()
|
|
|
|
|
|
|
|
|
|
| 413 |
if use_hyp and loss_hyp.requires_grad:
|
| 414 |
# Lema 2: gradient surgery
|
| 415 |
apply_gradient_surgery(self.model, loss_main_total, loss_hyp)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 416 |
# Clip
|
| 417 |
torch.nn.utils.clip_grad_norm_(self.model.parameters(), cfg.max_grad_norm)
|
| 418 |
# HW optimizer: step aceita loss para logging
|
|
@@ -432,6 +496,23 @@ class BiGRU_T_Trainer:
|
|
| 432 |
self.optimizer.step()
|
| 433 |
self.step_timer.stop(step_label=str(self.global_step))
|
| 434 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 435 |
self.micro_batch_idx += 1
|
| 436 |
|
| 437 |
# Log
|
|
@@ -464,12 +545,25 @@ class BiGRU_T_Trainer:
|
|
| 464 |
1, min(cfg.grad_accum, len(self.losses_history))
|
| 465 |
)
|
| 466 |
ppl = math.exp(min(20, avg_loss)) if avg_loss < 20 else float("inf")
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 467 |
logger.info(
|
| 468 |
f" step {self.global_step} | epoch {epoch+1} | "
|
| 469 |
-
f"loss={avg_loss:.4f} ppl={ppl:.2f} | "
|
| 470 |
f"RAM {state.ram_pct:.1f}% disk {state.disk_free_gb:.1f}GB | "
|
| 471 |
f"modules={active} | T={self.meta_cfg.temperature:.3f} "
|
| 472 |
f"τ={self.meta_cfg.tau:.3f} | "
|
|
|
|
|
|
|
|
|
|
| 473 |
f"elapsed={elapsed:.0f}s"
|
| 474 |
)
|
| 475 |
|
|
@@ -511,7 +605,10 @@ class BiGRU_T_Trainer:
|
|
| 511 |
partial_result = {
|
| 512 |
"epochs_completed": epoch + 1,
|
| 513 |
"killed": False,
|
|
|
|
|
|
|
| 514 |
"global_step": self.global_step,
|
|
|
|
| 515 |
"best_loss": self.best_loss,
|
| 516 |
"final_loss": epoch_losses[-1] if epoch_losses else float("nan"),
|
| 517 |
"final_avg_loss": epoch_avg,
|
|
@@ -555,6 +652,26 @@ class BiGRU_T_Trainer:
|
|
| 555 |
"step_timer": self.step_timer.summary(),
|
| 556 |
}
|
| 557 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 558 |
logger.info("=" * 70)
|
| 559 |
logger.info("TREINO CONCLUÍDO")
|
| 560 |
logger.info(f" epochs completed: {result['epochs_completed']}")
|
|
@@ -568,6 +685,9 @@ class BiGRU_T_Trainer:
|
|
| 568 |
logger.info(f" elapsed: {result['elapsed_s']:.1f}s")
|
| 569 |
logger.info(f" max RAM: {ks_summary.get('max_ram_pct', 0):.1f}%")
|
| 570 |
logger.info(f" min disk free: {ks_summary.get('min_disk_free_gb', 0):.2f}GB")
|
|
|
|
|
|
|
|
|
|
| 571 |
logger.info("=" * 70)
|
| 572 |
|
| 573 |
# Salva modelo final
|
|
@@ -580,6 +700,90 @@ class BiGRU_T_Trainer:
|
|
| 580 |
|
| 581 |
return result
|
| 582 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 583 |
def _run_meta_update(self):
|
| 584 |
"""Executa um passo do MetaConfigurator (Lema 4)."""
|
| 585 |
# Pega um batch de validação
|
|
@@ -592,11 +796,19 @@ class BiGRU_T_Trainer:
|
|
| 592 |
# Target = último token (mesma lógica de _forward_loss)
|
| 593 |
target = input_ids[:, -1].clone() # (batch,)
|
| 594 |
meta_result = self.meta_cfg.forward_with_meta(input_ids, target)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 595 |
logger.info(
|
| 596 |
f" META step {self.global_step}: "
|
| 597 |
f"val_loss={meta_result['loss_val']:.4f} "
|
| 598 |
f"sharpness={meta_result['sharpness']:.4f} "
|
| 599 |
-
f"T={meta_result['T_new']:.4f} τ={meta_result['tau_new']:.4f}"
|
|
|
|
|
|
|
| 600 |
)
|
| 601 |
except Exception as e:
|
| 602 |
logger.warning(f"Meta update failed: {e}")
|
|
|
|
| 43 |
from .gradient_surgery import apply_gradient_surgery
|
| 44 |
from .meta_configurator import MetaConfigurator
|
| 45 |
from .kill_switch import KillSwitch
|
| 46 |
+
from .dpo import compute_dynamic_beta, dpo_step, compute_sequence_logps
|
| 47 |
+
from .hypothesis_monitor import HypothesisMonitor
|
| 48 |
from ..optim.hamiltonian_wasserstein import HamiltonianWassersteinOptimizer
|
| 49 |
from ..utils.memory_cleanup import aggressive_cleanup, production_cleanup, TimeBudget, StepTimer, get_rss_mb
|
| 50 |
|
|
|
|
| 53 |
|
| 54 |
@dataclass
|
| 55 |
class TrainerConfig:
|
| 56 |
+
"""Configuração do treino (V3: 5 épocas + DPO + CircularReasoning + Monitor)."""
|
| 57 |
+
# Epochs (V3: 5 épocas conforme especificação do usuário)
|
| 58 |
+
epochs: int = 5
|
| 59 |
|
| 60 |
# Datasets
|
| 61 |
datasets: str = "CEIA-POSITIVO/ultrachat_br_clustred_balanced_v1,Madras1/corpus-ptbr-v2,dominguesm/restore-punctuation-ptbr-dataset"
|
|
|
|
| 79 |
hw_prune_every_n: int = 0 # pruning desativado por padrão (0 = off)
|
| 80 |
hw_ref_buffer_size: int = 2 # buffer de referência reduzido (era 10 → 2 para economizar RAM)
|
| 81 |
|
| 82 |
+
# DPO (V3: ATIVADO para datasets compatíveis)
|
| 83 |
+
# DPO é ativado automaticamente quando amostras têm label_text
|
| 84 |
+
# (instruction/response format). DPO loss é adicionada à loss principal.
|
| 85 |
+
use_dpo: bool = True # V3: ativo por padrão (auto-detecção)
|
| 86 |
dpo_beta_min: float = 0.05
|
| 87 |
dpo_beta_max: float = 0.5
|
| 88 |
dpo_warmup_steps: int = 100
|
| 89 |
+
dpo_loss_weight: float = 0.3 # peso da loss DPO na loss total
|
| 90 |
+
dpo_use_ipo: bool = False # IPO em vez de DPO clássico
|
| 91 |
+
dpo_label_smoothing: float = 0.0
|
| 92 |
+
dpo_every_n_steps: int = 2 # computar DPO a cada N steps (economia de compute)
|
| 93 |
+
|
| 94 |
+
# Circular Reasoning (Teorema 4)
|
| 95 |
+
use_circular_reasoning: bool = True # já está no UnifiedModelConfig também
|
| 96 |
|
| 97 |
# Meta-configurator (Lema 4)
|
| 98 |
meta_interval: int = 10 # a cada N micro-batches
|
|
|
|
| 109 |
loss_patience: int = 30
|
| 110 |
|
| 111 |
# Time budget (estilo Xavante — streaming com timed steps)
|
| 112 |
+
# V3: aumentado para 5 épocas (2700s total = 45 min, 540s/epoch = 9 min)
|
| 113 |
+
max_total_time_s: float = 2700.0 # 45 min máx total (5 épocas)
|
| 114 |
+
max_per_epoch_s: float = 540.0 # 9 min máx por época
|
| 115 |
|
| 116 |
# Memory cleanup agressivo (estilo Xavante)
|
| 117 |
+
cleanup_every_n_steps: int = 5 # aggressive_cleanup a cada N optimizer steps
|
| 118 |
|
| 119 |
# Logging
|
| 120 |
log_every: int = 5
|
|
|
|
| 207 |
log_dir=config.temp_dir,
|
| 208 |
)
|
| 209 |
|
| 210 |
+
# Hypothesis Monitor (V3) — detecta bugs na ativação dos 4 lemas
|
| 211 |
+
self.hyp_monitor = HypothesisMonitor(
|
| 212 |
+
max_modules=model.config.max_modules,
|
| 213 |
+
delta_epsilon=1e-4,
|
| 214 |
+
meta_stagnation_epsilon=1e-4,
|
| 215 |
+
history_size=200,
|
| 216 |
+
)
|
| 217 |
+
|
| 218 |
+
# DPO: detecta se dataset é compatível (tem label_text em pelo menos 30% das amostras)
|
| 219 |
+
self._dpo_compatible = False
|
| 220 |
+
if config.use_dpo and train_samples:
|
| 221 |
+
n_with_label = sum(1 for s in train_samples if getattr(s, "label_text", ""))
|
| 222 |
+
ratio = n_with_label / len(train_samples)
|
| 223 |
+
self._dpo_compatible = ratio >= 0.3
|
| 224 |
+
logger.info(
|
| 225 |
+
f"DPO auto-detecção: {n_with_label}/{len(train_samples)} ({ratio:.0%}) "
|
| 226 |
+
f"amostras com label_text → "
|
| 227 |
+
f"{'COMPATÍVEL (DPO ativado)' if self._dpo_compatible else 'não compatível (SFT only)'}"
|
| 228 |
+
)
|
| 229 |
+
|
| 230 |
# Diretórios
|
| 231 |
Path(config.output_dir).mkdir(parents=True, exist_ok=True)
|
| 232 |
Path(config.temp_dir).mkdir(parents=True, exist_ok=True)
|
|
|
|
| 237 |
self.best_loss = float("inf")
|
| 238 |
self.losses_history: list[float] = []
|
| 239 |
self.killed_reason: Optional[str] = None
|
| 240 |
+
# DPO reference logps (computados uma vez no início, com modelo congelado)
|
| 241 |
+
self._ref_logps_computed = False
|
| 242 |
|
| 243 |
def _collate(self, samples: list) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
| 244 |
"""Tokeniza samples → (input_ids, attention_mask, labels).
|
|
|
|
| 385 |
# Target = último token
|
| 386 |
target = input_ids[:, -1].clone()
|
| 387 |
|
| 388 |
+
# SINGLE forward: retorna y_hat, delta, e aux (entropy_reg, alpha, circular_info)
|
| 389 |
y_hat, delta, aux = self.model(
|
| 390 |
input_ids,
|
| 391 |
temperature=self.meta_cfg.temperature,
|
| 392 |
use_hypothesis=True, # sempre computa delta (barato se stop_grad)
|
| 393 |
stop_grad_hyp=self.config.stop_grad_hyp,
|
| 394 |
return_aux=True,
|
| 395 |
+
use_circular=cfg.use_circular_reasoning,
|
| 396 |
)
|
| 397 |
|
| 398 |
# Loss principal: CrossEntropy sobre y_hat + entropy_reg (Lema 1)
|
|
|
|
| 419 |
else:
|
| 420 |
loss_hyp = torch.zeros((), device=y_hat.device)
|
| 421 |
|
| 422 |
+
# V3: DPO loss — se dataset é compatível e estamos num step de DPO
|
| 423 |
+
loss_dpo_val = 0.0
|
| 424 |
+
if (
|
| 425 |
+
self._dpo_compatible
|
| 426 |
+
and self.global_step % cfg.dpo_every_n_steps == 0
|
| 427 |
+
and len(batch_samples) > 0
|
| 428 |
+
):
|
| 429 |
+
try:
|
| 430 |
+
loss_dpo = self._compute_dpo_loss(batch_samples, y_hat, target)
|
| 431 |
+
if loss_dpo is not None:
|
| 432 |
+
loss_main_total = loss_main_total + cfg.dpo_loss_weight * loss_dpo
|
| 433 |
+
loss_dpo_val = float(loss_dpo.item())
|
| 434 |
+
except Exception as dpo_err:
|
| 435 |
+
logger.debug(f"DPO step falhou (não crítico): {dpo_err}")
|
| 436 |
+
|
| 437 |
except RuntimeError as e:
|
| 438 |
err = str(e).lower()
|
| 439 |
if "out of memory" in err:
|
|
|
|
| 446 |
traceback.print_exc()
|
| 447 |
break
|
| 448 |
|
| 449 |
+
# DPO beta dinâmico (se DPO ativo) — para logging no optimizer
|
| 450 |
+
if self._dpo_compatible:
|
| 451 |
beta_t = compute_dynamic_beta(
|
| 452 |
self.global_step,
|
| 453 |
warmup_steps=cfg.dpo_warmup_steps,
|
|
|
|
| 459 |
|
| 460 |
# Backward + step
|
| 461 |
self.step_timer.start()
|
| 462 |
+
l2_applied = False
|
| 463 |
+
grad_hyp_norm = 0.0
|
| 464 |
+
orth_proj_norm = 0.0
|
| 465 |
if use_hyp and loss_hyp.requires_grad:
|
| 466 |
# Lema 2: gradient surgery
|
| 467 |
apply_gradient_surgery(self.model, loss_main_total, loss_hyp)
|
| 468 |
+
l2_applied = True
|
| 469 |
+
# Mede gradientes para o monitor (após surgery)
|
| 470 |
+
try:
|
| 471 |
+
grad_norms = [
|
| 472 |
+
p.grad.norm().item()
|
| 473 |
+
for p in self.model.parameters()
|
| 474 |
+
if p.grad is not None
|
| 475 |
+
]
|
| 476 |
+
if grad_norms:
|
| 477 |
+
grad_hyp_norm = sum(n ** 2 for n in grad_norms) ** 0.5
|
| 478 |
+
except Exception:
|
| 479 |
+
pass
|
| 480 |
# Clip
|
| 481 |
torch.nn.utils.clip_grad_norm_(self.model.parameters(), cfg.max_grad_norm)
|
| 482 |
# HW optimizer: step aceita loss para logging
|
|
|
|
| 496 |
self.optimizer.step()
|
| 497 |
self.step_timer.stop(step_label=str(self.global_step))
|
| 498 |
|
| 499 |
+
# V3: Hypothesis Monitor — registra estado do step
|
| 500 |
+
self.hyp_monitor.record(
|
| 501 |
+
step=self.global_step,
|
| 502 |
+
alpha=aux.get("alpha"),
|
| 503 |
+
entropy_reg=float(aux.get("entropy_reg", 0).item()
|
| 504 |
+
if hasattr(aux.get("entropy_reg", 0), "item") else 0),
|
| 505 |
+
loss_main=float(loss_main.item()),
|
| 506 |
+
tau=float(self.meta_cfg.tau),
|
| 507 |
+
use_hyp=use_hyp,
|
| 508 |
+
delta=delta if use_hyp else None,
|
| 509 |
+
temperature=self.meta_cfg.temperature,
|
| 510 |
+
circular_info=aux.get("circular_info"),
|
| 511 |
+
l2_applied=l2_applied,
|
| 512 |
+
grad_hyp_norm=grad_hyp_norm,
|
| 513 |
+
orth_proj_norm=orth_proj_norm,
|
| 514 |
+
)
|
| 515 |
+
|
| 516 |
self.micro_batch_idx += 1
|
| 517 |
|
| 518 |
# Log
|
|
|
|
| 545 |
1, min(cfg.grad_accum, len(self.losses_history))
|
| 546 |
)
|
| 547 |
ppl = math.exp(min(20, avg_loss)) if avg_loss < 20 else float("inf")
|
| 548 |
+
# V3: info do Circular Reasoning (último step registrado)
|
| 549 |
+
crw_str = ""
|
| 550 |
+
last_state = self.hyp_monitor._history[-1] if self.hyp_monitor._history else None
|
| 551 |
+
if last_state and last_state.crw_w2_final > 0:
|
| 552 |
+
crw_str = (
|
| 553 |
+
f" | CRW: w2={last_state.crw_w2_final:.4f} "
|
| 554 |
+
f"q̄={last_state.crw_q_avg:.3f} "
|
| 555 |
+
f"{'✓' if last_state.crw_contraction_valid else '⚠'}"
|
| 556 |
+
)
|
| 557 |
+
dpo_str = f" dpo={loss_dpo_val:.3f}" if loss_dpo_val > 0 else ""
|
| 558 |
logger.info(
|
| 559 |
f" step {self.global_step} | epoch {epoch+1} | "
|
| 560 |
+
f"loss={avg_loss:.4f} ppl={ppl:.2f}{dpo_str} | "
|
| 561 |
f"RAM {state.ram_pct:.1f}% disk {state.disk_free_gb:.1f}GB | "
|
| 562 |
f"modules={active} | T={self.meta_cfg.temperature:.3f} "
|
| 563 |
f"τ={self.meta_cfg.tau:.3f} | "
|
| 564 |
+
f"α_ent={last_state.alpha_entropy:.3f} "
|
| 565 |
+
f"α_max={last_state.alpha_max:.3f}"
|
| 566 |
+
f"{crw_str} | "
|
| 567 |
f"elapsed={elapsed:.0f}s"
|
| 568 |
)
|
| 569 |
|
|
|
|
| 605 |
partial_result = {
|
| 606 |
"epochs_completed": epoch + 1,
|
| 607 |
"killed": False,
|
| 608 |
+
"kill_reason": None, # V3 FIX: era missing → KeyError no _save_final_model
|
| 609 |
+
"global_reason": None,
|
| 610 |
"global_step": self.global_step,
|
| 611 |
+
"micro_batches": self.micro_batch_idx,
|
| 612 |
"best_loss": self.best_loss,
|
| 613 |
"final_loss": epoch_losses[-1] if epoch_losses else float("nan"),
|
| 614 |
"final_avg_loss": epoch_avg,
|
|
|
|
| 652 |
"step_timer": self.step_timer.summary(),
|
| 653 |
}
|
| 654 |
|
| 655 |
+
# V3: Log final do Hypothesis Monitor (detecção de bugs)
|
| 656 |
+
self.hyp_monitor.log_summary()
|
| 657 |
+
hyp_report = self.hyp_monitor.summary()
|
| 658 |
+
result["hypothesis_monitor"] = {
|
| 659 |
+
"total_steps": hyp_report.total_steps,
|
| 660 |
+
"punishment_steps": hyp_report.punishment_steps,
|
| 661 |
+
"hyp_activated_steps": hyp_report.hyp_activated_steps,
|
| 662 |
+
"hyp_activated_during_punishment": hyp_report.hyp_activated_during_punishment,
|
| 663 |
+
"delta_zero_when_activated": hyp_report.delta_zero_when_activated,
|
| 664 |
+
"l1_entropy_saturated_steps": hyp_report.l1_entropy_saturated_steps,
|
| 665 |
+
"l2_zero_grad_hyp_steps": hyp_report.l2_zero_grad_hyp_steps,
|
| 666 |
+
"crw_converged_steps": hyp_report.crw_converged_steps,
|
| 667 |
+
"crw_contraction_invalid_steps": hyp_report.crw_contraction_invalid_steps,
|
| 668 |
+
"last_T": hyp_report.last_T,
|
| 669 |
+
"last_tau": hyp_report.last_tau,
|
| 670 |
+
"has_bug": hyp_report.has_bug,
|
| 671 |
+
"bugs": hyp_report.bug_summary,
|
| 672 |
+
}
|
| 673 |
+
result["dpo_activated"] = self._dpo_compatible
|
| 674 |
+
|
| 675 |
logger.info("=" * 70)
|
| 676 |
logger.info("TREINO CONCLUÍDO")
|
| 677 |
logger.info(f" epochs completed: {result['epochs_completed']}")
|
|
|
|
| 685 |
logger.info(f" elapsed: {result['elapsed_s']:.1f}s")
|
| 686 |
logger.info(f" max RAM: {ks_summary.get('max_ram_pct', 0):.1f}%")
|
| 687 |
logger.info(f" min disk free: {ks_summary.get('min_disk_free_gb', 0):.2f}GB")
|
| 688 |
+
logger.info(f" DPO ativado: {result['dpo_activated']}")
|
| 689 |
+
logger.info(f" CircularReason: {cfg.use_circular_reasoning}")
|
| 690 |
+
logger.info(f" Hipótese monitor: {'⚠️ BUGS' if hyp_report.has_bug else '✓ OK'}")
|
| 691 |
logger.info("=" * 70)
|
| 692 |
|
| 693 |
# Salva modelo final
|
|
|
|
| 700 |
|
| 701 |
return result
|
| 702 |
|
| 703 |
+
def _compute_dpo_loss(
|
| 704 |
+
self,
|
| 705 |
+
batch_samples: list,
|
| 706 |
+
y_hat: torch.Tensor,
|
| 707 |
+
target: torch.Tensor,
|
| 708 |
+
) -> Optional[torch.Tensor]:
|
| 709 |
+
"""V3: Computa DPO loss para o batch.
|
| 710 |
+
|
| 711 |
+
Estratégia: para datasets compatíveis (com label_text = resposta correta),
|
| 712 |
+
sintetiza pares de preferência:
|
| 713 |
+
- chosen = resposta correta (label_text do próprio sample)
|
| 714 |
+
- rejected = resposta de OUTRO sample (mistura) — proxy para "resposta errada"
|
| 715 |
+
|
| 716 |
+
Como o modelo produz apenas 1 logit por amostra (não LM completo),
|
| 717 |
+
usamos uma aproximação DPO simplificada:
|
| 718 |
+
- policy_chosen_logps = log π(target | y_hat) (log-softmax do target)
|
| 719 |
+
- policy_rejected_logps = log π(rejected_token | y_hat) (de outro sample)
|
| 720 |
+
- reference = zeros (modelo congelado no início = policy inicial)
|
| 721 |
+
|
| 722 |
+
Isto NÃO é DPO canônico (que requer sequence logps), mas captura a
|
| 723 |
+
essência: punir o modelo por preferir "resposta errada" e recompensar
|
| 724 |
+
por preferir "resposta certa". Adequado para bug-detection.
|
| 725 |
+
|
| 726 |
+
Args:
|
| 727 |
+
batch_samples: lista de ProcessedSample
|
| 728 |
+
y_hat: (B, vocab) logits do modelo para o batch
|
| 729 |
+
target: (B,) token alvo (último token da sequência)
|
| 730 |
+
|
| 731 |
+
Returns:
|
| 732 |
+
loss_dpo (escalar com grad) ou None se não aplicável
|
| 733 |
+
"""
|
| 734 |
+
if not batch_samples or not hasattr(batch_samples[0], "label_text"):
|
| 735 |
+
return None
|
| 736 |
+
|
| 737 |
+
B = y_hat.size(0)
|
| 738 |
+
if B < 1:
|
| 739 |
+
return None
|
| 740 |
+
|
| 741 |
+
# Policy log-probs (com grad)
|
| 742 |
+
log_probs = F.log_softmax(y_hat, dim=-1) # (B, vocab)
|
| 743 |
+
|
| 744 |
+
# chosen = target do próprio sample
|
| 745 |
+
chosen_logps = log_probs.gather(
|
| 746 |
+
1, target.clamp(0, log_probs.size(-1) - 1).unsqueeze(-1)
|
| 747 |
+
).squeeze(-1) # (B,)
|
| 748 |
+
|
| 749 |
+
# rejected = target de OUTRO sample (roll-by-1) — proxy para "resposta errada"
|
| 750 |
+
rejected_target = target.roll(1, dims=0) # shift circular
|
| 751 |
+
# Se B == 1, não há outro sample — usa token aleatório como rejected
|
| 752 |
+
if B == 1:
|
| 753 |
+
rejected_target = torch.randint(
|
| 754 |
+
0, log_probs.size(-1), (1,), dtype=target.dtype
|
| 755 |
+
)
|
| 756 |
+
rejected_logps = log_probs.gather(
|
| 757 |
+
1, rejected_target.clamp(0, log_probs.size(-1) - 1).unsqueeze(-1)
|
| 758 |
+
).squeeze(-1) # (B,)
|
| 759 |
+
|
| 760 |
+
# Reference log-probs (sem grad) — política inicial = uniforme
|
| 761 |
+
# log π_ref(token) = -log(vocab_size)
|
| 762 |
+
vocab_size = log_probs.size(-1)
|
| 763 |
+
ref_chosen_logps = torch.full_like(chosen_logps, -math.log(vocab_size))
|
| 764 |
+
ref_rejected_logps = torch.full_like(rejected_logps, -math.log(vocab_size))
|
| 765 |
+
|
| 766 |
+
# β dinâmico
|
| 767 |
+
beta_t = compute_dynamic_beta(
|
| 768 |
+
self.global_step,
|
| 769 |
+
warmup_steps=self.config.dpo_warmup_steps,
|
| 770 |
+
beta_min=self.config.dpo_beta_min,
|
| 771 |
+
beta_max=self.config.dpo_beta_max,
|
| 772 |
+
)
|
| 773 |
+
|
| 774 |
+
# DPO loss
|
| 775 |
+
from .dpo import dpo_loss
|
| 776 |
+
loss = dpo_loss(
|
| 777 |
+
policy_chosen_logps=chosen_logps,
|
| 778 |
+
policy_rejected_logps=rejected_logps,
|
| 779 |
+
reference_chosen_logps=ref_chosen_logps,
|
| 780 |
+
reference_rejected_logps=ref_rejected_logps,
|
| 781 |
+
beta=beta_t,
|
| 782 |
+
label_smoothing=self.config.dpo_label_smoothing,
|
| 783 |
+
use_ipo=self.config.dpo_use_ipo,
|
| 784 |
+
)
|
| 785 |
+
return loss
|
| 786 |
+
|
| 787 |
def _run_meta_update(self):
|
| 788 |
"""Executa um passo do MetaConfigurator (Lema 4)."""
|
| 789 |
# Pega um batch de validação
|
|
|
|
| 796 |
# Target = último token (mesma lógica de _forward_loss)
|
| 797 |
target = input_ids[:, -1].clone() # (batch,)
|
| 798 |
meta_result = self.meta_cfg.forward_with_meta(input_ids, target)
|
| 799 |
+
# V3: registra no monitor para detectar estagnação do L4
|
| 800 |
+
self.hyp_monitor.record_meta(
|
| 801 |
+
T_new=meta_result["T_new"],
|
| 802 |
+
tau_new=meta_result["tau_new"],
|
| 803 |
+
sharpness=meta_result["sharpness"],
|
| 804 |
+
)
|
| 805 |
logger.info(
|
| 806 |
f" META step {self.global_step}: "
|
| 807 |
f"val_loss={meta_result['loss_val']:.4f} "
|
| 808 |
f"sharpness={meta_result['sharpness']:.4f} "
|
| 809 |
+
f"T={meta_result['T_new']:.4f} τ={meta_result['tau_new']:.4f} "
|
| 810 |
+
f"ΔT={meta_result['grad_log_T']:.6f} "
|
| 811 |
+
f"Δτ={meta_result['grad_log_tau']:.6f}"
|
| 812 |
)
|
| 813 |
except Exception as e:
|
| 814 |
logger.warning(f"Meta update failed: {e}")
|
training_report.json
CHANGED
|
@@ -1,76 +1,75 @@
|
|
| 1 |
{
|
| 2 |
-
"version": "
|
| 3 |
"training_params": {
|
| 4 |
"epochs": 2,
|
| 5 |
"per_device_batch_size": 1,
|
| 6 |
"grad_accum": 2,
|
| 7 |
"lr": 0.001,
|
| 8 |
-
"
|
| 9 |
-
"
|
| 10 |
-
"hw_lr_freq": 0.01,
|
| 11 |
-
"hw_ref_buffer_size": 2,
|
| 12 |
"max_seq_len": 32,
|
| 13 |
-
"max_modules": 8,
|
| 14 |
-
"d_model": 128,
|
| 15 |
"use_hypothesis": true,
|
| 16 |
"stop_grad_hyp": true,
|
| 17 |
"meta_interval": 5,
|
| 18 |
-
"
|
| 19 |
-
"
|
| 20 |
},
|
| 21 |
"results": {
|
| 22 |
"epochs_completed": 2,
|
| 23 |
-
"
|
| 24 |
-
"
|
| 25 |
-
"
|
| 26 |
-
"
|
| 27 |
-
"
|
| 28 |
-
"
|
| 29 |
-
"
|
| 30 |
-
"
|
| 31 |
-
|
| 32 |
-
|
| 33 |
-
|
| 34 |
-
|
| 35 |
-
|
| 36 |
-
|
| 37 |
-
|
| 38 |
-
|
| 39 |
-
|
| 40 |
-
|
| 41 |
-
|
| 42 |
-
|
| 43 |
-
296,
|
| 44 |
-
853
|
| 45 |
-
],
|
| 46 |
-
"meta_step_5": {
|
| 47 |
-
"val_loss": 9.41,
|
| 48 |
-
"sharpness": 157.21,
|
| 49 |
-
"T": 1.01,
|
| 50 |
-
"tau": 1.0
|
| 51 |
},
|
| 52 |
-
"
|
| 53 |
-
|
| 54 |
-
|
| 55 |
-
|
| 56 |
-
|
| 57 |
-
|
| 58 |
-
|
| 59 |
-
|
| 60 |
-
|
| 61 |
-
"
|
| 62 |
-
|
| 63 |
-
|
| 64 |
-
|
| 65 |
-
|
| 66 |
-
|
| 67 |
-
|
| 68 |
-
"
|
| 69 |
-
|
| 70 |
-
|
| 71 |
-
|
| 72 |
-
|
| 73 |
-
|
| 74 |
-
|
| 75 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 76 |
}
|
|
|
|
| 1 |
{
|
| 2 |
+
"version": "BiGRU_T_version",
|
| 3 |
"training_params": {
|
| 4 |
"epochs": 2,
|
| 5 |
"per_device_batch_size": 1,
|
| 6 |
"grad_accum": 2,
|
| 7 |
"lr": 0.001,
|
| 8 |
+
"weight_decay": 0.01,
|
| 9 |
+
"max_grad_norm": 5.0,
|
|
|
|
|
|
|
| 10 |
"max_seq_len": 32,
|
|
|
|
|
|
|
| 11 |
"use_hypothesis": true,
|
| 12 |
"stop_grad_hyp": true,
|
| 13 |
"meta_interval": 5,
|
| 14 |
+
"meta_lr": 0.01,
|
| 15 |
+
"sharpness_lambda": 0.01
|
| 16 |
},
|
| 17 |
"results": {
|
| 18 |
"epochs_completed": 2,
|
| 19 |
+
"killed": false,
|
| 20 |
+
"kill_reason": null,
|
| 21 |
+
"global_reason": null,
|
| 22 |
+
"global_step": 7,
|
| 23 |
+
"micro_batches": 14,
|
| 24 |
+
"best_loss": 5.260426998138428,
|
| 25 |
+
"final_loss": 7.809796333312988,
|
| 26 |
+
"final_avg_loss": 8.098814010620117,
|
| 27 |
+
"final_perplexity": 3290.5631871698165,
|
| 28 |
+
"elapsed_s": 97.07497835159302,
|
| 29 |
+
"params_total": 13033867,
|
| 30 |
+
"params_M": 13.033867,
|
| 31 |
+
"monitor_summary": {
|
| 32 |
+
"total_steps": 7,
|
| 33 |
+
"min_loss": 6.0627899169921875,
|
| 34 |
+
"max_ram_pct": 44.9,
|
| 35 |
+
"min_disk_free_gb": 8.102731776,
|
| 36 |
+
"max_ppl": 54961.20955507722,
|
| 37 |
+
"killed": false,
|
| 38 |
+
"kill_reason": null
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 39 |
},
|
| 40 |
+
"optimizer": "hamiltonian_wasserstein",
|
| 41 |
+
"time_budget": {
|
| 42 |
+
"elapsed_s": 97.07337236404419,
|
| 43 |
+
"remaining_s": 802.9266257286072,
|
| 44 |
+
"epoch_elapsed_s": 56.31318664550781,
|
| 45 |
+
"epoch_remaining_s": 393.6868121623993,
|
| 46 |
+
"is_over": false,
|
| 47 |
+
"is_epoch_over": false
|
| 48 |
+
},
|
| 49 |
+
"step_timer": {
|
| 50 |
+
"count": 14,
|
| 51 |
+
"total_s": 75.34578108787537,
|
| 52 |
+
"avg_s": 5.381841506276812,
|
| 53 |
+
"max_s": 6.357912063598633,
|
| 54 |
+
"expected_s": 2.0
|
| 55 |
+
},
|
| 56 |
+
"hypothesis_monitor": {
|
| 57 |
+
"total_steps": 14,
|
| 58 |
+
"punishment_steps": 14,
|
| 59 |
+
"hyp_activated_steps": 12,
|
| 60 |
+
"hyp_activated_during_punishment": 12,
|
| 61 |
+
"delta_zero_when_activated": 0,
|
| 62 |
+
"l1_entropy_saturated_steps": 14,
|
| 63 |
+
"l2_zero_grad_hyp_steps": 0,
|
| 64 |
+
"crw_converged_steps": 0,
|
| 65 |
+
"crw_contraction_invalid_steps": 1,
|
| 66 |
+
"last_T": 1.0100501775741577,
|
| 67 |
+
"last_tau": 1.0100501775741577,
|
| 68 |
+
"has_bug": true,
|
| 69 |
+
"bugs": [
|
| 70 |
+
"ALERTA: L1 entropy saturada em 14/14 passos — module_selector não está selecionando"
|
| 71 |
+
]
|
| 72 |
+
},
|
| 73 |
+
"dpo_activated": true
|
| 74 |
+
}
|
| 75 |
}
|
v32_test_report.json
ADDED
|
@@ -0,0 +1,74 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"version": "V3.2",
|
| 3 |
+
"timestamp": 1786012938.6738098,
|
| 4 |
+
"validations": [
|
| 5 |
+
{
|
| 6 |
+
"bug": "#1 MetaConfigurator T e \u03c4 movem",
|
| 7 |
+
"passed": true,
|
| 8 |
+
"detail": "T: 1.0 \u2192 1.0101 (moved=True) | \u03c4: 1.0 \u2192 1.0101 (moved=True)"
|
| 9 |
+
},
|
| 10 |
+
{
|
| 11 |
+
"bug": "#2 module_selector torch.clamp + randn init",
|
| 12 |
+
"passed": true,
|
| 13 |
+
"detail": "treino completou 7 steps sem crash"
|
| 14 |
+
},
|
| 15 |
+
{
|
| 16 |
+
"bug": "#3 partial_result kill_reason field",
|
| 17 |
+
"passed": true,
|
| 18 |
+
"detail": "killed=False, kill_reason=None"
|
| 19 |
+
},
|
| 20 |
+
{
|
| 21 |
+
"bug": "#4 u8cell_T import case-sensitivity",
|
| 22 |
+
"passed": true,
|
| 23 |
+
"detail": "imports OK (treino inicializado sem erro)"
|
| 24 |
+
},
|
| 25 |
+
{
|
| 26 |
+
"bug": "#5 OpenMathInstruct-2 generated_solution label_field",
|
| 27 |
+
"passed": true,
|
| 28 |
+
"detail": "8/8 (100%) samples com label_text"
|
| 29 |
+
},
|
| 30 |
+
{
|
| 31 |
+
"bug": "#6 DPO auto-detec\u00e7\u00e3o (\u226530% samples com label)",
|
| 32 |
+
"passed": true,
|
| 33 |
+
"detail": "dpo_activated=True"
|
| 34 |
+
}
|
| 35 |
+
],
|
| 36 |
+
"n_pass": 6,
|
| 37 |
+
"n_total": 6,
|
| 38 |
+
"all_passed": true,
|
| 39 |
+
"training_result": {
|
| 40 |
+
"epochs_completed": 2,
|
| 41 |
+
"global_step": 7,
|
| 42 |
+
"micro_batches": 14,
|
| 43 |
+
"best_loss": 5.260426998138428,
|
| 44 |
+
"final_avg_loss": 8.098814010620117,
|
| 45 |
+
"final_perplexity": 3290.5631871698165,
|
| 46 |
+
"elapsed_s": 97.07497835159302,
|
| 47 |
+
"dpo_activated": true,
|
| 48 |
+
"T_final": 1.0100501775741577,
|
| 49 |
+
"tau_final": 1.0100501775741577,
|
| 50 |
+
"hypothesis_monitor": {
|
| 51 |
+
"total_steps": 14,
|
| 52 |
+
"punishment_steps": 14,
|
| 53 |
+
"hyp_activated_steps": 12,
|
| 54 |
+
"hyp_activated_during_punishment": 12,
|
| 55 |
+
"delta_zero_when_activated": 0,
|
| 56 |
+
"l1_entropy_saturated_steps": 14,
|
| 57 |
+
"l2_zero_grad_hyp_steps": 0,
|
| 58 |
+
"crw_converged_steps": 0,
|
| 59 |
+
"crw_contraction_invalid_steps": 1,
|
| 60 |
+
"last_T": 1.0100501775741577,
|
| 61 |
+
"last_tau": 1.0100501775741577,
|
| 62 |
+
"has_bug": true,
|
| 63 |
+
"bugs": [
|
| 64 |
+
"ALERTA: L1 entropy saturada em 14/14 passos \u2014 module_selector n\u00e3o est\u00e1 selecionando"
|
| 65 |
+
]
|
| 66 |
+
}
|
| 67 |
+
},
|
| 68 |
+
"model_config": {
|
| 69 |
+
"max_modules": 8,
|
| 70 |
+
"d_model": 128,
|
| 71 |
+
"params_total": 13033867,
|
| 72 |
+
"params_M": 13.033867
|
| 73 |
+
}
|
| 74 |
+
}
|
v3_report.json
ADDED
|
@@ -0,0 +1,119 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"version": "V3",
|
| 3 |
+
"multimodal_test": {
|
| 4 |
+
"text_encoder": {
|
| 5 |
+
"ok": true,
|
| 6 |
+
"shape": "(2, 16, 128)"
|
| 7 |
+
},
|
| 8 |
+
"image_encoder": {
|
| 9 |
+
"ok": true,
|
| 10 |
+
"shape": "(2, 128)"
|
| 11 |
+
},
|
| 12 |
+
"audio_encoder": {
|
| 13 |
+
"ok": true,
|
| 14 |
+
"shape": "(2, 128)"
|
| 15 |
+
},
|
| 16 |
+
"video_encoder": {
|
| 17 |
+
"ok": true,
|
| 18 |
+
"shape": "(2, 128)"
|
| 19 |
+
},
|
| 20 |
+
"modal_router_fusion": {
|
| 21 |
+
"ok": true,
|
| 22 |
+
"router_shape": "(1, 128)",
|
| 23 |
+
"fusion_shape": "(2, 128)"
|
| 24 |
+
}
|
| 25 |
+
},
|
| 26 |
+
"smoke_test": {
|
| 27 |
+
"passed": true
|
| 28 |
+
},
|
| 29 |
+
"xeon_optimizations": {
|
| 30 |
+
"ipex_active": true,
|
| 31 |
+
"ipex_version": "2.8.0+cpu",
|
| 32 |
+
"torch_version": "2.8.0+cpu",
|
| 33 |
+
"cpu_isa": "AMX",
|
| 34 |
+
"omp_num_threads": "2",
|
| 35 |
+
"mkl_num_threads": "2",
|
| 36 |
+
"onednn_max_cpu_isa": "AMX_INT8",
|
| 37 |
+
"fp16_benchmark": {
|
| 38 |
+
"fp16_tflops": 0.6224,
|
| 39 |
+
"fp32_tflops": 0.2203,
|
| 40 |
+
"speedup": 2.82
|
| 41 |
+
}
|
| 42 |
+
},
|
| 43 |
+
"model_config": {
|
| 44 |
+
"max_modules": 8,
|
| 45 |
+
"d_model": 128,
|
| 46 |
+
"bigru_hidden": 32,
|
| 47 |
+
"d_transformer": 64,
|
| 48 |
+
"params_total": 13033867,
|
| 49 |
+
"params_M": 13.033867,
|
| 50 |
+
"use_circular_reasoning": true,
|
| 51 |
+
"circular_n_cycles": 3,
|
| 52 |
+
"optimizer": "hamiltonian_wasserstein"
|
| 53 |
+
},
|
| 54 |
+
"training_config": {
|
| 55 |
+
"epochs": 5,
|
| 56 |
+
"use_dpo": true,
|
| 57 |
+
"use_hypothesis": true,
|
| 58 |
+
"meta_interval": 5,
|
| 59 |
+
"max_total_time_s": 2700.0
|
| 60 |
+
},
|
| 61 |
+
"results": {
|
| 62 |
+
"epochs_completed": 5,
|
| 63 |
+
"killed": false,
|
| 64 |
+
"kill_reason": null,
|
| 65 |
+
"global_reason": null,
|
| 66 |
+
"global_step": 25,
|
| 67 |
+
"micro_batches": 50,
|
| 68 |
+
"best_loss": 2.1027424335479736,
|
| 69 |
+
"final_loss": 4.172516345977783,
|
| 70 |
+
"final_avg_loss": 3.9418057203292847,
|
| 71 |
+
"final_perplexity": 51.511532793090105,
|
| 72 |
+
"elapsed_s": 427.32008838653564,
|
| 73 |
+
"params_total": 13033867,
|
| 74 |
+
"params_M": 13.033867,
|
| 75 |
+
"monitor_summary": {
|
| 76 |
+
"total_steps": 25,
|
| 77 |
+
"min_loss": 2.3843162059783936,
|
| 78 |
+
"max_ram_pct": 48.5,
|
| 79 |
+
"min_disk_free_gb": 3.452616704,
|
| 80 |
+
"max_ppl": 40034.00867026023,
|
| 81 |
+
"killed": false,
|
| 82 |
+
"kill_reason": null
|
| 83 |
+
},
|
| 84 |
+
"optimizer": "hamiltonian_wasserstein",
|
| 85 |
+
"time_budget": {
|
| 86 |
+
"elapsed_s": 427.3184747695923,
|
| 87 |
+
"remaining_s": 2272.6815235614777,
|
| 88 |
+
"epoch_elapsed_s": 93.83768510818481,
|
| 89 |
+
"epoch_remaining_s": 446.16231393814087,
|
| 90 |
+
"is_over": false,
|
| 91 |
+
"is_epoch_over": false
|
| 92 |
+
},
|
| 93 |
+
"step_timer": {
|
| 94 |
+
"count": 50,
|
| 95 |
+
"total_s": 330.579580783844,
|
| 96 |
+
"avg_s": 6.61159161567688,
|
| 97 |
+
"max_s": 8.381056547164917,
|
| 98 |
+
"expected_s": 2.0
|
| 99 |
+
},
|
| 100 |
+
"hypothesis_monitor": {
|
| 101 |
+
"total_steps": 50,
|
| 102 |
+
"punishment_steps": 50,
|
| 103 |
+
"hyp_activated_steps": 48,
|
| 104 |
+
"hyp_activated_during_punishment": 48,
|
| 105 |
+
"delta_zero_when_activated": 0,
|
| 106 |
+
"l1_entropy_saturated_steps": 50,
|
| 107 |
+
"l2_zero_grad_hyp_steps": 0,
|
| 108 |
+
"crw_converged_steps": 0,
|
| 109 |
+
"crw_contraction_invalid_steps": 2,
|
| 110 |
+
"last_T": 1.0299465656280518,
|
| 111 |
+
"last_tau": 1.051407814025879,
|
| 112 |
+
"has_bug": true,
|
| 113 |
+
"bugs": [
|
| 114 |
+
"ALERTA: L1 entropy saturada em 50/50 passos \u2014 module_selector n\u00e3o est\u00e1 selecionando"
|
| 115 |
+
]
|
| 116 |
+
},
|
| 117 |
+
"dpo_activated": true
|
| 118 |
+
}
|
| 119 |
+
}
|