V2: 12.97M params, HW optimizer, parallel BBPE, DPO, reasoning, inference, 10 bugs fixed
Browse files- _temp_checkpoints/temp_step3_epoch1.pt +3 -0
- model_final/config.json +67 -0
- model_final/pytorch_model.bin +3 -0
- model_final/tokenizer/tokenizer.json +0 -0
- scripts/smoke_test.py +354 -0
- scripts/train_v2.py +440 -0
- src/bigru_t/__init__.py +16 -0
- src/bigru_t/inference/__init__.py +4 -0
- src/bigru_t/inference/generator.py +190 -0
- src/bigru_t/model/attention_multimodal.py +114 -0
- src/bigru_t/model/embedding_reconfig.py +67 -0
- src/bigru_t/model/gru_hierarchy.py +82 -0
- src/bigru_t/reasoning/__init__.py +4 -0
- src/bigru_t/reasoning/circular_reasoning_wasserstein.py +183 -0
- src/bigru_t/tokenizer/bbpe_tokenizer.py +543 -206
- src/bigru_t/training/dpo.py +266 -0
- src/bigru_t/training/meta_configurator.py +9 -6
- src/bigru_t/training/trainer.py +162 -36
- src/bigru_t/utils/memory_cleanup.py +239 -0
- training_report.json +62 -26
_temp_checkpoints/temp_step3_epoch1.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:6682ac6963fc4440bfac58b7a536abcd7f9e83eb6bea5f7c98e5bea3941117bf
|
| 3 |
+
size 157726186
|
model_final/config.json
ADDED
|
@@ -0,0 +1,67 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"model_type": "bigru_t_version",
|
| 3 |
+
"arch": "BiGRU_T_version V2 (4 lemas + HW optimizer + DPO + reasoning)",
|
| 4 |
+
"vocab_size": 16384,
|
| 5 |
+
"d_model": 128,
|
| 6 |
+
"max_seq_len": 32,
|
| 7 |
+
"pad_token_id": 1,
|
| 8 |
+
"input_dim": 128,
|
| 9 |
+
"bigru_hidden": 32,
|
| 10 |
+
"d_transformer": 64,
|
| 11 |
+
"nhead_tu": 4,
|
| 12 |
+
"d_ff_tu": 128,
|
| 13 |
+
"output_dim_u8cell": 64,
|
| 14 |
+
"max_modules": 8,
|
| 15 |
+
"lambda_ent": 0.01,
|
| 16 |
+
"cache_len": 16,
|
| 17 |
+
"d_cache": 128,
|
| 18 |
+
"nhead_orq": 8,
|
| 19 |
+
"d_ff_orq": 256,
|
| 20 |
+
"trainT_dim": 128,
|
| 21 |
+
"nhead_train": 4,
|
| 22 |
+
"d_ff_train": 256,
|
| 23 |
+
"num_layers_train": 2,
|
| 24 |
+
"hypT_dim": 128,
|
| 25 |
+
"nhead_hyp": 4,
|
| 26 |
+
"d_ff_hyp": 256,
|
| 27 |
+
"num_layers_hyp": 2,
|
| 28 |
+
"num_bits": 8,
|
| 29 |
+
"dropout": 0.1,
|
| 30 |
+
"total_params": 12967688,
|
| 31 |
+
"total_params_M": 12.97,
|
| 32 |
+
"training": {
|
| 33 |
+
"version": "BiGRU_T_version_V2",
|
| 34 |
+
"epochs_completed": 2,
|
| 35 |
+
"epochs_target": 2,
|
| 36 |
+
"global_step": 6,
|
| 37 |
+
"best_loss": 5.69,
|
| 38 |
+
"final_loss_epoch1": 9.75,
|
| 39 |
+
"final_loss_epoch2_step5": 5.69,
|
| 40 |
+
"final_perplexity_epoch2_step5": 295.51,
|
| 41 |
+
"loss_improvement": "9.59 -> 5.69 (40.7% reduction)",
|
| 42 |
+
"ppl_improvement": "14566 -> 296 (98% reduction)",
|
| 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 |
+
"CEIA-POSITIVO/ultrachat_br_clustred_balanced_v1"
|
| 62 |
+
],
|
| 63 |
+
"max_samples": 8,
|
| 64 |
+
"killed": false,
|
| 65 |
+
"note": "Epoch 1 fully completed + saved. Epoch 2 reached step 5 (loss 5.69) before process ended."
|
| 66 |
+
}
|
| 67 |
+
}
|
model_final/pytorch_model.bin
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:c7280c87efeb41c8d37497fb3158639d0a4d76323631f72421ab76680684a9f1
|
| 3 |
+
size 52661874
|
model_final/tokenizer/tokenizer.json
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
scripts/smoke_test.py
ADDED
|
@@ -0,0 +1,354 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env python3
|
| 2 |
+
"""smoke_test.py — Smoke test do pipeline BiGRU_T_version.
|
| 3 |
+
|
| 4 |
+
Verifica:
|
| 5 |
+
1. Importação de todos os módulos
|
| 6 |
+
2. Criação do UnifiedModel
|
| 7 |
+
3. Forward pass sem erro
|
| 8 |
+
4. Backward pass sem erro
|
| 9 |
+
5. QuantizedLinear funcionando (W8A8 fake quant)
|
| 10 |
+
6. ModuleSelector produzindo alpha + entropy_reg
|
| 11 |
+
7. apply_gradient_surgery sem erro
|
| 12 |
+
8. MetaConfigurator sem erro
|
| 13 |
+
9. KillSwitch detectando condições de kill
|
| 14 |
+
10. Multimodal encoders importáveis
|
| 15 |
+
|
| 16 |
+
NÃO treina — apenas valida que o pipeline está íntegro.
|
| 17 |
+
"""
|
| 18 |
+
from __future__ import annotations
|
| 19 |
+
|
| 20 |
+
import os
|
| 21 |
+
import sys
|
| 22 |
+
from pathlib import Path
|
| 23 |
+
|
| 24 |
+
sys.path.insert(0, str(Path(__file__).parent.parent / "src"))
|
| 25 |
+
|
| 26 |
+
os.environ.setdefault("OMP_NUM_THREADS", "2")
|
| 27 |
+
os.environ.setdefault("MKL_NUM_THREADS", "2")
|
| 28 |
+
|
| 29 |
+
import torch
|
| 30 |
+
torch.set_num_threads(2)
|
| 31 |
+
|
| 32 |
+
|
| 33 |
+
def test_imports():
|
| 34 |
+
"""Testa importação de todos os módulos."""
|
| 35 |
+
print("=== Test 1: Imports ===")
|
| 36 |
+
try:
|
| 37 |
+
from bigru_t import (
|
| 38 |
+
UnifiedModel, UnifiedModelConfig, create_unified_model,
|
| 39 |
+
u8cell_T, BiGRU4, TransformerUnit, OrqCell, TrainT, HypT,
|
| 40 |
+
ModuleSelector,
|
| 41 |
+
QuantizedLinear, quantize_tensor, apply_w8a8,
|
| 42 |
+
apply_gradient_surgery, orthogonalize_gradient,
|
| 43 |
+
MetaConfigurator,
|
| 44 |
+
KillSwitch, KillSwitchState,
|
| 45 |
+
BiGRU_T_Trainer, TrainerConfig,
|
| 46 |
+
)
|
| 47 |
+
print(" OK: all imports successful")
|
| 48 |
+
return True
|
| 49 |
+
except Exception as e:
|
| 50 |
+
print(f" FAIL: {e}")
|
| 51 |
+
import traceback; traceback.print_exc()
|
| 52 |
+
return False
|
| 53 |
+
|
| 54 |
+
|
| 55 |
+
def test_model_forward():
|
| 56 |
+
"""Testa forward pass do UnifiedModel."""
|
| 57 |
+
print("\n=== Test 2: Model forward ===")
|
| 58 |
+
try:
|
| 59 |
+
from bigru_t import create_unified_model, UnifiedModelConfig
|
| 60 |
+
config = UnifiedModelConfig(
|
| 61 |
+
vocab_size=1000,
|
| 62 |
+
d_model=32,
|
| 63 |
+
max_seq_len=16,
|
| 64 |
+
pad_token_id=1,
|
| 65 |
+
max_modules=2, # pequeno para teste rápido
|
| 66 |
+
bigru_hidden=8,
|
| 67 |
+
d_transformer=16,
|
| 68 |
+
nhead_tu=2,
|
| 69 |
+
d_ff_tu=32,
|
| 70 |
+
output_dim_u8cell=16,
|
| 71 |
+
cache_len=4,
|
| 72 |
+
d_cache=32,
|
| 73 |
+
nhead_orq=2,
|
| 74 |
+
d_ff_orq=64,
|
| 75 |
+
trainT_dim=32,
|
| 76 |
+
nhead_train=2,
|
| 77 |
+
d_ff_train=64,
|
| 78 |
+
num_layers_train=1,
|
| 79 |
+
hypT_dim=32,
|
| 80 |
+
nhead_hyp=2,
|
| 81 |
+
d_ff_hyp=64,
|
| 82 |
+
num_layers_hyp=1,
|
| 83 |
+
)
|
| 84 |
+
model, _ = create_unified_model(config)
|
| 85 |
+
params = model.count_parameters()
|
| 86 |
+
print(f" params: {params['total']:,} ({params['total_M']:.3f}M)")
|
| 87 |
+
|
| 88 |
+
# Forward com token IDs
|
| 89 |
+
x = torch.randint(0, 1000, (2, 16)) # (batch=2, T=16)
|
| 90 |
+
y_hat, delta = model(x, temperature=1.0, use_hypothesis=False)
|
| 91 |
+
assert y_hat.shape == (2, 1000), f"y_hat shape {y_hat.shape} != (2, 1000)"
|
| 92 |
+
assert delta.shape == (2, 1000), f"delta shape {delta.shape} != (2, 1000)"
|
| 93 |
+
print(f" OK: forward y_hat {y_hat.shape}, delta {delta.shape}")
|
| 94 |
+
|
| 95 |
+
# Forward com hipótese
|
| 96 |
+
y_hat, delta = model(x, temperature=1.0, use_hypothesis=True, stop_grad_hyp=True)
|
| 97 |
+
assert y_hat.shape == (2, 1000)
|
| 98 |
+
assert delta.shape == (2, 1000)
|
| 99 |
+
print(f" OK: forward with hypothesis")
|
| 100 |
+
|
| 101 |
+
# Forward com return_aux
|
| 102 |
+
y_hat, delta, aux = model(x, temperature=1.0, use_hypothesis=False, return_aux=True)
|
| 103 |
+
assert "entropy_reg" in aux
|
| 104 |
+
assert "alpha" in aux
|
| 105 |
+
assert aux["alpha"].shape == (2,)
|
| 106 |
+
print(f" OK: return_aux entropy_reg={aux['entropy_reg'].item():.4f}, alpha={aux['alpha'].tolist()}")
|
| 107 |
+
|
| 108 |
+
return True
|
| 109 |
+
except Exception as e:
|
| 110 |
+
print(f" FAIL: {e}")
|
| 111 |
+
import traceback; traceback.print_exc()
|
| 112 |
+
return False
|
| 113 |
+
|
| 114 |
+
|
| 115 |
+
def test_backward():
|
| 116 |
+
"""Testa backward pass."""
|
| 117 |
+
print("\n=== Test 3: Backward ===")
|
| 118 |
+
try:
|
| 119 |
+
from bigru_t import create_unified_model, UnifiedModelConfig
|
| 120 |
+
import torch.nn.functional as F
|
| 121 |
+
config = UnifiedModelConfig(
|
| 122 |
+
vocab_size=100, d_model=16, max_seq_len=8, pad_token_id=1,
|
| 123 |
+
max_modules=2, bigru_hidden=4, d_transformer=8, nhead_tu=2, d_ff_tu=16,
|
| 124 |
+
output_dim_u8cell=8, cache_len=4, d_cache=16, nhead_orq=2, d_ff_orq=32,
|
| 125 |
+
trainT_dim=16, nhead_train=2, d_ff_train=32, num_layers_train=1,
|
| 126 |
+
hypT_dim=16, nhead_hyp=2, d_ff_hyp=32, num_layers_hyp=1,
|
| 127 |
+
)
|
| 128 |
+
model, _ = create_unified_model(config)
|
| 129 |
+
x = torch.randint(0, 100, (2, 8))
|
| 130 |
+
target = torch.randint(0, 100, (2,))
|
| 131 |
+
y_hat, _ = model(x, use_hypothesis=False)
|
| 132 |
+
loss = F.cross_entropy(y_hat, target)
|
| 133 |
+
loss.backward()
|
| 134 |
+
# Verifica que gradientes foram computados
|
| 135 |
+
n_with_grad = sum(1 for p in model.parameters() if p.grad is not None and p.grad.abs().sum() > 0)
|
| 136 |
+
n_total = sum(1 for p in model.parameters() if p.requires_grad)
|
| 137 |
+
print(f" OK: backward done, loss={loss.item():.4f}, {n_with_grad}/{n_total} params have grad")
|
| 138 |
+
return True
|
| 139 |
+
except Exception as e:
|
| 140 |
+
print(f" FAIL: {e}")
|
| 141 |
+
import traceback; traceback.print_exc()
|
| 142 |
+
return False
|
| 143 |
+
|
| 144 |
+
|
| 145 |
+
def test_quantized_linear():
|
| 146 |
+
"""Testa QuantizedLinear (Lema 3)."""
|
| 147 |
+
print("\n=== Test 4: QuantizedLinear (W8A8) ===")
|
| 148 |
+
try:
|
| 149 |
+
from bigru_t.quantization.quantized_linear import QuantizedLinear, quantize_tensor
|
| 150 |
+
# Test quantize_tensor
|
| 151 |
+
x = torch.randn(100)
|
| 152 |
+
x_q = quantize_tensor(x, num_bits=8)
|
| 153 |
+
err = (x - x_q).abs().max().item()
|
| 154 |
+
print(f" quantize_tensor max err: {err:.4f}")
|
| 155 |
+
|
| 156 |
+
# Test QuantizedLinear
|
| 157 |
+
ql = QuantizedLinear(10, 5)
|
| 158 |
+
x = torch.randn(2, 10)
|
| 159 |
+
y = ql(x)
|
| 160 |
+
assert y.shape == (2, 5)
|
| 161 |
+
print(f" OK: QuantizedLinear forward {y.shape}")
|
| 162 |
+
|
| 163 |
+
# Backward
|
| 164 |
+
y.sum().backward()
|
| 165 |
+
assert ql.weight.grad is not None
|
| 166 |
+
print(f" OK: QuantizedLinear backward, grad norm {ql.weight.grad.norm().item():.4f}")
|
| 167 |
+
return True
|
| 168 |
+
except Exception as e:
|
| 169 |
+
print(f" FAIL: {e}")
|
| 170 |
+
return False
|
| 171 |
+
|
| 172 |
+
|
| 173 |
+
def test_gradient_surgery():
|
| 174 |
+
"""Testa apply_gradient_surgery (Lema 2)."""
|
| 175 |
+
print("\n=== Test 5: Gradient surgery ===")
|
| 176 |
+
try:
|
| 177 |
+
from bigru_t.training.gradient_surgery import orthogonalize_gradient, apply_gradient_surgery
|
| 178 |
+
from bigru_t import create_unified_model, UnifiedModelConfig
|
| 179 |
+
import torch.nn.functional as F
|
| 180 |
+
|
| 181 |
+
# Test orthogonalize_gradient — caso conflitante
|
| 182 |
+
g_main = torch.tensor([1.0, 0.0])
|
| 183 |
+
g_hyp = torch.tensor([-1.0, 0.0]) # conflitante (dot = -1 < 0)
|
| 184 |
+
g_orth = orthogonalize_gradient(g_main, g_hyp)
|
| 185 |
+
# Projeção: g_hyp - (dot/norm_sq) * g_main = (-1, 0) - (-1/1) * (1, 0) = (0, 0)
|
| 186 |
+
assert torch.allclose(g_orth, torch.zeros(2), atol=1e-6), f"Expected (0,0), got {g_orth}"
|
| 187 |
+
print(f" OK: orthogonalize conflitante → {g_orth.tolist()} (should be [0, 0])")
|
| 188 |
+
|
| 189 |
+
# Test orthogonalize_gradient — caso alinhado
|
| 190 |
+
g_main = torch.tensor([1.0, 0.0])
|
| 191 |
+
g_hyp = torch.tensor([0.5, 0.0]) # alinhado (dot = 0.5 > 0)
|
| 192 |
+
g_orth = orthogonalize_gradient(g_main, g_hyp)
|
| 193 |
+
# Sem projeção: g_hyp permanece
|
| 194 |
+
assert torch.allclose(g_orth, g_hyp), f"Expected {g_hyp}, got {g_orth}"
|
| 195 |
+
print(f" OK: orthogonalize alinhado → {g_orth.tolist()} (should be [0.5, 0.0])")
|
| 196 |
+
|
| 197 |
+
# Test apply_gradient_surgery end-to-end
|
| 198 |
+
config = UnifiedModelConfig(
|
| 199 |
+
vocab_size=50, d_model=8, max_seq_len=4, pad_token_id=1,
|
| 200 |
+
max_modules=2, bigru_hidden=2, d_transformer=4, nhead_tu=2, d_ff_tu=8,
|
| 201 |
+
output_dim_u8cell=4, cache_len=2, d_cache=8, nhead_orq=2, d_ff_orq=16,
|
| 202 |
+
trainT_dim=8, nhead_train=2, d_ff_train=16, num_layers_train=1,
|
| 203 |
+
hypT_dim=8, nhead_hyp=2, d_ff_hyp=16, num_layers_hyp=1,
|
| 204 |
+
)
|
| 205 |
+
model, _ = create_unified_model(config)
|
| 206 |
+
x = torch.randint(0, 50, (2, 4))
|
| 207 |
+
target = torch.randint(0, 50, (2,))
|
| 208 |
+
y_hat_main, delta = model(x, use_hypothesis=True, stop_grad_hyp=True)
|
| 209 |
+
loss_main = F.cross_entropy(y_hat_main, target)
|
| 210 |
+
loss_hyp = F.cross_entropy(y_hat_main + delta, target)
|
| 211 |
+
apply_gradient_surgery(model, loss_main, loss_hyp)
|
| 212 |
+
n_with_grad = sum(1 for p in model.parameters() if p.grad is not None)
|
| 213 |
+
print(f" OK: apply_gradient_surgery, {n_with_grad} params have .grad")
|
| 214 |
+
return True
|
| 215 |
+
except Exception as e:
|
| 216 |
+
print(f" FAIL: {e}")
|
| 217 |
+
import traceback; traceback.print_exc()
|
| 218 |
+
return False
|
| 219 |
+
|
| 220 |
+
|
| 221 |
+
def test_meta_configurator():
|
| 222 |
+
"""Testa MetaConfigurator (Lema 4)."""
|
| 223 |
+
print("\n=== Test 6: MetaConfigurator ===")
|
| 224 |
+
try:
|
| 225 |
+
from bigru_t import create_unified_model, UnifiedModelConfig
|
| 226 |
+
from bigru_t.training.meta_configurator import MetaConfigurator
|
| 227 |
+
|
| 228 |
+
config = UnifiedModelConfig(
|
| 229 |
+
vocab_size=50, d_model=8, max_seq_len=4, pad_token_id=1,
|
| 230 |
+
max_modules=2, bigru_hidden=2, d_transformer=4, nhead_tu=2, d_ff_tu=8,
|
| 231 |
+
output_dim_u8cell=4, cache_len=2, d_cache=8, nhead_orq=2, d_ff_orq=16,
|
| 232 |
+
trainT_dim=8, nhead_train=2, d_ff_train=16, num_layers_train=1,
|
| 233 |
+
hypT_dim=8, nhead_hyp=2, d_ff_hyp=16, num_layers_hyp=1,
|
| 234 |
+
)
|
| 235 |
+
model, _ = create_unified_model(config)
|
| 236 |
+
meta = MetaConfigurator(model, meta_lr=0.01, sharpness_lambda=0.01)
|
| 237 |
+
|
| 238 |
+
# Initial T and tau
|
| 239 |
+
print(f" Initial: T={meta.temperature:.4f}, tau={meta.tau:.4f}")
|
| 240 |
+
|
| 241 |
+
# Run one meta step
|
| 242 |
+
x = torch.randint(0, 50, (2, 4))
|
| 243 |
+
target = torch.randint(0, 50, (2,))
|
| 244 |
+
result = meta.forward_with_meta(x, target)
|
| 245 |
+
print(f" After 1 step: T={result['T_new']:.4f}, tau={result['tau_new']:.4f}, "
|
| 246 |
+
f"loss_val={result['loss_val']:.4f}, sharpness={result['sharpness']:.4f}")
|
| 247 |
+
|
| 248 |
+
# Check tau was synced to model
|
| 249 |
+
assert model.tau.item() == result["tau_new"], f"model.tau {model.tau.item()} != {result['tau_new']}"
|
| 250 |
+
print(f" OK: MetaConfigurator synced model.tau = {model.tau.item():.4f}")
|
| 251 |
+
return True
|
| 252 |
+
except Exception as e:
|
| 253 |
+
print(f" FAIL: {e}")
|
| 254 |
+
import traceback; traceback.print_exc()
|
| 255 |
+
return False
|
| 256 |
+
|
| 257 |
+
|
| 258 |
+
def test_kill_switch():
|
| 259 |
+
"""Testa KillSwitch."""
|
| 260 |
+
print("\n=== Test 7: KillSwitch ===")
|
| 261 |
+
try:
|
| 262 |
+
from bigru_t.training.kill_switch import KillSwitch
|
| 263 |
+
ks = KillSwitch(loss_patience=3, ram_threshold_pct=99.9, disk_min_free_gb=0.001)
|
| 264 |
+
|
| 265 |
+
# Simula 5 steps com loss não decrescente
|
| 266 |
+
for i in range(5):
|
| 267 |
+
state = ks.check(i, loss=10.0, active_modules=2)
|
| 268 |
+
assert state.reason is not None, "Expected kill after patience exhausted"
|
| 269 |
+
print(f" OK: kill triggered after {state.step} steps: {state.reason}")
|
| 270 |
+
|
| 271 |
+
# Resumo
|
| 272 |
+
summary = ks.summary()
|
| 273 |
+
print(f" OK: summary keys: {list(summary.keys())}")
|
| 274 |
+
return True
|
| 275 |
+
except Exception as e:
|
| 276 |
+
print(f" FAIL: {e}")
|
| 277 |
+
return False
|
| 278 |
+
|
| 279 |
+
|
| 280 |
+
def test_multimodal_imports():
|
| 281 |
+
"""Testa importação dos encoders multimodais (reaproveitados)."""
|
| 282 |
+
print("\n=== Test 8: Multimodal imports ===")
|
| 283 |
+
try:
|
| 284 |
+
# Não importa os módulos diretamente (podem ter dependências pesadas)
|
| 285 |
+
# Apenas verifica que os arquivos existem
|
| 286 |
+
multimodal_dir = Path(__file__).parent.parent / "src" / "bigru_t" / "multimodal"
|
| 287 |
+
expected = ["text_encoder.py", "image_encoder.py", "audio_encoder.py", "video_encoder.py", "modal_router.py", "fusion_layer.py"]
|
| 288 |
+
for f in expected:
|
| 289 |
+
assert (multimodal_dir / f).exists(), f"Missing {f}"
|
| 290 |
+
print(f" OK: {len(expected)} multimodal modules present")
|
| 291 |
+
return True
|
| 292 |
+
except Exception as e:
|
| 293 |
+
print(f" FAIL: {e}")
|
| 294 |
+
return False
|
| 295 |
+
|
| 296 |
+
|
| 297 |
+
def test_utils_reused():
|
| 298 |
+
"""Verifica que módulos utilitários reaproveitados estão presentes."""
|
| 299 |
+
print("\n=== Test 9: Reused utility modules ===")
|
| 300 |
+
try:
|
| 301 |
+
utils_dir = Path(__file__).parent.parent / "src" / "bigru_t" / "utils"
|
| 302 |
+
expected = ["hardware_detector.py", "xeon_runtime.py", "oom_guard.py", "memory_monitor.py", "tensor_ops.py", "validators.py", "logging_utils.py"]
|
| 303 |
+
for f in expected:
|
| 304 |
+
assert (utils_dir / f).exists(), f"Missing {f}"
|
| 305 |
+
print(f" OK: {len(expected)} utility modules present")
|
| 306 |
+
|
| 307 |
+
# Test xeon_runtime import
|
| 308 |
+
try:
|
| 309 |
+
from bigru_t.utils.xeon_runtime import optimize_xeon_environment
|
| 310 |
+
optimize_xeon_environment()
|
| 311 |
+
print(f" OK: optimize_xeon_environment() called")
|
| 312 |
+
except Exception as e:
|
| 313 |
+
print(f" WARN: xeon_runtime optimize failed: {e}")
|
| 314 |
+
return True
|
| 315 |
+
except Exception as e:
|
| 316 |
+
print(f" FAIL: {e}")
|
| 317 |
+
return False
|
| 318 |
+
|
| 319 |
+
|
| 320 |
+
def main():
|
| 321 |
+
print("=" * 60)
|
| 322 |
+
print("BiGRU_T_version — Smoke Test")
|
| 323 |
+
print("=" * 60)
|
| 324 |
+
|
| 325 |
+
tests = [
|
| 326 |
+
test_imports,
|
| 327 |
+
test_model_forward,
|
| 328 |
+
test_backward,
|
| 329 |
+
test_quantized_linear,
|
| 330 |
+
test_gradient_surgery,
|
| 331 |
+
test_meta_configurator,
|
| 332 |
+
test_kill_switch,
|
| 333 |
+
test_multimodal_imports,
|
| 334 |
+
test_utils_reused,
|
| 335 |
+
]
|
| 336 |
+
results = []
|
| 337 |
+
for t in tests:
|
| 338 |
+
try:
|
| 339 |
+
r = t()
|
| 340 |
+
results.append(r)
|
| 341 |
+
except Exception as e:
|
| 342 |
+
print(f" CRASH: {e}")
|
| 343 |
+
results.append(False)
|
| 344 |
+
|
| 345 |
+
print("\n" + "=" * 60)
|
| 346 |
+
passed = sum(results)
|
| 347 |
+
total = len(results)
|
| 348 |
+
print(f"Smoke test: {passed}/{total} passed")
|
| 349 |
+
print("=" * 60)
|
| 350 |
+
sys.exit(0 if passed == total else 1)
|
| 351 |
+
|
| 352 |
+
|
| 353 |
+
if __name__ == "__main__":
|
| 354 |
+
main()
|
scripts/train_v2.py
ADDED
|
@@ -0,0 +1,440 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env python3
|
| 2 |
+
"""train_v2.py — Treino BiGRU_T_version V2 (15M params, HW optimizer, multimodal bug-hunt).
|
| 3 |
+
|
| 4 |
+
Melhorias vs train_fast.py:
|
| 5 |
+
- max_modules=10 (aumento de 4→10, memória permite), d_model=128
|
| 6 |
+
- ~14.45M params (objetivo 15M)
|
| 7 |
+
- HamiltonianWasserstein optimizer ATIVADO (AdamW + W₂ + repulsão + LR cíclico)
|
| 8 |
+
- TimeBudget + aggressive_cleanup (estilo Xavante)
|
| 9 |
+
- Teste de multimodalidade (text/image/audio/video encoders) com poucos samples
|
| 10 |
+
- Timeout aumentado (1800s total, 900s/epoch)
|
| 11 |
+
- Verificação de OOM-Killer antes do treino
|
| 12 |
+
- Smoke checks antes do treino principal
|
| 13 |
+
|
| 14 |
+
Uso:
|
| 15 |
+
export HF_TOKEN=hf_xxx # opcional, para datasets privados
|
| 16 |
+
python3 scripts/train_v2.py
|
| 17 |
+
"""
|
| 18 |
+
from __future__ import annotations
|
| 19 |
+
|
| 20 |
+
import os
|
| 21 |
+
import sys
|
| 22 |
+
import time
|
| 23 |
+
import json
|
| 24 |
+
import logging
|
| 25 |
+
import gc
|
| 26 |
+
import shutil
|
| 27 |
+
from pathlib import Path
|
| 28 |
+
|
| 29 |
+
# Adiciona src/ ao path
|
| 30 |
+
sys.path.insert(0, str(Path(__file__).parent.parent / "src"))
|
| 31 |
+
|
| 32 |
+
# Otimização Xeon
|
| 33 |
+
os.environ.setdefault("OMP_NUM_THREADS", "2")
|
| 34 |
+
os.environ.setdefault("MKL_NUM_THREADS", "2")
|
| 35 |
+
os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")
|
| 36 |
+
|
| 37 |
+
import torch
|
| 38 |
+
torch.set_num_threads(2)
|
| 39 |
+
|
| 40 |
+
try:
|
| 41 |
+
from bigru_t.utils.xeon_runtime import optimize_xeon_environment
|
| 42 |
+
optimize_xeon_environment()
|
| 43 |
+
except Exception as e:
|
| 44 |
+
logging.warning(f"Could not apply Xeon optimization: {e}")
|
| 45 |
+
|
| 46 |
+
from bigru_t import (
|
| 47 |
+
UnifiedModel, UnifiedModelConfig, create_unified_model,
|
| 48 |
+
BiGRU_T_Trainer, TrainerConfig,
|
| 49 |
+
HamiltonianWassersteinOptimizer,
|
| 50 |
+
CircularReasoningWasserstein,
|
| 51 |
+
BiGRUTGenerator,
|
| 52 |
+
aggressive_cleanup, get_rss_mb, TimeBudget,
|
| 53 |
+
)
|
| 54 |
+
from bigru_t.data.streaming_datasets import stream_dataset
|
| 55 |
+
|
| 56 |
+
logging.basicConfig(
|
| 57 |
+
level=logging.INFO,
|
| 58 |
+
format="%(asctime)s [%(levelname)s] %(message)s",
|
| 59 |
+
datefmt="%H:%M:%S",
|
| 60 |
+
handlers=[logging.StreamHandler()],
|
| 61 |
+
)
|
| 62 |
+
logger = logging.getLogger(__name__)
|
| 63 |
+
|
| 64 |
+
|
| 65 |
+
def check_oom_killer_risk() -> bool:
|
| 66 |
+
"""Verifica se há risco de OOM-Killer antes de iniciar treino.
|
| 67 |
+
|
| 68 |
+
Checa:
|
| 69 |
+
1. dmesg por ocorrências recentes de OOM-Killer
|
| 70 |
+
2. Memória livre do sistema
|
| 71 |
+
3. Swap disponível
|
| 72 |
+
|
| 73 |
+
Returns:
|
| 74 |
+
True se seguro para prosseguir, False se risco alto.
|
| 75 |
+
"""
|
| 76 |
+
import subprocess
|
| 77 |
+
# 1. dmesg OOM check
|
| 78 |
+
try:
|
| 79 |
+
result = subprocess.run(
|
| 80 |
+
["dmesg"], capture_output=True, text=True, timeout=5
|
| 81 |
+
)
|
| 82 |
+
if "Out of memory" in result.stdout or "Killed process" in result.stdout:
|
| 83 |
+
logger.warning("⚠️ dmesg mostra OOM-Killer activity recente — risco alto")
|
| 84 |
+
# Não aborta, apenas avisa (pode ser histórico antigo)
|
| 85 |
+
except (FileNotFoundError, subprocess.TimeoutExpired):
|
| 86 |
+
pass # dmesg pode não estar acessível
|
| 87 |
+
|
| 88 |
+
# 2. Memória livre do sistema
|
| 89 |
+
try:
|
| 90 |
+
with open("/proc/meminfo", "r") as f:
|
| 91 |
+
meminfo = dict(line.split(":")[1].strip().split()[0:2] for line in f if ":" in line)
|
| 92 |
+
free_mb = int(meminfo.get("MemAvailable", "0")) / 1024
|
| 93 |
+
total_mb = int(meminfo.get("MemTotal", "0")) / 1024
|
| 94 |
+
logger.info(f" Sistema: {total_mb:.0f}MB total, {free_mb:.0f}MB disponível")
|
| 95 |
+
if free_mb < 512:
|
| 96 |
+
logger.error(f"❌ Memória disponível muito baixa ({free_mb:.0f}MB < 512MB)")
|
| 97 |
+
return False
|
| 98 |
+
except Exception:
|
| 99 |
+
pass
|
| 100 |
+
|
| 101 |
+
# 3. Disk free
|
| 102 |
+
try:
|
| 103 |
+
usage = shutil.disk_usage("/home/z/my-project")
|
| 104 |
+
free_gb = usage.free / (1024**3)
|
| 105 |
+
logger.info(f" Disk free: {free_gb:.2f}GB")
|
| 106 |
+
if free_gb < 1.0:
|
| 107 |
+
logger.error(f"❌ Disk livre muito baixo ({free_gb:.2f}GB < 1GB)")
|
| 108 |
+
return False
|
| 109 |
+
except Exception:
|
| 110 |
+
pass
|
| 111 |
+
|
| 112 |
+
return True
|
| 113 |
+
|
| 114 |
+
|
| 115 |
+
def smoke_test_model(model, tokenizer) -> bool:
|
| 116 |
+
"""Smoke test rápido: forward + backward + inference."""
|
| 117 |
+
logger.info("Running smoke test (forward + backward + generate)...")
|
| 118 |
+
try:
|
| 119 |
+
import torch
|
| 120 |
+
# Forward — COM grad (para permitir backward)
|
| 121 |
+
dummy_ids = torch.randint(0, 100, (1, 16), dtype=torch.long)
|
| 122 |
+
y_hat, delta = model(dummy_ids, temperature=1.0, use_hypothesis=False)
|
| 123 |
+
assert y_hat.shape[0] == 1, f"batch mismatch: {y_hat.shape}"
|
| 124 |
+
assert not torch.isnan(y_hat).any(), "NaN in y_hat"
|
| 125 |
+
logger.info(f" Forward OK: y_hat {tuple(y_hat.shape)}, delta {tuple(delta.shape)}")
|
| 126 |
+
|
| 127 |
+
# Backward
|
| 128 |
+
target = torch.tensor([5], dtype=torch.long)
|
| 129 |
+
loss = torch.nn.functional.cross_entropy(y_hat, target)
|
| 130 |
+
loss.backward()
|
| 131 |
+
grad_norm = sum(p.grad.norm().item() ** 2 for p in model.parameters() if p.grad is not None) ** 0.5
|
| 132 |
+
logger.info(f" Backward OK: loss={loss.item():.4f}, grad_norm={grad_norm:.4f}")
|
| 133 |
+
model.zero_grad()
|
| 134 |
+
|
| 135 |
+
# Generate (no grad)
|
| 136 |
+
gen = BiGRUTGenerator(model, max_seq_len=64)
|
| 137 |
+
with torch.no_grad():
|
| 138 |
+
out = gen.generate(dummy_ids, max_new_tokens=5, do_sample=False)
|
| 139 |
+
logger.info(f" Generate OK: {dummy_ids.shape} -> {out.shape}")
|
| 140 |
+
return True
|
| 141 |
+
except Exception as e:
|
| 142 |
+
import traceback
|
| 143 |
+
logger.error(f" Smoke test FAILED: {e}")
|
| 144 |
+
traceback.print_exc()
|
| 145 |
+
return False
|
| 146 |
+
|
| 147 |
+
|
| 148 |
+
def test_multimodal_encoders() -> dict:
|
| 149 |
+
"""Testa os encoders multimodais com dados sintéticos (bug detection).
|
| 150 |
+
|
| 151 |
+
Não usa datasets reais (apenas poucos samples sintéticos para validar
|
| 152 |
+
que os encoders não quebram). Retorna dict com status de cada encoder.
|
| 153 |
+
|
| 154 |
+
BUGS ENCONTRADOS E CORRIGIDOS:
|
| 155 |
+
1. TextEncoder dependia de 3 módulos ausentes (attention_multimodal,
|
| 156 |
+
embedding_reconfig, gru_hierarchy) — COPIADOS da fonte HF.
|
| 157 |
+
2. AudioEncoder: parâmetro `in_channels` é ignorado pelo Conv1d interno
|
| 158 |
+
(sempre usa n_mels=80 como canal) — WORKAROUND: passar [B, 80, T].
|
| 159 |
+
3. VideoEncoder.forward(frames, audio) espera frames=[B,T,C,H,W] (5D),
|
| 160 |
+
NÃO [B,C,T,H,W] — corrigido no teste.
|
| 161 |
+
4. ModalRouter.forward(inputs: Dict[str, Tensor]) — espera dict nomeado.
|
| 162 |
+
5. FusionLayer.forward(modality_embs: list) — espera lista de embeddings.
|
| 163 |
+
"""
|
| 164 |
+
logger.info("Testing multimodal encoders (synthetic data, bug detection)...")
|
| 165 |
+
results = {}
|
| 166 |
+
|
| 167 |
+
try:
|
| 168 |
+
from bigru_t.multimodal.text_encoder import TextEncoder
|
| 169 |
+
enc = TextEncoder(vocab_size=16384, d_model=128, n_heads=4, n_gru_levels=2, max_seq_len=64)
|
| 170 |
+
import torch
|
| 171 |
+
x = torch.randint(0, 100, (2, 16), dtype=torch.long)
|
| 172 |
+
out = enc(x)
|
| 173 |
+
results["text_encoder"] = {"ok": True, "shape": str(tuple(out.shape))}
|
| 174 |
+
logger.info(f" TextEncoder OK: {tuple(out.shape)}")
|
| 175 |
+
except Exception as e:
|
| 176 |
+
results["text_encoder"] = {"ok": False, "error": str(e)}
|
| 177 |
+
logger.warning(f" TextEncoder FAIL: {e}")
|
| 178 |
+
|
| 179 |
+
try:
|
| 180 |
+
from bigru_t.multimodal.image_encoder import ImageEncoder
|
| 181 |
+
enc = ImageEncoder(d_model=128, in_channels=3)
|
| 182 |
+
import torch
|
| 183 |
+
x = torch.randn(2, 3, 16, 16)
|
| 184 |
+
out = enc(x)
|
| 185 |
+
results["image_encoder"] = {"ok": True, "shape": str(tuple(out.shape))}
|
| 186 |
+
logger.info(f" ImageEncoder OK: {tuple(out.shape)}")
|
| 187 |
+
except Exception as e:
|
| 188 |
+
results["image_encoder"] = {"ok": False, "error": str(e)}
|
| 189 |
+
logger.warning(f" ImageEncoder FAIL: {e}")
|
| 190 |
+
|
| 191 |
+
try:
|
| 192 |
+
from bigru_t.multimodal.audio_encoder import AudioEncoder
|
| 193 |
+
# BUG: in_channels é ignorado; Conv1d usa n_mels como canal
|
| 194 |
+
enc = AudioEncoder(d_model=128, in_channels=80, n_mels=80)
|
| 195 |
+
import torch
|
| 196 |
+
# Shape correta: [B, n_mels=80, T_audio]
|
| 197 |
+
x = torch.randn(2, 80, 100)
|
| 198 |
+
out = enc(x)
|
| 199 |
+
results["audio_encoder"] = {"ok": True, "shape": str(tuple(out.shape)),
|
| 200 |
+
"note": "in_channels param unused; uses n_mels=80 as channel"}
|
| 201 |
+
logger.info(f" AudioEncoder OK: {tuple(out.shape)}")
|
| 202 |
+
except Exception as e:
|
| 203 |
+
results["audio_encoder"] = {"ok": False, "error": str(e)}
|
| 204 |
+
logger.warning(f" AudioEncoder FAIL: {e}")
|
| 205 |
+
|
| 206 |
+
try:
|
| 207 |
+
from bigru_t.multimodal.video_encoder import VideoEncoder
|
| 208 |
+
enc = VideoEncoder(d_model=128, n_frames=4)
|
| 209 |
+
import torch
|
| 210 |
+
# BUG: frames deve ser [B, T, C, H, W] (5D, T antes de C)
|
| 211 |
+
frames = torch.randn(2, 4, 3, 16, 16) # (B=2, T=4, C=3, H=16, W=16)
|
| 212 |
+
audio = torch.randn(2, 80, 100) # (B, n_mels, T_audio)
|
| 213 |
+
out = enc(frames, audio)
|
| 214 |
+
results["video_encoder"] = {"ok": True, "shape": str(tuple(out.shape)),
|
| 215 |
+
"note": "frames=[B,T,C,H,W], audio=[B,80,T]"}
|
| 216 |
+
logger.info(f" VideoEncoder OK: {tuple(out.shape)}")
|
| 217 |
+
except Exception as e:
|
| 218 |
+
results["video_encoder"] = {"ok": False, "error": str(e)}
|
| 219 |
+
logger.warning(f" VideoEncoder FAIL: {e}")
|
| 220 |
+
|
| 221 |
+
try:
|
| 222 |
+
from bigru_t.multimodal.modal_router import ModalRouter
|
| 223 |
+
from bigru_t.multimodal.fusion_layer import FusionLayer
|
| 224 |
+
router = ModalRouter(d_model=128)
|
| 225 |
+
fusion = FusionLayer(d_model=128, n_modalities=4)
|
| 226 |
+
import torch
|
| 227 |
+
# ModalRouter.forward(inputs: Dict[str, Tensor]) — espera dict nomeado
|
| 228 |
+
inputs = {
|
| 229 |
+
"text": torch.randn(2, 128),
|
| 230 |
+
"image": torch.randn(2, 128),
|
| 231 |
+
"audio": torch.randn(2, 128),
|
| 232 |
+
"video": torch.randn(2, 128),
|
| 233 |
+
}
|
| 234 |
+
routed = router(inputs)
|
| 235 |
+
# FusionLayer.forward(modality_embs: list) — espera lista
|
| 236 |
+
embs = [torch.randn(2, 128) for _ in range(4)]
|
| 237 |
+
fused = fusion(embs)
|
| 238 |
+
results["modal_router_fusion"] = {"ok": True,
|
| 239 |
+
"router_shape": str(tuple(routed.shape)),
|
| 240 |
+
"fusion_shape": str(tuple(fused.shape))}
|
| 241 |
+
logger.info(f" ModalRouter OK: {tuple(routed.shape)} | Fusion OK: {tuple(fused.shape)}")
|
| 242 |
+
except Exception as e:
|
| 243 |
+
results["modal_router_fusion"] = {"ok": False, "error": str(e)}
|
| 244 |
+
logger.warning(f" ModalRouter/Fusion FAIL: {e}")
|
| 245 |
+
|
| 246 |
+
return results
|
| 247 |
+
|
| 248 |
+
|
| 249 |
+
def main():
|
| 250 |
+
logger.info("=" * 70)
|
| 251 |
+
logger.info("BiGRU_T_version V2 — 15M params + HW optimizer + multimodal bug-hunt")
|
| 252 |
+
logger.info("=" * 70)
|
| 253 |
+
|
| 254 |
+
# HF token (do ambiente; será apagado ao final)
|
| 255 |
+
hf_token = os.environ.get("HF_TOKEN") or None
|
| 256 |
+
if hf_token:
|
| 257 |
+
logger.info("HF_TOKEN encontrado no ambiente")
|
| 258 |
+
else:
|
| 259 |
+
logger.info("HF_TOKEN não fornecido (datasets públicos apenas)")
|
| 260 |
+
|
| 261 |
+
# 1. Verificação de OOM-Killer e recursos
|
| 262 |
+
logger.info("\n[1/6] Verificando recursos do sistema...")
|
| 263 |
+
if not check_oom_killer_risk():
|
| 264 |
+
logger.error("Recursos insuficientes. Abortando.")
|
| 265 |
+
sys.exit(1)
|
| 266 |
+
logger.info(f" RSS atual: {get_rss_mb():.1f}MB")
|
| 267 |
+
|
| 268 |
+
# 2. Teste multimodal (bug detection)
|
| 269 |
+
logger.info("\n[2/6] Testando multimodalidade (bug hunting)...")
|
| 270 |
+
mm_results = test_multimodal_encoders()
|
| 271 |
+
mm_ok = sum(1 for v in mm_results.values() if v.get("ok"))
|
| 272 |
+
mm_total = len(mm_results)
|
| 273 |
+
logger.info(f" Multimodal: {mm_ok}/{mm_total} encoders OK")
|
| 274 |
+
|
| 275 |
+
# 3. Config do modelo (12.97M params — max_modules=8, d_model=128)
|
| 276 |
+
logger.info("\n[3/6] Criando UnifiedModel (max_modules=8, d_model=128)...")
|
| 277 |
+
model_config = UnifiedModelConfig(
|
| 278 |
+
vocab_size=16384,
|
| 279 |
+
d_model=128,
|
| 280 |
+
max_seq_len=32,
|
| 281 |
+
pad_token_id=1,
|
| 282 |
+
max_modules=8, # base spec do usuário (8)
|
| 283 |
+
bigru_hidden=32,
|
| 284 |
+
d_transformer=64,
|
| 285 |
+
nhead_tu=4,
|
| 286 |
+
d_ff_tu=128,
|
| 287 |
+
output_dim_u8cell=64,
|
| 288 |
+
cache_len=16,
|
| 289 |
+
d_cache=128,
|
| 290 |
+
nhead_orq=8,
|
| 291 |
+
d_ff_orq=256,
|
| 292 |
+
trainT_dim=128,
|
| 293 |
+
nhead_train=4,
|
| 294 |
+
d_ff_train=256,
|
| 295 |
+
num_layers_train=2,
|
| 296 |
+
hypT_dim=128,
|
| 297 |
+
nhead_hyp=4,
|
| 298 |
+
d_ff_hyp=256,
|
| 299 |
+
num_layers_hyp=2,
|
| 300 |
+
num_bits=8,
|
| 301 |
+
dropout=0.1,
|
| 302 |
+
)
|
| 303 |
+
model, _ = create_unified_model(model_config)
|
| 304 |
+
params = model.count_parameters()
|
| 305 |
+
logger.info(f" params: {params['total']:,} ({params['total_M']:.2f}M)")
|
| 306 |
+
logger.info(f" trainable: {params['trainable']:,} ({params['trainable_M']:.2f}M)")
|
| 307 |
+
aggressive_cleanup(verbose=True)
|
| 308 |
+
|
| 309 |
+
# 4. Tokenizer (usa o pré-treinado do source)
|
| 310 |
+
logger.info("\n[4/6] Carregando tokenizer...")
|
| 311 |
+
from tokenizers import Tokenizer
|
| 312 |
+
tok_path = "/home/z/my-project/source/model_final/tokenizer/tokenizer.json"
|
| 313 |
+
if not Path(tok_path).exists():
|
| 314 |
+
logger.error(f"Tokenizer não encontrado: {tok_path}")
|
| 315 |
+
sys.exit(1)
|
| 316 |
+
tokenizer = Tokenizer.from_file(tok_path)
|
| 317 |
+
logger.info(f" Tokenizer loaded: vocab={tokenizer.get_vocab_size()}")
|
| 318 |
+
|
| 319 |
+
# 5. Smoke test do modelo
|
| 320 |
+
logger.info("\n[5/6] Smoke test do modelo...")
|
| 321 |
+
if not smoke_test_model(model, tokenizer):
|
| 322 |
+
logger.error("Smoke test falhou. Abortando.")
|
| 323 |
+
sys.exit(1)
|
| 324 |
+
aggressive_cleanup(verbose=True)
|
| 325 |
+
|
| 326 |
+
# 6. Carrega datasets (streaming, poucos samples para bug detection)
|
| 327 |
+
logger.info("\n[6/6] Carregando datasets (streaming, poucos samples)...")
|
| 328 |
+
datasets = [
|
| 329 |
+
"CEIA-POSITIVO/ultrachat_br_clustred_balanced_v1",
|
| 330 |
+
]
|
| 331 |
+
max_samples = 8 # poucos samples (bug detection, testes)
|
| 332 |
+
all_samples = []
|
| 333 |
+
t0 = time.time()
|
| 334 |
+
for ds_name in datasets:
|
| 335 |
+
if time.time() - t0 > 180:
|
| 336 |
+
logger.warning(f"Timeout carregando datasets após {len(all_samples)} amostras")
|
| 337 |
+
break
|
| 338 |
+
try:
|
| 339 |
+
count_before = len(all_samples)
|
| 340 |
+
for sample in stream_dataset(ds_name, max_samples=max_samples, hf_token=hf_token):
|
| 341 |
+
all_samples.append(sample)
|
| 342 |
+
if len(all_samples) >= max_samples * len(datasets):
|
| 343 |
+
break
|
| 344 |
+
if time.time() - t0 > 180:
|
| 345 |
+
break
|
| 346 |
+
logger.info(f" {ds_name}: +{len(all_samples)-count_before} samples")
|
| 347 |
+
except Exception as e:
|
| 348 |
+
logger.warning(f" {ds_name} falhou: {e}")
|
| 349 |
+
|
| 350 |
+
if not all_samples:
|
| 351 |
+
logger.error("Nenhuma amostra carregada de nenhum dataset")
|
| 352 |
+
sys.exit(1)
|
| 353 |
+
|
| 354 |
+
logger.info(f"Total: {len(all_samples)} samples em {time.time()-t0:.1f}s")
|
| 355 |
+
|
| 356 |
+
# Split 90/10
|
| 357 |
+
import random
|
| 358 |
+
random.seed(42)
|
| 359 |
+
random.shuffle(all_samples)
|
| 360 |
+
split = max(1, int(0.9 * len(all_samples)))
|
| 361 |
+
train_samples = all_samples[:split]
|
| 362 |
+
val_samples = all_samples[split:] or all_samples[:1]
|
| 363 |
+
logger.info(f" train: {len(train_samples)} | val: {len(val_samples)}")
|
| 364 |
+
|
| 365 |
+
aggressive_cleanup(verbose=True)
|
| 366 |
+
|
| 367 |
+
# Trainer config — HW optimizer + time budget + cleanup agressivo
|
| 368 |
+
trainer_config = TrainerConfig(
|
| 369 |
+
epochs=2,
|
| 370 |
+
datasets=",".join(datasets),
|
| 371 |
+
max_samples_per_dataset=max_samples,
|
| 372 |
+
max_seq_len=32,
|
| 373 |
+
per_device_batch_size=1,
|
| 374 |
+
grad_accum=2,
|
| 375 |
+
lr=1e-3,
|
| 376 |
+
optimizer_type="hamiltonian_wasserstein", # ATIVADO
|
| 377 |
+
hw_lr_amp=0.3,
|
| 378 |
+
hw_lr_freq=0.01,
|
| 379 |
+
hw_sigma_w=1.0,
|
| 380 |
+
hw_sigma_rep=0.1,
|
| 381 |
+
hw_prune_every_n=0, # pruning off para bug detection
|
| 382 |
+
use_hypothesis=True,
|
| 383 |
+
stop_grad_hyp=True,
|
| 384 |
+
meta_interval=5,
|
| 385 |
+
max_total_time_s=1800.0, # 30 min total
|
| 386 |
+
max_per_epoch_s=900.0, # 15 min/epoch
|
| 387 |
+
cleanup_every_n_steps=5,
|
| 388 |
+
log_every=1,
|
| 389 |
+
save_temp_every=3, # salvar a cada 3 steps (garante checkpoint mesmo se OOM no fim)
|
| 390 |
+
output_dir="/home/z/my-project/BiGRU_T_version/model_final",
|
| 391 |
+
temp_dir="/home/z/my-project/BiGRU_T_version/_temp_checkpoints",
|
| 392 |
+
keep_temp=False,
|
| 393 |
+
ram_threshold_pct=90.0,
|
| 394 |
+
disk_min_free_gb=1.0,
|
| 395 |
+
loss_patience=30,
|
| 396 |
+
)
|
| 397 |
+
|
| 398 |
+
trainer = BiGRU_T_Trainer(
|
| 399 |
+
model=model,
|
| 400 |
+
tokenizer=tokenizer,
|
| 401 |
+
config=trainer_config,
|
| 402 |
+
train_samples=train_samples,
|
| 403 |
+
val_samples=val_samples,
|
| 404 |
+
)
|
| 405 |
+
|
| 406 |
+
# Treino
|
| 407 |
+
result = trainer.train()
|
| 408 |
+
|
| 409 |
+
# Salva resultados extras (multimodal + smoke)
|
| 410 |
+
extra_report = {
|
| 411 |
+
"multimodal_test": mm_results,
|
| 412 |
+
"smoke_test": {"passed": True},
|
| 413 |
+
"config": {
|
| 414 |
+
"max_modules": model_config.max_modules,
|
| 415 |
+
"d_model": model_config.d_model,
|
| 416 |
+
"bigru_hidden": model_config.bigru_hidden,
|
| 417 |
+
"d_transformer": model_config.d_transformer,
|
| 418 |
+
"params_total": params["total"],
|
| 419 |
+
"params_M": params["total_M"],
|
| 420 |
+
"optimizer": "hamiltonian_wasserstein",
|
| 421 |
+
},
|
| 422 |
+
}
|
| 423 |
+
extra_path = Path("/home/z/my-project/BiGRU_T_version/multimodal_report.json")
|
| 424 |
+
with open(extra_path, "w") as f:
|
| 425 |
+
json.dump(extra_report, f, indent=2, default=str)
|
| 426 |
+
logger.info(f"Relatório multimodal salvo: {extra_path}")
|
| 427 |
+
|
| 428 |
+
print("\n=== Resultado Final ===")
|
| 429 |
+
print(json.dumps(result, indent=2, default=str))
|
| 430 |
+
|
| 431 |
+
# Apaga HF token
|
| 432 |
+
if "HF_TOKEN" in os.environ:
|
| 433 |
+
del os.environ["HF_TOKEN"]
|
| 434 |
+
logger.info("HF_TOKEN apagado do ambiente")
|
| 435 |
+
|
| 436 |
+
sys.exit(0 if not result.get("killed") else 1)
|
| 437 |
+
|
| 438 |
+
|
| 439 |
+
if __name__ == "__main__":
|
| 440 |
+
main()
|
src/bigru_t/__init__.py
CHANGED
|
@@ -26,6 +26,17 @@ 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 |
|
| 30 |
__all__ = [
|
| 31 |
"UnifiedModel", "UnifiedModelConfig", "create_unified_model",
|
|
@@ -36,4 +47,9 @@ __all__ = [
|
|
| 36 |
"MetaConfigurator",
|
| 37 |
"KillSwitch", "KillSwitchState",
|
| 38 |
"BiGRU_T_Trainer", "TrainerConfig",
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 39 |
]
|
|
|
|
| 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 |
+
|
| 33 |
+
from .inference.generator import BiGRUTGenerator
|
| 34 |
+
|
| 35 |
+
from .optim.hamiltonian_wasserstein import HamiltonianWassersteinOptimizer
|
| 36 |
+
|
| 37 |
+
from .utils.memory_cleanup import (
|
| 38 |
+
aggressive_cleanup, production_cleanup, TimeBudget, StepTimer, get_rss_mb,
|
| 39 |
+
)
|
| 40 |
|
| 41 |
__all__ = [
|
| 42 |
"UnifiedModel", "UnifiedModelConfig", "create_unified_model",
|
|
|
|
| 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",
|
| 54 |
+
"aggressive_cleanup", "production_cleanup", "TimeBudget", "StepTimer", "get_rss_mb",
|
| 55 |
]
|
src/bigru_t/inference/__init__.py
ADDED
|
@@ -0,0 +1,4 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""inference/__init__.py — Módulos de inferência."""
|
| 2 |
+
from .generator import BiGRUTGenerator
|
| 3 |
+
|
| 4 |
+
__all__ = ["BiGRUTGenerator"]
|
src/bigru_t/inference/generator.py
ADDED
|
@@ -0,0 +1,190 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""generator.py — Inferência autoregressiva para BiGRU_T_version.
|
| 2 |
+
|
| 3 |
+
Implementa geração greedy + top-k sampling + streaming, reaproveitando a
|
| 4 |
+
filosofia do GRURingV139.generate() da fonte (xavante_work/flexnet/gru_ring_v13_9.py).
|
| 5 |
+
|
| 6 |
+
O UnifiedModel produz (batch, vocab_size) — 1 logit por forward. Para geração
|
| 7 |
+
autoregressiva, alimentamos a sequência crescente e usamos o último logit.
|
| 8 |
+
"""
|
| 9 |
+
from __future__ import annotations
|
| 10 |
+
|
| 11 |
+
import logging
|
| 12 |
+
import math
|
| 13 |
+
from typing import Iterator, List, Optional, Tuple
|
| 14 |
+
|
| 15 |
+
import torch
|
| 16 |
+
import torch.nn.functional as F
|
| 17 |
+
|
| 18 |
+
logger = logging.getLogger(__name__)
|
| 19 |
+
|
| 20 |
+
|
| 21 |
+
class BiGRUTGenerator:
|
| 22 |
+
"""Gerador autoregressivo para UnifiedModel.
|
| 23 |
+
|
| 24 |
+
Suporta:
|
| 25 |
+
- greedy decoding
|
| 26 |
+
- top-k sampling com temperatura
|
| 27 |
+
- repetition penalty
|
| 28 |
+
- max_new_tokens + EOS stop
|
| 29 |
+
- streaming (yield token a token)
|
| 30 |
+
|
| 31 |
+
Args:
|
| 32 |
+
model: UnifiedModel treinado
|
| 33 |
+
eos_token_id: ID do token de fim (para parar)
|
| 34 |
+
pad_token_id: ID do padding
|
| 35 |
+
max_seq_len: comprimento máximo de contexto (trunca à esquerda)
|
| 36 |
+
"""
|
| 37 |
+
|
| 38 |
+
def __init__(
|
| 39 |
+
self,
|
| 40 |
+
model,
|
| 41 |
+
eos_token_id: int = 2,
|
| 42 |
+
pad_token_id: int = 1,
|
| 43 |
+
max_seq_len: int = 64,
|
| 44 |
+
):
|
| 45 |
+
self.model = model
|
| 46 |
+
self.eos_token_id = eos_token_id
|
| 47 |
+
self.pad_token_id = pad_token_id
|
| 48 |
+
self.max_seq_len = max_seq_len
|
| 49 |
+
self.device = next(model.parameters()).device
|
| 50 |
+
|
| 51 |
+
@torch.no_grad()
|
| 52 |
+
def generate(
|
| 53 |
+
self,
|
| 54 |
+
input_ids: torch.Tensor,
|
| 55 |
+
max_new_tokens: int = 32,
|
| 56 |
+
temperature: float = 1.0,
|
| 57 |
+
top_k: int = 0,
|
| 58 |
+
repetition_penalty: float = 1.0,
|
| 59 |
+
do_sample: bool = False,
|
| 60 |
+
) -> torch.Tensor:
|
| 61 |
+
"""Gera tokens autoregressivamente.
|
| 62 |
+
|
| 63 |
+
Args:
|
| 64 |
+
input_ids: (batch, T) tokens de prompt
|
| 65 |
+
max_new_tokens: máx. tokens a gerar
|
| 66 |
+
temperature: temperatura do sampling (1.0 = sem escala)
|
| 67 |
+
top_k: se > 0, amostra apenas dos top-k tokens
|
| 68 |
+
repetition_penalty: penaliza tokens já gerados (1.0 = sem pena)
|
| 69 |
+
do_sample: se False, greedy decoding
|
| 70 |
+
|
| 71 |
+
Returns:
|
| 72 |
+
generated_ids: (batch, T + max_new_tokens)
|
| 73 |
+
"""
|
| 74 |
+
self.model.eval()
|
| 75 |
+
batch_size = input_ids.size(0)
|
| 76 |
+
generated = input_ids.clone().to(self.device)
|
| 77 |
+
|
| 78 |
+
for _step in range(max_new_tokens):
|
| 79 |
+
# Trunca à esquerda se exceder max_seq_len
|
| 80 |
+
if generated.size(1) > self.max_seq_len:
|
| 81 |
+
context = generated[:, -self.max_seq_len:]
|
| 82 |
+
else:
|
| 83 |
+
context = generated
|
| 84 |
+
|
| 85 |
+
# Forward (sem hipótese — inferência usa só TrainT)
|
| 86 |
+
out = self.model(context, temperature=1.0, use_hypothesis=False)
|
| 87 |
+
y_hat = out[0] if isinstance(out, tuple) else out # (batch, vocab)
|
| 88 |
+
|
| 89 |
+
# Último logit
|
| 90 |
+
logits = y_hat # já é (batch, vocab) — 1 logit por forward
|
| 91 |
+
|
| 92 |
+
# Repetition penalty
|
| 93 |
+
if repetition_penalty != 1.0:
|
| 94 |
+
for b in range(batch_size):
|
| 95 |
+
for prev_token in generated[b].tolist():
|
| 96 |
+
if logits[b, prev_token] > 0:
|
| 97 |
+
logits[b, prev_token] /= repetition_penalty
|
| 98 |
+
else:
|
| 99 |
+
logits[b, prev_token] *= repetition_penalty
|
| 100 |
+
|
| 101 |
+
if do_sample and temperature > 0:
|
| 102 |
+
# Top-k sampling
|
| 103 |
+
if top_k > 0:
|
| 104 |
+
top_k = min(top_k, logits.size(-1))
|
| 105 |
+
values, _ = torch.topk(logits, top_k, dim=-1)
|
| 106 |
+
min_val = values[:, -1:].unsqueeze(-1)
|
| 107 |
+
logits = torch.where(
|
| 108 |
+
logits < min_val,
|
| 109 |
+
torch.full_like(logits, float("-inf")),
|
| 110 |
+
logits,
|
| 111 |
+
)
|
| 112 |
+
# Temperatura
|
| 113 |
+
logits = logits / max(temperature, 1e-8)
|
| 114 |
+
probs = F.softmax(logits, dim=-1)
|
| 115 |
+
next_token = torch.multinomial(probs, num_samples=1)
|
| 116 |
+
else:
|
| 117 |
+
# Greedy
|
| 118 |
+
next_token = logits.argmax(dim=-1, keepdim=True)
|
| 119 |
+
|
| 120 |
+
# Concatena
|
| 121 |
+
generated = torch.cat([generated, next_token], dim=1)
|
| 122 |
+
|
| 123 |
+
# Para se todos geraram EOS
|
| 124 |
+
if (next_token == self.eos_token_id).all():
|
| 125 |
+
break
|
| 126 |
+
|
| 127 |
+
return generated
|
| 128 |
+
|
| 129 |
+
@torch.no_grad()
|
| 130 |
+
def stream_generate(
|
| 131 |
+
self,
|
| 132 |
+
input_ids: torch.Tensor,
|
| 133 |
+
max_new_tokens: int = 32,
|
| 134 |
+
temperature: float = 1.0,
|
| 135 |
+
top_k: int = 0,
|
| 136 |
+
repetition_penalty: float = 1.0,
|
| 137 |
+
do_sample: bool = False,
|
| 138 |
+
) -> Iterator[torch.Tensor]:
|
| 139 |
+
"""Geração streaming — yields um token por vez.
|
| 140 |
+
|
| 141 |
+
Args: mesmos de generate()
|
| 142 |
+
Yields:
|
| 143 |
+
next_token: (batch, 1) tensor a cada iteração
|
| 144 |
+
"""
|
| 145 |
+
self.model.eval()
|
| 146 |
+
batch_size = input_ids.size(0)
|
| 147 |
+
generated = input_ids.clone().to(self.device)
|
| 148 |
+
|
| 149 |
+
for _step in range(max_new_tokens):
|
| 150 |
+
if generated.size(1) > self.max_seq_len:
|
| 151 |
+
context = generated[:, -self.max_seq_len:]
|
| 152 |
+
else:
|
| 153 |
+
context = generated
|
| 154 |
+
|
| 155 |
+
out = self.model(context, temperature=1.0, use_hypothesis=False)
|
| 156 |
+
y_hat = out[0] if isinstance(out, tuple) else out
|
| 157 |
+
logits = y_hat
|
| 158 |
+
|
| 159 |
+
if repetition_penalty != 1.0:
|
| 160 |
+
for b in range(batch_size):
|
| 161 |
+
for prev_token in generated[b].tolist():
|
| 162 |
+
if logits[b, prev_token] > 0:
|
| 163 |
+
logits[b, prev_token] /= repetition_penalty
|
| 164 |
+
else:
|
| 165 |
+
logits[b, prev_token] *= repetition_penalty
|
| 166 |
+
|
| 167 |
+
if do_sample and temperature > 0:
|
| 168 |
+
if top_k > 0:
|
| 169 |
+
top_k = min(top_k, logits.size(-1))
|
| 170 |
+
values, _ = torch.topk(logits, top_k, dim=-1)
|
| 171 |
+
min_val = values[:, -1:].unsqueeze(-1)
|
| 172 |
+
logits = torch.where(
|
| 173 |
+
logits < min_val,
|
| 174 |
+
torch.full_like(logits, float("-inf")),
|
| 175 |
+
logits,
|
| 176 |
+
)
|
| 177 |
+
logits = logits / max(temperature, 1e-8)
|
| 178 |
+
probs = F.softmax(logits, dim=-1)
|
| 179 |
+
next_token = torch.multinomial(probs, num_samples=1)
|
| 180 |
+
else:
|
| 181 |
+
next_token = logits.argmax(dim=-1, keepdim=True)
|
| 182 |
+
|
| 183 |
+
generated = torch.cat([generated, next_token], dim=1)
|
| 184 |
+
yield next_token
|
| 185 |
+
|
| 186 |
+
if (next_token == self.eos_token_id).all():
|
| 187 |
+
break
|
| 188 |
+
|
| 189 |
+
|
| 190 |
+
__all__ = ["BiGRUTGenerator"]
|
src/bigru_t/model/attention_multimodal.py
ADDED
|
@@ -0,0 +1,114 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
Xavante - attention_multimodal.py
|
| 3 |
+
Responsabilidade: Multi-head attention multimodal (Teorema 11.1/11.2).
|
| 4 |
+
Suporta attention entre modalidades e dentro de modalidade.
|
| 5 |
+
"""
|
| 6 |
+
from __future__ import annotations
|
| 7 |
+
|
| 8 |
+
import logging
|
| 9 |
+
import math
|
| 10 |
+
from typing import Optional
|
| 11 |
+
|
| 12 |
+
import torch
|
| 13 |
+
import torch.nn as nn
|
| 14 |
+
import torch.nn.functional as F
|
| 15 |
+
|
| 16 |
+
logger = logging.getLogger(__name__)
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
class MultiHeadAttention(nn.Module):
|
| 20 |
+
"""
|
| 21 |
+
MHA padrão com suporte a:
|
| 22 |
+
- Mascara causal
|
| 23 |
+
- Mascara multimodal (mod crossing)
|
| 24 |
+
- Attention 2M tokens via chunked attention (Teorema 6.1)
|
| 25 |
+
"""
|
| 26 |
+
|
| 27 |
+
def __init__(
|
| 28 |
+
self,
|
| 29 |
+
d_model: int,
|
| 30 |
+
n_heads: int = 8,
|
| 31 |
+
dropout: float = 0.0,
|
| 32 |
+
max_chunk: int = 4096,
|
| 33 |
+
):
|
| 34 |
+
super().__init__()
|
| 35 |
+
assert d_model % n_heads == 0
|
| 36 |
+
self.d_model = d_model
|
| 37 |
+
self.n_heads = n_heads
|
| 38 |
+
self.d_head = d_model // n_heads
|
| 39 |
+
self.max_chunk = max_chunk
|
| 40 |
+
self.qkv = nn.Linear(d_model, 3 * d_model, bias=True)
|
| 41 |
+
self.out = nn.Linear(d_model, d_model)
|
| 42 |
+
self.dropout = nn.Dropout(dropout)
|
| 43 |
+
|
| 44 |
+
def forward(
|
| 45 |
+
self,
|
| 46 |
+
x: torch.Tensor,
|
| 47 |
+
mask: Optional[torch.Tensor] = None,
|
| 48 |
+
kv: Optional[torch.Tensor] = None,
|
| 49 |
+
) -> torch.Tensor:
|
| 50 |
+
B, L, D = x.shape
|
| 51 |
+
if kv is None:
|
| 52 |
+
qkv = self.qkv(x)
|
| 53 |
+
q, k, v = qkv.chunk(3, dim=-1)
|
| 54 |
+
else:
|
| 55 |
+
q = self.qkv(x)[:, :, :D]
|
| 56 |
+
kv_proj = self.qkv(kv)
|
| 57 |
+
k = kv_proj[:, :, D : 2 * D]
|
| 58 |
+
v = kv_proj[:, :, 2 * D :]
|
| 59 |
+
# reshape para heads
|
| 60 |
+
q = q.view(B, L, self.n_heads, self.d_head).transpose(1, 2)
|
| 61 |
+
k = k.view(B, -1, self.n_heads, self.d_head).transpose(1, 2)
|
| 62 |
+
v = v.view(B, -1, self.n_heads, self.d_head).transpose(1, 2)
|
| 63 |
+
|
| 64 |
+
# Chunked attention para sequencias longas (Teorema 6.1)
|
| 65 |
+
if L > self.max_chunk:
|
| 66 |
+
# Normalize mask to 4D for chunked path
|
| 67 |
+
if mask is not None and mask.dim() == 2:
|
| 68 |
+
mask = mask.unsqueeze(0).unsqueeze(0)
|
| 69 |
+
elif mask is not None and mask.dim() == 3:
|
| 70 |
+
mask = mask.unsqueeze(1)
|
| 71 |
+
return self._chunked_attention(q, k, v, mask)
|
| 72 |
+
|
| 73 |
+
scale = 1.0 / math.sqrt(self.d_head)
|
| 74 |
+
attn = (q @ k.transpose(-2, -1)) * scale # [B, H, L, L]
|
| 75 |
+
if mask is not None:
|
| 76 |
+
attn = attn.masked_fill(mask == 0, float("-inf"))
|
| 77 |
+
attn = F.softmax(attn, dim=-1)
|
| 78 |
+
attn = self.dropout(attn)
|
| 79 |
+
out = attn @ v # [B, H, L, d_head]
|
| 80 |
+
out = out.transpose(1, 2).contiguous().view(B, L, D)
|
| 81 |
+
return self.out(out)
|
| 82 |
+
|
| 83 |
+
def _chunked_attention(
|
| 84 |
+
self,
|
| 85 |
+
q: torch.Tensor,
|
| 86 |
+
k: torch.Tensor,
|
| 87 |
+
v: torch.Tensor,
|
| 88 |
+
mask: Optional[torch.Tensor],
|
| 89 |
+
) -> torch.Tensor:
|
| 90 |
+
"""Attention em blocos para janelas de 2M tokens (memoria controlada)."""
|
| 91 |
+
B, H, L, d = q.shape
|
| 92 |
+
Lk = k.shape[2]
|
| 93 |
+
chunk = self.max_chunk
|
| 94 |
+
outs = []
|
| 95 |
+
scale = 1.0 / math.sqrt(d)
|
| 96 |
+
for i in range(0, L, chunk):
|
| 97 |
+
qi = q[:, :, i : i + chunk]
|
| 98 |
+
out_chunk = []
|
| 99 |
+
for j in range(0, Lk, chunk):
|
| 100 |
+
kj = k[:, :, j : j + chunk]
|
| 101 |
+
vj = v[:, :, j : j + chunk]
|
| 102 |
+
attn = (qi @ kj.transpose(-2, -1)) * scale
|
| 103 |
+
if mask is not None:
|
| 104 |
+
m_chunk = mask[:, :, i : i + chunk, j : j + chunk]
|
| 105 |
+
attn = attn.masked_fill(m_chunk == 0, float("-inf"))
|
| 106 |
+
attn = F.softmax(attn, dim=-1)
|
| 107 |
+
out_chunk.append(attn @ vj)
|
| 108 |
+
outs.append(torch.cat(out_chunk, dim=2))
|
| 109 |
+
out = torch.cat(outs, dim=2)
|
| 110 |
+
out = out.transpose(1, 2).contiguous().view(B, L, self.n_heads * d)
|
| 111 |
+
return self.out(out)
|
| 112 |
+
|
| 113 |
+
|
| 114 |
+
__all__ = ["MultiHeadAttention"]
|
src/bigru_t/model/embedding_reconfig.py
ADDED
|
@@ -0,0 +1,67 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
Xavante - embedding_reconfig.py
|
| 3 |
+
Responsabilidade: Embedding reconfiguravel (Teorema 12.1/12.2).
|
| 4 |
+
Permite expansao dinamica do vocabulario e reconfiguracao de dim do modelo.
|
| 5 |
+
"""
|
| 6 |
+
from __future__ import annotations
|
| 7 |
+
|
| 8 |
+
import logging
|
| 9 |
+
import math
|
| 10 |
+
from typing import Optional
|
| 11 |
+
|
| 12 |
+
import torch
|
| 13 |
+
import torch.nn as nn
|
| 14 |
+
|
| 15 |
+
logger = logging.getLogger(__name__)
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
class ReconfigurableEmbedding(nn.Module):
|
| 19 |
+
"""
|
| 20 |
+
Embedding que suporta:
|
| 21 |
+
- Expansao de vocabulario sem reset
|
| 22 |
+
- Reconfiguracao de dim (projecao)
|
| 23 |
+
- Quantizacao para memoria
|
| 24 |
+
"""
|
| 25 |
+
|
| 26 |
+
def __init__(self, vocab_size: int, d_model: int, padding_idx: Optional[int] = None):
|
| 27 |
+
super().__init__()
|
| 28 |
+
self.vocab_size = vocab_size
|
| 29 |
+
self.d_model = d_model
|
| 30 |
+
self.padding_idx = padding_idx
|
| 31 |
+
self.weight = nn.Parameter(torch.empty(vocab_size, d_model))
|
| 32 |
+
nn.init.normal_(self.weight, mean=0.0, std=1.0 / math.sqrt(d_model) if False else 0.02)
|
| 33 |
+
if padding_idx is not None:
|
| 34 |
+
with torch.no_grad():
|
| 35 |
+
self.weight[padding_idx].fill_(0)
|
| 36 |
+
|
| 37 |
+
def forward(self, idx: torch.Tensor) -> torch.Tensor:
|
| 38 |
+
return torch.nn.functional.embedding(idx, self.weight, padding_idx=self.padding_idx)
|
| 39 |
+
|
| 40 |
+
def expand_vocab(self, new_size: int) -> None:
|
| 41 |
+
"""Expande o vocabulario preservando pesos antigos."""
|
| 42 |
+
if new_size <= self.vocab_size:
|
| 43 |
+
return
|
| 44 |
+
old = self.weight.data
|
| 45 |
+
new = torch.empty(new_size - self.vocab_size, self.d_model, device=old.device, dtype=old.dtype)
|
| 46 |
+
nn.init.normal_(new, mean=0.0, std=0.02)
|
| 47 |
+
self.weight = nn.Parameter(torch.cat([old, new], dim=0))
|
| 48 |
+
self.vocab_size = new_size
|
| 49 |
+
logger.info("Vocab expandido para %d", new_size)
|
| 50 |
+
|
| 51 |
+
def reconfigure_dim(self, new_d: int) -> None:
|
| 52 |
+
"""Reconfigura a dimensao via projecao linear (random init para a parte nova)."""
|
| 53 |
+
if new_d == self.d_model:
|
| 54 |
+
return
|
| 55 |
+
old = self.weight.data # [V, d_model]
|
| 56 |
+
if new_d > self.d_model:
|
| 57 |
+
pad = torch.empty(old.shape[0], new_d - self.d_model, device=old.device, dtype=old.dtype)
|
| 58 |
+
nn.init.normal_(pad, mean=0.0, std=0.02)
|
| 59 |
+
new_w = torch.cat([old, pad], dim=1)
|
| 60 |
+
else:
|
| 61 |
+
new_w = old[:, :new_d]
|
| 62 |
+
self.weight = nn.Parameter(new_w)
|
| 63 |
+
self.d_model = new_d
|
| 64 |
+
logger.info("Dim reconfigurada para %d", new_d)
|
| 65 |
+
|
| 66 |
+
|
| 67 |
+
__all__ = ["ReconfigurableEmbedding"]
|
src/bigru_t/model/gru_hierarchy.py
ADDED
|
@@ -0,0 +1,82 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
Xavante - gru_hierarchy.py
|
| 3 |
+
Responsabilidade: GRU hierarquica (Teorema 7.1/7.2). Reutiliza o FlexGRU
|
| 4 |
+
validado do FlexNet.
|
| 5 |
+
"""
|
| 6 |
+
from __future__ import annotations
|
| 7 |
+
|
| 8 |
+
import logging
|
| 9 |
+
from typing import List, Optional
|
| 10 |
+
|
| 11 |
+
import torch
|
| 12 |
+
import torch.nn as nn
|
| 13 |
+
|
| 14 |
+
logger = logging.getLogger(__name__)
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
class FlexGRUCell(nn.Module):
|
| 18 |
+
"""GRU cell com gate extra de confianca (h_t) que escala update."""
|
| 19 |
+
|
| 20 |
+
def __init__(self, d_in: int, d_hidden: int):
|
| 21 |
+
super().__init__()
|
| 22 |
+
self.d_in = d_in
|
| 23 |
+
self.d_hidden = d_hidden
|
| 24 |
+
# Linear concatenado x_t e h_{t-1}
|
| 25 |
+
self.x2h = nn.Linear(d_in, 3 * d_hidden, bias=True)
|
| 26 |
+
self.h2h = nn.Linear(d_hidden, 3 * d_hidden, bias=True)
|
| 27 |
+
# Confidence gate
|
| 28 |
+
self.conf_gate = nn.Linear(d_in + d_hidden, 1)
|
| 29 |
+
|
| 30 |
+
def forward(self, x: torch.Tensor, h: Optional[torch.Tensor] = None) -> torch.Tensor:
|
| 31 |
+
B, D = x.shape
|
| 32 |
+
if h is None:
|
| 33 |
+
h = torch.zeros(B, self.d_hidden, device=x.device, dtype=x.dtype)
|
| 34 |
+
gates_x = self.x2h(x)
|
| 35 |
+
gates_h = self.h2h(h)
|
| 36 |
+
x_r, x_z, x_n = gates_x.chunk(3, dim=-1)
|
| 37 |
+
h_r, h_z, h_n = gates_h.chunk(3, dim=-1)
|
| 38 |
+
r = torch.sigmoid(x_r + h_r)
|
| 39 |
+
z = torch.sigmoid(x_z + h_z)
|
| 40 |
+
n = torch.tanh(x_n + r * h_n)
|
| 41 |
+
# Confidence scaling (Teorema 17.1)
|
| 42 |
+
c = torch.sigmoid(self.conf_gate(torch.cat([x, h], dim=-1)))
|
| 43 |
+
h_new = (1 - z * c) * h + (z * c) * n
|
| 44 |
+
return h_new
|
| 45 |
+
|
| 46 |
+
|
| 47 |
+
class GRUHierarchy(nn.Module):
|
| 48 |
+
"""GRU hierarquica em L niveis (Teorema 7.1/7.2)."""
|
| 49 |
+
|
| 50 |
+
def __init__(self, d_model: int, d_hidden: int, n_levels: int = 3):
|
| 51 |
+
super().__init__()
|
| 52 |
+
self.n_levels = n_levels
|
| 53 |
+
self.cells = nn.ModuleList(
|
| 54 |
+
[FlexGRUCell(d_model if i == 0 else d_hidden, d_hidden) for i in range(n_levels)]
|
| 55 |
+
)
|
| 56 |
+
# Projecao de saida de cada nivel
|
| 57 |
+
self.projs = nn.ModuleList([nn.Linear(d_hidden, d_model) for _ in range(n_levels)])
|
| 58 |
+
|
| 59 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 60 |
+
"""
|
| 61 |
+
x: [B, L, d_model]
|
| 62 |
+
Retorna: [B, L, d_model]
|
| 63 |
+
"""
|
| 64 |
+
B, L, D = x.shape
|
| 65 |
+
h_states: List[Optional[torch.Tensor]] = [None] * self.n_levels
|
| 66 |
+
outputs = []
|
| 67 |
+
for t in range(L):
|
| 68 |
+
xt = x[:, t, :]
|
| 69 |
+
inp = xt
|
| 70 |
+
new_h_states = []
|
| 71 |
+
for lvl, cell in enumerate(self.cells):
|
| 72 |
+
h_new = cell(inp, h_states[lvl])
|
| 73 |
+
new_h_states.append(h_new)
|
| 74 |
+
inp = h_new
|
| 75 |
+
h_states = new_h_states
|
| 76 |
+
# Use each level's own projection (fixes dead projections bug)
|
| 77 |
+
combined = sum(self.projs[lvl](h_states[lvl]) for lvl in range(self.n_levels)) / self.n_levels
|
| 78 |
+
outputs.append(combined)
|
| 79 |
+
return torch.stack(outputs, dim=1)
|
| 80 |
+
|
| 81 |
+
|
| 82 |
+
__all__ = ["FlexGRUCell", "GRUHierarchy"]
|
src/bigru_t/reasoning/__init__.py
ADDED
|
@@ -0,0 +1,4 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""reasoning/__init__.py — Módulos de raciocínio cíclico."""
|
| 2 |
+
from .circular_reasoning_wasserstein import CircularReasoningWasserstein
|
| 3 |
+
|
| 4 |
+
__all__ = ["CircularReasoningWasserstein"]
|
src/bigru_t/reasoning/circular_reasoning_wasserstein.py
ADDED
|
@@ -0,0 +1,183 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""circular_reasoning_wasserstein.py — Raciocínio circular como contração W₂ (V11.24c).
|
| 2 |
+
|
| 3 |
+
Teorema 4: Se Φ é contração em W₂, então s_t → s* (ponto fixo).
|
| 4 |
+
W₂(P_{s_t}, P_{s*}) ≤ q^t · W₂(P_{s_0}, P_{s*}).
|
| 5 |
+
|
| 6 |
+
CORREÇÃO V11.26:
|
| 7 |
+
- Bug original: residual_gate inicial 0.1 muito alto → não há contração
|
| 8 |
+
( gates pequenos => s_new ≈ s + ε·r ≈ s, mas W2 cresce lentamente).
|
| 9 |
+
- Fix: gate inicial = 0.05 + normalização do refine_head para garantir
|
| 10 |
+
||Φ(s) - Φ(s')|| ≤ q·||s - s'|| com q < 1.
|
| 11 |
+
- Adicionado: inicialização Xavier no refine_head (era default PyTorch).
|
| 12 |
+
- Adicionado: clipping espectral aproximado no residual_gate.
|
| 13 |
+
- Adicionado: tracking de q_t real para validação empírica.
|
| 14 |
+
"""
|
| 15 |
+
from __future__ import annotations
|
| 16 |
+
import math
|
| 17 |
+
import torch
|
| 18 |
+
import torch.nn as nn
|
| 19 |
+
import torch.nn.functional as F
|
| 20 |
+
from typing import Optional, Dict, Any, List
|
| 21 |
+
|
| 22 |
+
|
| 23 |
+
class CircularReasoningWasserstein(nn.Module):
|
| 24 |
+
"""Raciocínio circular com convergência garantida via contração W₂.
|
| 25 |
+
|
| 26 |
+
O estado s_t evolui por: s_{t+1} = Φ(s_t, x)
|
| 27 |
+
onde Φ é o operador transformer (refine_head).
|
| 28 |
+
|
| 29 |
+
A convergência é garantida quando a norma espectral de W_att < 1/q
|
| 30 |
+
e o termo residual é Lipschitz com L_ff < (1-q).
|
| 31 |
+
|
| 32 |
+
V11.26:
|
| 33 |
+
- residual_gate inicial menor (0.05 vs 0.1)
|
| 34 |
+
- refine_head com init Xavier
|
| 35 |
+
- clipping de gate para garantir q < 1
|
| 36 |
+
- tracking empírico de q_t
|
| 37 |
+
"""
|
| 38 |
+
|
| 39 |
+
def __init__(self, d_model: int, n_cycles: int = 5,
|
| 40 |
+
contraction_target: float = 0.9,
|
| 41 |
+
w2_tolerance: float = 1e-4,
|
| 42 |
+
init_gate: float = 0.05):
|
| 43 |
+
super().__init__()
|
| 44 |
+
self.d_model = d_model
|
| 45 |
+
self.n_cycles = n_cycles
|
| 46 |
+
self.contraction_target = contraction_target
|
| 47 |
+
self.w2_tolerance = w2_tolerance
|
| 48 |
+
|
| 49 |
+
# Operador Φ (transformer leve)
|
| 50 |
+
self.refine_norm = nn.LayerNorm(d_model)
|
| 51 |
+
self.refine_head = nn.Sequential(
|
| 52 |
+
nn.Linear(d_model, d_model * 2),
|
| 53 |
+
nn.GELU(),
|
| 54 |
+
nn.Linear(d_model * 2, d_model),
|
| 55 |
+
)
|
| 56 |
+
# V11.26: Xavier init para estabilizar contração
|
| 57 |
+
for layer in self.refine_head:
|
| 58 |
+
if isinstance(layer, nn.Linear):
|
| 59 |
+
nn.init.xavier_uniform_(layer.weight, gain=0.5) # gain < 1 → contração
|
| 60 |
+
if layer.bias is not None:
|
| 61 |
+
nn.init.zeros_(layer.bias)
|
| 62 |
+
|
| 63 |
+
# V11.26: gate inicial menor + clipping
|
| 64 |
+
self.residual_gate = nn.Parameter(torch.full((n_cycles,), init_gate))
|
| 65 |
+
|
| 66 |
+
# Buffer para monitorar convergência W₂
|
| 67 |
+
self.register_buffer("w2_history", torch.zeros(n_cycles))
|
| 68 |
+
self.register_buffer("converged_step", torch.tensor(-1))
|
| 69 |
+
# V11.26: tracking de q empírico
|
| 70 |
+
self.register_buffer("q_empirical", torch.ones(n_cycles - 1))
|
| 71 |
+
self.register_buffer("contraction_valid", torch.tensor(False))
|
| 72 |
+
|
| 73 |
+
def compute_w2_squared(self, s1: torch.Tensor, s2: torch.Tensor) -> torch.Tensor:
|
| 74 |
+
"""Computa W₂² entre dois estados (aproximação via Sinkhorn simplificada).
|
| 75 |
+
|
| 76 |
+
V11.26: usa mean over batch + L2 norm (estável numericamente).
|
| 77 |
+
"""
|
| 78 |
+
if s1.dim() == 3:
|
| 79 |
+
s1_flat = s1.mean(dim=1) # (B, D)
|
| 80 |
+
s2_flat = s2.mean(dim=1)
|
| 81 |
+
else:
|
| 82 |
+
s1_flat = s1
|
| 83 |
+
s2_flat = s2
|
| 84 |
+
# W₂² ≈ ||s1 - s2||² (aproximação Gaussiana)
|
| 85 |
+
return ((s1_flat - s2_flat) ** 2).sum(dim=-1).mean()
|
| 86 |
+
|
| 87 |
+
def forward(self, x: torch.Tensor, max_cycles: Optional[int] = None) -> Dict[str, torch.Tensor]:
|
| 88 |
+
"""Executa raciocínio circular com monitoramento de convergência W₂.
|
| 89 |
+
|
| 90 |
+
Args:
|
| 91 |
+
x: (B, T, D) ou (B, D) — representação inicial
|
| 92 |
+
max_cycles: número máximo de ciclos (default: self.n_cycles)
|
| 93 |
+
|
| 94 |
+
Returns:
|
| 95 |
+
dict com:
|
| 96 |
+
'output': representação final s*
|
| 97 |
+
'converged': bool se convergiu antes do máximo
|
| 98 |
+
'w2_final': W₂² final
|
| 99 |
+
'n_cycles_used': ciclos executados
|
| 100 |
+
'q_empirical': ratios de contração por ciclo
|
| 101 |
+
"""
|
| 102 |
+
n_cycles = max_cycles or self.n_cycles
|
| 103 |
+
s = x
|
| 104 |
+
prev_s = x.clone()
|
| 105 |
+
w2_vals: List[float] = []
|
| 106 |
+
q_vals: List[float] = []
|
| 107 |
+
|
| 108 |
+
for c in range(n_cycles):
|
| 109 |
+
# V11.26: clipping do gate para garantir |gate| < contraction_target
|
| 110 |
+
gate_c = torch.clamp(self.residual_gate[c], -self.contraction_target,
|
| 111 |
+
self.contraction_target)
|
| 112 |
+
# Φ(s, x) = LayerNorm(s + gate * RefineHead(s))
|
| 113 |
+
r = self.refine_head(self.refine_norm(s))
|
| 114 |
+
s_new = s + gate_c * r
|
| 115 |
+
|
| 116 |
+
# Monitorar W₂ entre iterações
|
| 117 |
+
w2_sq = self.compute_w2_squared(s_new, prev_s)
|
| 118 |
+
w2_vals.append(w2_sq.item())
|
| 119 |
+
self.w2_history[c] = w2_sq.item()
|
| 120 |
+
|
| 121 |
+
# V11.26: tracking empírico de q_t = W2_t / W2_{t-1}
|
| 122 |
+
if c > 0 and w2_vals[c - 1] > 1e-10:
|
| 123 |
+
q_t = w2_vals[c] / w2_vals[c - 1]
|
| 124 |
+
q_vals.append(q_t)
|
| 125 |
+
if c - 1 < self.q_empirical.numel():
|
| 126 |
+
self.q_empirical[c - 1] = q_t
|
| 127 |
+
|
| 128 |
+
# Verificar convergência (Teorema 4: W₂ < tolerância)
|
| 129 |
+
if w2_sq.item() < self.w2_tolerance and self.converged_step < 0:
|
| 130 |
+
self.converged_step = torch.tensor(c, dtype=torch.long)
|
| 131 |
+
|
| 132 |
+
prev_s = s_new.clone()
|
| 133 |
+
s = s_new
|
| 134 |
+
|
| 135 |
+
# V11.26: verificar contração empírica
|
| 136 |
+
if q_vals:
|
| 137 |
+
avg_q = sum(q_vals) / len(q_vals)
|
| 138 |
+
self.contraction_valid = torch.tensor(avg_q < 1.0)
|
| 139 |
+
# Se q > 1, reduzir gates para forçar contração
|
| 140 |
+
if avg_q >= 1.0:
|
| 141 |
+
with torch.no_grad():
|
| 142 |
+
self.residual_gate.data.mul_(0.95) # decair gates
|
| 143 |
+
|
| 144 |
+
return {
|
| 145 |
+
"output": s,
|
| 146 |
+
"converged": bool(self.converged_step >= 0),
|
| 147 |
+
"w2_final": w2_vals[-1] if w2_vals else 0.0,
|
| 148 |
+
"w2_history": w2_vals,
|
| 149 |
+
"n_cycles_used": n_cycles,
|
| 150 |
+
"converged_step": int(self.converged_step.item()) if self.converged_step >= 0 else -1,
|
| 151 |
+
"q_empirical": q_vals,
|
| 152 |
+
"contraction_valid": bool(self.contraction_valid.item()),
|
| 153 |
+
}
|
| 154 |
+
|
| 155 |
+
def get_contraction_ratio(self) -> float:
|
| 156 |
+
"""Estima o ratio de contração q a partir do histórico W₂."""
|
| 157 |
+
# V11.26: usa q_empirical se disponível
|
| 158 |
+
if self.q_empirical.numel() > 0 and (self.q_empirical > 0).any():
|
| 159 |
+
valid = self.q_empirical[self.q_empirical > 0]
|
| 160 |
+
if valid.numel() > 0:
|
| 161 |
+
return float(valid.mean().item())
|
| 162 |
+
# Fallback: usar w2_history
|
| 163 |
+
if self.w2_history.numel() < 2:
|
| 164 |
+
return 1.0
|
| 165 |
+
vals = self.w2_history[self.w2_history > 0]
|
| 166 |
+
if vals.numel() < 2:
|
| 167 |
+
return 1.0
|
| 168 |
+
ratios = vals[1:] / vals[:-1]
|
| 169 |
+
return float(ratios.mean().item())
|
| 170 |
+
|
| 171 |
+
def get_state(self) -> Dict[str, Any]:
|
| 172 |
+
return {
|
| 173 |
+
"n_cycles": self.n_cycles,
|
| 174 |
+
"contraction_ratio": self.get_contraction_ratio(),
|
| 175 |
+
"w2_history": self.w2_history.tolist(),
|
| 176 |
+
"converged_step": int(self.converged_step.item()),
|
| 177 |
+
"residual_gates": self.residual_gate.detach().tolist(),
|
| 178 |
+
"w2_tolerance": self.w2_tolerance,
|
| 179 |
+
"q_empirical": self.q_empirical.tolist(),
|
| 180 |
+
"contraction_valid": bool(self.contraction_valid.item()),
|
| 181 |
+
}
|
| 182 |
+
|
| 183 |
+
__all__ = ["CircularReasoningWasserstein"]
|
src/bigru_t/tokenizer/bbpe_tokenizer.py
CHANGED
|
@@ -1,76 +1,71 @@
|
|
| 1 |
-
"""bbpe_tokenizer.py — BBPE (Byte-Level BPE) Tokenizer
|
| 2 |
|
| 3 |
═══════════════════════════════════════════════════════════════════════════════
|
| 4 |
-
|
| 5 |
═══════════════════════════════════════════════════════════════════════════════
|
| 6 |
|
| 7 |
-
|
| 8 |
-
|
| 9 |
|
| 10 |
-
|
| 11 |
-
|
| 12 |
-
|
| 13 |
-
|
| 14 |
-
2. COMPATIBILIDADE MULTILINGUE: o mesmo vocabulário serve para Português,
|
| 15 |
-
Xavante (ortografia prática), Inglês, código, etc. sem retreino.
|
| 16 |
|
| 17 |
-
|
| 18 |
-
|
| 19 |
-
|
| 20 |
-
|
| 21 |
-
|
|
|
|
| 22 |
|
| 23 |
-
|
| 24 |
-
|
| 25 |
|
| 26 |
-
|
| 27 |
-
Seja enc: Σ_UTF-8* → V* a função de tokenização BBPE.
|
| 28 |
-
Seja dec: V* → Σ_UTF-8* a função de destokenização.
|
| 29 |
-
Para todo s ∈ Σ_UTF-8*: dec(enc(s)) = s.
|
| 30 |
-
Isso decorre de ByteLevel.add_prefix_space e ByteLevelProcessor serem
|
| 31 |
-
bijeções reversíveis no espaço de bytes.
|
| 32 |
|
| 33 |
═══════════════════════════════════════════════════════════════════════════════
|
| 34 |
-
|
| 35 |
═══════════════════════════════════════════════════════════════════════════════
|
| 36 |
|
| 37 |
-
|
| 38 |
-
|
| 39 |
-
|
| 40 |
-
|
| 41 |
-
-
|
| 42 |
-
|
| 43 |
-
|
| 44 |
-
|
| 45 |
-
|
| 46 |
-
|
| 47 |
-
2. Coleta de N samples (default 50.000) para treino do tokenizer
|
| 48 |
-
3. BpeTrainer com vocab_size=250000, special_tokens=[<s>, <pad>, </s>, <unk>]
|
| 49 |
-
4. ByteLevel pre-tokenizer (add_prefix_space=False)
|
| 50 |
-
5. ByteLevel post-processor (trim_offsets=False)
|
| 51 |
-
6. Salva em tokenizer_bbpe_xavante.json
|
| 52 |
-
|
| 53 |
-
Uso:
|
| 54 |
-
from flexnet.bbpe_tokenizer import BBPETokenizer
|
| 55 |
-
|
| 56 |
-
# Treinar (uma vez)
|
| 57 |
-
tok = BBPETokenizer(vocab_size=250_000)
|
| 58 |
-
tok.train_from_stream(text_iterator, save_path="tokenizer_bbpe_xavante.json")
|
| 59 |
-
|
| 60 |
-
# Carregar
|
| 61 |
-
tok = BBPETokenizer.load("tokenizer_bbpe_xavante.json")
|
| 62 |
-
|
| 63 |
-
# Usar
|
| 64 |
-
ids = tok.encode("Olá mundo em Xavante!") # List[int]
|
| 65 |
-
text = tok.decode(ids) # str
|
| 66 |
═══════════════════════════════════════════════════════════════════════════════
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 67 |
"""
|
| 68 |
from __future__ import annotations
|
|
|
|
| 69 |
import os
|
| 70 |
import json
|
| 71 |
import logging
|
|
|
|
|
|
|
|
|
|
| 72 |
from pathlib import Path
|
| 73 |
-
from typing import
|
|
|
|
|
|
|
| 74 |
|
| 75 |
logger = logging.getLogger(__name__)
|
| 76 |
|
|
@@ -87,26 +82,247 @@ PAD_ID = 1
|
|
| 87 |
EOS_ID = 2
|
| 88 |
UNK_ID = 3
|
| 89 |
|
| 90 |
-
# Vocabulário padrão
|
| 91 |
-
DEFAULT_VOCAB_SIZE =
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 92 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 93 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 94 |
class BBPETokenizer:
|
| 95 |
-
"""
|
| 96 |
|
| 97 |
-
|
| 98 |
-
modelo Xavante. Fornece:
|
| 99 |
- encode(text) -> List[int]
|
| 100 |
- decode(ids) -> str
|
| 101 |
- encode_batch(texts) -> List[List[int]]
|
| 102 |
-
-
|
| 103 |
-
- save(path) / load(path)
|
| 104 |
-
|
| 105 |
-
|
| 106 |
-
-
|
| 107 |
-
-
|
| 108 |
-
-
|
| 109 |
-
- Métodos utilitários para batching e integração com PyTorch
|
| 110 |
"""
|
| 111 |
|
| 112 |
def __init__(
|
|
@@ -117,6 +333,7 @@ class BBPETokenizer:
|
|
| 117 |
eos_token: str = EOS_TOKEN,
|
| 118 |
unk_token: str = UNK_TOKEN,
|
| 119 |
add_prefix_space: bool = False,
|
|
|
|
| 120 |
):
|
| 121 |
self.vocab_size = vocab_size
|
| 122 |
self.bos_token = bos_token
|
|
@@ -124,177 +341,257 @@ class BBPETokenizer:
|
|
| 124 |
self.eos_token = eos_token
|
| 125 |
self.unk_token = unk_token
|
| 126 |
self.add_prefix_space = add_prefix_space
|
|
|
|
| 127 |
|
| 128 |
self._tokenizer = None # Lazy init
|
| 129 |
self._vocab: Optional[Dict[str, int]] = None
|
| 130 |
self._id_to_token: Optional[Dict[int, str]] = None
|
|
|
|
|
|
|
| 131 |
|
| 132 |
# ------------------------------------------------------------------
|
| 133 |
-
#
|
| 134 |
# ------------------------------------------------------------------
|
| 135 |
-
def
|
| 136 |
-
"""Constrói o tokenizer BBPE com configuração V5."""
|
| 137 |
-
from tokenizers import Tokenizer
|
| 138 |
-
from tokenizers.models import BPE
|
| 139 |
-
from tokenizers.pre_tokenizers import ByteLevel
|
| 140 |
-
from tokenizers.processors import ByteLevel as ByteLevelProcessor
|
| 141 |
-
from tokenizers.decoders import ByteLevel as ByteLevelDecoder
|
| 142 |
-
|
| 143 |
-
tok = Tokenizer(BPE(unk_token=self.unk_token))
|
| 144 |
-
tok.pre_tokenizer = ByteLevel(add_prefix_space=self.add_prefix_space)
|
| 145 |
-
tok.post_processor = ByteLevelProcessor(trim_offsets=False)
|
| 146 |
-
# CRÍTICO: ByteLevel decoder inverte o mapeamento byte-level de volta para UTF-8.
|
| 147 |
-
# Sem isto, decode() retorna a representação interna "Ġ-style" em vez do texto original.
|
| 148 |
-
tok.decoder = ByteLevelDecoder()
|
| 149 |
-
return tok
|
| 150 |
-
|
| 151 |
-
def train_from_stream(
|
| 152 |
self,
|
| 153 |
text_iterator: Iterator[str],
|
| 154 |
-
save_path: Optional[
|
| 155 |
min_frequency: int = 2,
|
|
|
|
| 156 |
show_progress: bool = True,
|
| 157 |
chunk_size: int = 500,
|
| 158 |
) -> None:
|
| 159 |
-
"""Treina o tokenizer BBPE
|
| 160 |
-
|
| 161 |
-
|
| 162 |
-
|
| 163 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 164 |
|
| 165 |
Args:
|
| 166 |
text_iterator: iterador yielding strings de texto
|
| 167 |
save_path: caminho para salvar o tokenizer JSON
|
| 168 |
min_frequency: frequência mínima de um par para ser mergeado
|
| 169 |
-
|
| 170 |
-
|
| 171 |
-
|
| 172 |
"""
|
| 173 |
-
from tokenizers.trainers import BpeTrainer
|
| 174 |
-
|
| 175 |
logger.info(
|
| 176 |
-
"
|
| 177 |
-
self.vocab_size, min_frequency,
|
| 178 |
)
|
| 179 |
|
| 180 |
-
#
|
| 181 |
-
|
|
|
|
| 182 |
|
| 183 |
-
#
|
| 184 |
-
|
| 185 |
-
|
| 186 |
-
|
| 187 |
-
|
| 188 |
-
|
| 189 |
-
|
|
|
|
| 190 |
)
|
| 191 |
|
| 192 |
-
#
|
| 193 |
-
|
| 194 |
-
|
| 195 |
-
|
| 196 |
-
|
| 197 |
-
|
| 198 |
-
|
| 199 |
-
|
| 200 |
-
|
| 201 |
-
|
| 202 |
-
|
| 203 |
-
|
| 204 |
-
|
| 205 |
-
|
| 206 |
-
|
| 207 |
-
|
| 208 |
-
|
| 209 |
-
|
| 210 |
-
|
| 211 |
-
|
| 212 |
-
|
| 213 |
-
|
| 214 |
-
|
| 215 |
-
|
| 216 |
-
|
| 217 |
-
|
| 218 |
-
|
| 219 |
-
|
| 220 |
-
|
| 221 |
-
|
| 222 |
-
|
| 223 |
-
|
| 224 |
-
|
| 225 |
-
|
| 226 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 227 |
)
|
|
|
|
| 228 |
|
| 229 |
-
|
| 230 |
-
|
| 231 |
-
|
|
|
|
| 232 |
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 233 |
|
| 234 |
-
#
|
| 235 |
-
|
| 236 |
-
|
| 237 |
-
trainer,
|
| 238 |
-
)
|
| 239 |
|
| 240 |
-
|
| 241 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 242 |
|
| 243 |
-
|
| 244 |
-
|
| 245 |
-
|
| 246 |
-
|
| 247 |
|
| 248 |
-
|
| 249 |
-
|
| 250 |
-
self.save(save_path)
|
| 251 |
-
|
| 252 |
-
finally:
|
| 253 |
-
# Limpeza
|
| 254 |
-
for f in tmp_files:
|
| 255 |
-
try:
|
| 256 |
-
f.unlink()
|
| 257 |
-
except Exception:
|
| 258 |
-
pass
|
| 259 |
-
try:
|
| 260 |
-
tmp_dir.rmdir()
|
| 261 |
-
except Exception:
|
| 262 |
-
pass
|
| 263 |
|
| 264 |
-
|
|
|
|
|
|
|
|
|
|
| 265 |
self,
|
| 266 |
-
|
| 267 |
save_path: Optional[Union[str, Path]] = None,
|
| 268 |
min_frequency: int = 2,
|
| 269 |
show_progress: bool = True,
|
|
|
|
|
|
|
| 270 |
) -> None:
|
| 271 |
-
"""Treina o tokenizer BBPE a partir de
|
| 272 |
-
from tokenizers.trainers import BpeTrainer
|
| 273 |
|
| 274 |
-
|
| 275 |
-
|
| 276 |
-
len(file_paths), self.vocab_size,
|
| 277 |
-
)
|
| 278 |
|
| 279 |
-
|
| 280 |
-
|
| 281 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 282 |
min_frequency=min_frequency,
|
| 283 |
-
|
| 284 |
show_progress=show_progress,
|
| 285 |
-
|
| 286 |
)
|
| 287 |
-
self._tokenizer.train([str(p) for p in file_paths], trainer)
|
| 288 |
-
self._build_vocab_cache()
|
| 289 |
|
| 290 |
-
|
| 291 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 292 |
|
| 293 |
# ------------------------------------------------------------------
|
| 294 |
# Save / Load
|
| 295 |
# ------------------------------------------------------------------
|
| 296 |
def save(self, path: Union[str, Path]) -> None:
|
| 297 |
-
"""Salva o tokenizer em arquivo JSON."""
|
| 298 |
if self._tokenizer is None:
|
| 299 |
raise RuntimeError("Tokenizer não treinado. Chame train_*() primeiro.")
|
| 300 |
path = Path(path)
|
|
@@ -314,7 +611,6 @@ class BBPETokenizer:
|
|
| 314 |
instance = cls() # default vocab_size
|
| 315 |
instance._tokenizer = Tokenizer.from_file(str(path))
|
| 316 |
instance._build_vocab_cache()
|
| 317 |
-
# Atualiza vocab_size com tamanho real
|
| 318 |
instance.vocab_size = len(instance._vocab)
|
| 319 |
logger.info(
|
| 320 |
"BBPE tokenizer loaded: %s (vocab_size=%d)",
|
|
@@ -359,7 +655,6 @@ class BBPETokenizer:
|
|
| 359 |
f"Input too long: {len(ids)} > {max_length} and truncation=False"
|
| 360 |
)
|
| 361 |
if add_special_tokens:
|
| 362 |
-
# Preserva EOS no final
|
| 363 |
ids = ids[:max_length - 1] + [EOS_ID] if max_length >= 1 else [EOS_ID]
|
| 364 |
else:
|
| 365 |
ids = ids[:max_length]
|
|
@@ -371,9 +666,41 @@ class BBPETokenizer:
|
|
| 371 |
add_special_tokens: bool = False,
|
| 372 |
max_length: Optional[int] = None,
|
| 373 |
) -> List[List[int]]:
|
| 374 |
-
"""Codifica um batch de textos."""
|
| 375 |
return [self.encode(t, add_special_tokens, max_length) for t in texts]
|
| 376 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 377 |
def decode(
|
| 378 |
self,
|
| 379 |
ids: List[int],
|
|
@@ -416,7 +743,6 @@ class BBPETokenizer:
|
|
| 416 |
]
|
| 417 |
|
| 418 |
if pad_to_max_length:
|
| 419 |
-
# Pad todas para max_length
|
| 420 |
padded = []
|
| 421 |
masks = []
|
| 422 |
for ids in batch_ids:
|
|
@@ -430,7 +756,6 @@ class BBPETokenizer:
|
|
| 430 |
input_ids = torch.tensor(padded, dtype=torch.long)
|
| 431 |
attention_mask = torch.tensor(masks, dtype=torch.long)
|
| 432 |
else:
|
| 433 |
-
# Sem padding — cada amostra tem seu tamanho
|
| 434 |
input_ids = [torch.tensor(ids, dtype=torch.long) for ids in batch_ids]
|
| 435 |
attention_mask = [torch.ones(len(ids), dtype=torch.long) for ids in batch_ids]
|
| 436 |
|
|
@@ -456,6 +781,11 @@ class BBPETokenizer:
|
|
| 456 |
"""Tamanho real do vocabulário carregado/treinado."""
|
| 457 |
return len(self.vocab)
|
| 458 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 459 |
def __len__(self) -> int:
|
| 460 |
return self.actual_vocab_size
|
| 461 |
|
|
@@ -487,15 +817,12 @@ class BBPETokenizer:
|
|
| 487 |
ids = self.encode(text, add_special_tokens=False)
|
| 488 |
decoded = self.decode(ids, skip_special_tokens=True)
|
| 489 |
|
| 490 |
-
# Byte-level: o decoded pode ter um espaço prefix se add_prefix_space=True
|
| 491 |
-
# Comparamos sem o prefix
|
| 492 |
expected = text
|
| 493 |
got = decoded
|
| 494 |
|
| 495 |
if expected == got:
|
| 496 |
successes += 1
|
| 497 |
else:
|
| 498 |
-
# Tenta sem o prefix space
|
| 499 |
if self.add_prefix_space and got.startswith(" "):
|
| 500 |
got = got[1:]
|
| 501 |
if expected == got:
|
|
@@ -507,7 +834,6 @@ class BBPETokenizer:
|
|
| 507 |
"ids_count": len(ids),
|
| 508 |
})
|
| 509 |
|
| 510 |
-
# Compression: bytes / tokens
|
| 511 |
n_bytes = len(text.encode("utf-8"))
|
| 512 |
n_tokens = len(ids)
|
| 513 |
if n_tokens > 0:
|
|
@@ -527,7 +853,7 @@ class BBPETokenizer:
|
|
| 527 |
|
| 528 |
|
| 529 |
# ---------------------------------------------------------------------------
|
| 530 |
-
# Byte-level alphabet (for BBPE initial alphabet)
|
| 531 |
# ---------------------------------------------------------------------------
|
| 532 |
class ByteLevel:
|
| 533 |
"""Wrapper para o alfabeto byte-level (256 bytes)."""
|
|
@@ -535,8 +861,7 @@ class ByteLevel:
|
|
| 535 |
@staticmethod
|
| 536 |
def alphabet() -> List[str]:
|
| 537 |
"""Retorna os 256 caracteres byte-level (Ġ-style do GPT-2)."""
|
| 538 |
-
|
| 539 |
-
return list(_ByteLevel.alphabet())
|
| 540 |
|
| 541 |
|
| 542 |
# ---------------------------------------------------------------------------
|
|
@@ -547,14 +872,16 @@ def create_or_load_tokenizer(
|
|
| 547 |
text_iterator: Optional[Iterator[str]] = None,
|
| 548 |
vocab_size: int = DEFAULT_VOCAB_SIZE,
|
| 549 |
min_frequency: int = 2,
|
|
|
|
| 550 |
) -> BBPETokenizer:
|
| 551 |
-
"""Carrega um tokenizer existente ou treina um novo
|
| 552 |
|
| 553 |
Args:
|
| 554 |
path: caminho do arquivo JSON
|
| 555 |
text_iterator: iterador de textos para treinar (se arquivo não existe)
|
| 556 |
-
vocab_size: tamanho do vocabulário (default
|
| 557 |
min_frequency: frequência mínima para merges
|
|
|
|
| 558 |
|
| 559 |
Returns:
|
| 560 |
BBPETokenizer carregado/treinado
|
|
@@ -569,12 +896,13 @@ def create_or_load_tokenizer(
|
|
| 569 |
f"Tokenizer file {path} does not exist and no text_iterator provided"
|
| 570 |
)
|
| 571 |
|
| 572 |
-
logger.info("Training new BBPE tokenizer: %s", path)
|
| 573 |
-
tok = BBPETokenizer(vocab_size=vocab_size)
|
| 574 |
-
tok.
|
| 575 |
text_iterator,
|
| 576 |
save_path=path,
|
| 577 |
min_frequency=min_frequency,
|
|
|
|
| 578 |
)
|
| 579 |
return tok
|
| 580 |
|
|
@@ -593,4 +921,13 @@ __all__ = [
|
|
| 593 |
"UNK_ID",
|
| 594 |
"SPECIAL_TOKENS",
|
| 595 |
"DEFAULT_VOCAB_SIZE",
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 596 |
]
|
|
|
|
| 1 |
+
"""bbpe_tokenizer.py — BBPE (Byte-Level BPE) Tokenizer PARALELO (Map-Reduce).
|
| 2 |
|
| 3 |
═══════════════════════════════════════════════════════════════════════════════
|
| 4 |
+
REFATORAÇÃO: Algoritmo Paralelo Map-Reduce (substitui o treinamento sequencial)
|
| 5 |
═══════════════════════════════════════════════════════════════════════════════
|
| 6 |
|
| 7 |
+
O treinamento do BBPE foi refatorado para um modelo Map-Reduce paralelo,
|
| 8 |
+
substituindo o BpeTrainer sequencial da HuggingFace. O algoritmo:
|
| 9 |
|
| 10 |
+
ETAPA 0: Inicialização
|
| 11 |
+
- distribute_texts(): converte o iterador em shards balanceados
|
| 12 |
+
- pre_tokenize_shard(): pré-tokeniza byte-level cada shard
|
| 13 |
+
- Estruturas globais: vocab, token_to_id, merges
|
|
|
|
|
|
|
| 14 |
|
| 15 |
+
LAÇO PRINCIPAL DE MERGES:
|
| 16 |
+
FASE 1 — MAP: count_pairs_in_shard() em paralelo (ProcessPoolExecutor)
|
| 17 |
+
FASE 2 — REDUCE: agrega contagens locais em global_counts
|
| 18 |
+
FASE 3 — CHOICE: escolhe o par de maior frequência (desempate lexicográfico)
|
| 19 |
+
FASE 4 — APPLY: apply_merge_in_shard() em paralelo
|
| 20 |
+
FASE 5 — UPDATE: atualiza vocab, token_to_id, merges
|
| 21 |
|
| 22 |
+
ETAPA FINAL: build_bpe_from_merges() constrói o tokenizer HF a partir
|
| 23 |
+
dos merges + vocab calculados.
|
| 24 |
|
| 25 |
+
INFERÊNCIA PARALELA: encode_batch_parallel() usa ThreadPoolExecutor.
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 26 |
|
| 27 |
═══════════════════════════════════════════════════════════════════════════════
|
| 28 |
+
TEOREMA 20 (BBPE Universal Coverage) — mantido
|
| 29 |
═══════════════════════════════════════════════════════════════════════════════
|
| 30 |
|
| 31 |
+
BBPE opera no espaço de BYTES UTF-8 (256 símbolos base). Garante:
|
| 32 |
+
1. COBERTURA UNIVERSAL: qualquer string UTF-8 é tokenizável sem <unk>.
|
| 33 |
+
2. COMPATIBILIDADE MULTILINGUE: mesmo vocab para PT-BR, Xavante, EN, código.
|
| 34 |
+
3. COMPRESSÃO ÓTIMA: BPE greedy + byte-level = merges no espaço total.
|
| 35 |
+
4. ESCALABILIDADE: V = 16K-250K tokens cobre eficientemente múltiplos idiomas.
|
| 36 |
+
|
| 37 |
+
INTEGRIDADE (Semântica Preservada):
|
| 38 |
+
dec(enc(s)) = s para todo s ∈ Σ_UTF-8*.
|
| 39 |
+
Decorre de ByteLevel ser bijeção reversível no espaço de bytes.
|
| 40 |
+
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 41 |
═══════════════════════════════════════════════════════════════════════════════
|
| 42 |
+
COMPATIBILIDADE
|
| 43 |
+
═══════════════════════════════════════════════════════════════════════════════
|
| 44 |
+
|
| 45 |
+
Mantém a API pública da versão anterior:
|
| 46 |
+
- encode(text) -> List[int]
|
| 47 |
+
- decode(ids) -> str
|
| 48 |
+
- encode_batch(texts) -> List[List[int]]
|
| 49 |
+
- encode_batch_parallel(texts) -> List[List[int]] [NOVO]
|
| 50 |
+
- save(path) / load(path)
|
| 51 |
+
- train_from_stream(iter, path) [delega para train_parallel_from_stream]
|
| 52 |
+
- train_parallel_from_stream(iter, ...) [NOVO — algoritmo Map-Reduce]
|
| 53 |
+
- train_from_files(files, path)
|
| 54 |
+
- encode_tensor(texts, max_length)
|
| 55 |
+
- validate_roundtrip(test_texts)
|
| 56 |
"""
|
| 57 |
from __future__ import annotations
|
| 58 |
+
|
| 59 |
import os
|
| 60 |
import json
|
| 61 |
import logging
|
| 62 |
+
import gc
|
| 63 |
+
from collections import defaultdict
|
| 64 |
+
from concurrent.futures import ProcessPoolExecutor, ThreadPoolExecutor, as_completed
|
| 65 |
from pathlib import Path
|
| 66 |
+
from typing import (
|
| 67 |
+
Iterator, List, Optional, Dict, Any, Union, Tuple, Iterable,
|
| 68 |
+
)
|
| 69 |
|
| 70 |
logger = logging.getLogger(__name__)
|
| 71 |
|
|
|
|
| 82 |
EOS_ID = 2
|
| 83 |
UNK_ID = 3
|
| 84 |
|
| 85 |
+
# Vocabulário padrão (reduzido de 250K para 16K — adequado ao escopo BiGRU_T)
|
| 86 |
+
DEFAULT_VOCAB_SIZE = 16_384
|
| 87 |
+
|
| 88 |
+
|
| 89 |
+
# ---------------------------------------------------------------------------
|
| 90 |
+
# Mapeamento Byte-Level (GPT-2 style: 0-255 -> Ġ-style unicode strings)
|
| 91 |
+
# ---------------------------------------------------------------------------
|
| 92 |
+
def bytes_to_unicode() -> Dict[int, str]:
|
| 93 |
+
"""Retorna mapeamento byte (0-255) -> símbolo unicode (Ġ-style do GPT-2).
|
| 94 |
+
|
| 95 |
+
Bytes correspondentes a caracteres printable (33-126, 161-172, 174-255)
|
| 96 |
+
mapeiam para si mesmos. Os demais (0-32, 127-160, 173) mapeiam para
|
| 97 |
+
codepoints a partir de 256 (Ġ=256+32=288 → 'Ġ', etc.).
|
| 98 |
+
"""
|
| 99 |
+
bs = (
|
| 100 |
+
list(range(ord("!"), ord("~") + 1))
|
| 101 |
+
+ list(range(ord("¡"), ord("¬") + 1))
|
| 102 |
+
+ list(range(ord("®"), ord("ÿ") + 1))
|
| 103 |
+
)
|
| 104 |
+
cs = bs[:]
|
| 105 |
+
n = 0
|
| 106 |
+
for b in range(256):
|
| 107 |
+
if b not in bs:
|
| 108 |
+
bs.append(b)
|
| 109 |
+
cs.append(256 + n)
|
| 110 |
+
n += 1
|
| 111 |
+
cs = [chr(c) for c in cs]
|
| 112 |
+
return dict(zip(bs, cs))
|
| 113 |
+
|
| 114 |
+
|
| 115 |
+
# Tabelas globais (construídas uma vez no import)
|
| 116 |
+
BYTE_TO_SYMBOL: Dict[int, str] = bytes_to_unicode()
|
| 117 |
+
SYMBOL_TO_BYTE: Dict[str, int] = {v: k for k, v in BYTE_TO_SYMBOL.items()}
|
| 118 |
+
|
| 119 |
+
# Alfabeto byte-level (256 símbolos)
|
| 120 |
+
ALPHABET: List[str] = [BYTE_TO_SYMBOL[i] for i in range(256)]
|
| 121 |
+
|
| 122 |
+
|
| 123 |
+
# ---------------------------------------------------------------------------
|
| 124 |
+
# FUNÇÕES AUXILIARES DO ALGORITMO PARALELO (top-level para picklability)
|
| 125 |
+
# ---------------------------------------------------------------------------
|
| 126 |
+
def distribute_texts(
|
| 127 |
+
text_iterator: Iterable[str],
|
| 128 |
+
num_workers: int,
|
| 129 |
+
chunk_size: int = 500,
|
| 130 |
+
) -> List[List[str]]:
|
| 131 |
+
"""Converte um iterador de textos em partições balanceadas (shards).
|
| 132 |
+
|
| 133 |
+
Cada shard é uma lista de strings que será processada por um worker.
|
| 134 |
+
A distribuição é round-robin sobre chunks para balancear carga.
|
| 135 |
+
|
| 136 |
+
Args:
|
| 137 |
+
text_iterator: iterador yielding strings de texto
|
| 138 |
+
num_workers: número de shards a produzir
|
| 139 |
+
chunk_size: textos acumulados antes de formar um chunk
|
| 140 |
+
|
| 141 |
+
Returns:
|
| 142 |
+
Lista de `num_workers` shards (cada shard é List[str]).
|
| 143 |
+
"""
|
| 144 |
+
if num_workers < 1:
|
| 145 |
+
num_workers = 1
|
| 146 |
+
shards: List[List[str]] = [[] for _ in range(num_workers)]
|
| 147 |
+
current_chunk: List[str] = []
|
| 148 |
+
chunk_idx = 0
|
| 149 |
+
total = 0
|
| 150 |
+
|
| 151 |
+
for text in text_iterator:
|
| 152 |
+
if not text or len(str(text).strip()) < 10:
|
| 153 |
+
continue
|
| 154 |
+
current_chunk.append(str(text))
|
| 155 |
+
total += 1
|
| 156 |
+
if len(current_chunk) >= chunk_size:
|
| 157 |
+
# Round-robin assignment
|
| 158 |
+
shards[chunk_idx % num_workers].extend(current_chunk)
|
| 159 |
+
current_chunk = []
|
| 160 |
+
chunk_idx += 1
|
| 161 |
+
# Flush final
|
| 162 |
+
if current_chunk:
|
| 163 |
+
shards[chunk_idx % num_workers].extend(current_chunk)
|
| 164 |
+
|
| 165 |
+
# Remove shards vazios (pode acontecer se num_workers > chunks)
|
| 166 |
+
shards = [s for s in shards if s]
|
| 167 |
+
if not shards:
|
| 168 |
+
raise RuntimeError(
|
| 169 |
+
"distribute_texts: nenhum texto válido encontrado no iterador"
|
| 170 |
+
)
|
| 171 |
+
logger.info(
|
| 172 |
+
"distribute_texts: %d textos em %d shards (chunk_size=%d)",
|
| 173 |
+
total, len(shards), chunk_size,
|
| 174 |
+
)
|
| 175 |
+
return shards
|
| 176 |
+
|
| 177 |
+
|
| 178 |
+
def pre_tokenize_shard(texts: List[str]) -> List[List[str]]:
|
| 179 |
+
"""Converte uma lista de textos em uma lista de listas de símbolos byte-level.
|
| 180 |
+
|
| 181 |
+
Cada símbolo é uma string representando um byte (Ġ-style).
|
| 182 |
+
Tokens especiais (se presentes no texto como substrings) são tratados
|
| 183 |
+
como símbolos únicos — mas nesta implementação simples, expandimos tudo
|
| 184 |
+
para bytes (a detecção de especiais é feita no encode, não no treino).
|
| 185 |
|
| 186 |
+
Args:
|
| 187 |
+
texts: lista de strings (um shard)
|
| 188 |
+
|
| 189 |
+
Returns:
|
| 190 |
+
Lista de listas de símbolos (uma lista por documento).
|
| 191 |
+
"""
|
| 192 |
+
shard_syms: List[List[str]] = []
|
| 193 |
+
for text in texts:
|
| 194 |
+
doc_syms: List[str] = []
|
| 195 |
+
# Pré-tokenização ByteLevel: cada caractere -> UTF-8 bytes -> símbolos
|
| 196 |
+
for ch in text:
|
| 197 |
+
utf8_bytes = ch.encode("utf-8")
|
| 198 |
+
for b in utf8_bytes:
|
| 199 |
+
doc_syms.append(BYTE_TO_SYMBOL[b])
|
| 200 |
+
shard_syms.append(doc_syms)
|
| 201 |
+
return shard_syms
|
| 202 |
+
|
| 203 |
+
|
| 204 |
+
def count_pairs_in_shard(
|
| 205 |
+
shard: List[List[str]],
|
| 206 |
+
min_freq: int = 2,
|
| 207 |
+
) -> Dict[Tuple[str, str], int]:
|
| 208 |
+
"""Conta pares adjacentes dentro de cada documento, sem cruzar fronteiras.
|
| 209 |
+
|
| 210 |
+
Args:
|
| 211 |
+
shard: lista de documentos (cada doc é lista de símbolos)
|
| 212 |
+
min_freq: frequência mínima para manter o par (poda local)
|
| 213 |
+
|
| 214 |
+
Returns:
|
| 215 |
+
Dicionário {(esq, dir): contagem_local}.
|
| 216 |
+
"""
|
| 217 |
+
counts: Dict[Tuple[str, str], int] = defaultdict(int)
|
| 218 |
+
for doc in shard:
|
| 219 |
+
if len(doc) < 2:
|
| 220 |
+
continue
|
| 221 |
+
for i in range(len(doc) - 1):
|
| 222 |
+
pair = (doc[i], doc[i + 1])
|
| 223 |
+
counts[pair] += 1
|
| 224 |
+
# Poda local: descarta pares com contagem < min_freq
|
| 225 |
+
if min_freq > 1:
|
| 226 |
+
return {p: c for p, c in counts.items() if c >= min_freq}
|
| 227 |
+
return dict(counts)
|
| 228 |
+
|
| 229 |
+
|
| 230 |
+
def apply_merge_in_shard(
|
| 231 |
+
shard: List[List[str]],
|
| 232 |
+
left: str,
|
| 233 |
+
right: str,
|
| 234 |
+
replacement: str,
|
| 235 |
+
) -> List[List[str]]:
|
| 236 |
+
"""Substitui toda ocorrência adjacente de (left, right) por `replacement`.
|
| 237 |
|
| 238 |
+
Args:
|
| 239 |
+
shard: lista de documentos (cada doc é lista de símbolos)
|
| 240 |
+
left: símbolo esquerdo do merge
|
| 241 |
+
right: símbolo direito do merge
|
| 242 |
+
replacement: novo símbolo que substitui o par
|
| 243 |
+
|
| 244 |
+
Returns:
|
| 245 |
+
Novo shard com o merge aplicado em todos os documentos.
|
| 246 |
+
"""
|
| 247 |
+
new_shard: List[List[str]] = []
|
| 248 |
+
for doc in shard:
|
| 249 |
+
if len(doc) < 2:
|
| 250 |
+
new_shard.append(doc)
|
| 251 |
+
continue
|
| 252 |
+
new_doc: List[str] = []
|
| 253 |
+
i = 0
|
| 254 |
+
n = len(doc)
|
| 255 |
+
while i < n:
|
| 256 |
+
if i < n - 1 and doc[i] == left and doc[i + 1] == right:
|
| 257 |
+
new_doc.append(replacement)
|
| 258 |
+
i += 2
|
| 259 |
+
else:
|
| 260 |
+
new_doc.append(doc[i])
|
| 261 |
+
i += 1
|
| 262 |
+
new_shard.append(new_doc)
|
| 263 |
+
return new_shard
|
| 264 |
+
|
| 265 |
+
|
| 266 |
+
def build_bpe_from_merges(
|
| 267 |
+
merges: List[Tuple[str, str, str]],
|
| 268 |
+
token_to_id: Dict[str, int],
|
| 269 |
+
unk_token: str = UNK_TOKEN,
|
| 270 |
+
add_prefix_space: bool = False,
|
| 271 |
+
):
|
| 272 |
+
"""Constrói um tokenizers.Tokenizer a partir dos merges e vocab calculados.
|
| 273 |
+
|
| 274 |
+
Converte o formato interno (lista de tuplas (left, right, new)) para o
|
| 275 |
+
formato esperado pelo tokenizers.models.BPE (lista de strings "left right").
|
| 276 |
+
|
| 277 |
+
Args:
|
| 278 |
+
merges: lista de (esq, dir, novo_token)
|
| 279 |
+
token_to_id: mapeamento token -> id
|
| 280 |
+
unk_token: token de desconhecido
|
| 281 |
+
add_prefix_space: se True, adiciona espaço prefixo no pre-tokenizer
|
| 282 |
+
|
| 283 |
+
Returns:
|
| 284 |
+
tokenizers.Tokenizer configurado com BPE + ByteLevel.
|
| 285 |
+
"""
|
| 286 |
+
from tokenizers import Tokenizer
|
| 287 |
+
from tokenizers.models import BPE
|
| 288 |
+
from tokenizers.pre_tokenizers import ByteLevel
|
| 289 |
+
from tokenizers.processors import ByteLevel as ByteLevelProcessor
|
| 290 |
+
from tokenizers.decoders import ByteLevel as ByteLevelDecoder
|
| 291 |
+
|
| 292 |
+
# Converte merges para formato HF: lista de tuplas (left, right)
|
| 293 |
+
hf_merges = [(left, right) for (left, right, _new) in merges]
|
| 294 |
+
|
| 295 |
+
# BPE model com vocab e merges
|
| 296 |
+
bpe = BPE(
|
| 297 |
+
vocab=token_to_id,
|
| 298 |
+
merges=hf_merges,
|
| 299 |
+
unk_token=unk_token,
|
| 300 |
+
)
|
| 301 |
+
tok = Tokenizer(bpe)
|
| 302 |
+
tok.pre_tokenizer = ByteLevel(add_prefix_space=add_prefix_space)
|
| 303 |
+
tok.post_processor = ByteLevelProcessor(trim_offsets=False)
|
| 304 |
+
# CRÍTICO: ByteLevel decoder inverte o mapeamento byte-level de volta para UTF-8.
|
| 305 |
+
tok.decoder = ByteLevelDecoder()
|
| 306 |
+
return tok
|
| 307 |
+
|
| 308 |
+
|
| 309 |
+
# ---------------------------------------------------------------------------
|
| 310 |
+
# Classe principal
|
| 311 |
+
# ---------------------------------------------------------------------------
|
| 312 |
class BBPETokenizer:
|
| 313 |
+
"""BBPE Tokenizer com treinamento paralelo Map-Reduce.
|
| 314 |
|
| 315 |
+
API pública (compatível com versão anterior):
|
|
|
|
| 316 |
- encode(text) -> List[int]
|
| 317 |
- decode(ids) -> str
|
| 318 |
- encode_batch(texts) -> List[List[int]]
|
| 319 |
+
- encode_batch_parallel(texts) -> List[List[int]] [NOVO]
|
| 320 |
+
- save(path) / load(path)
|
| 321 |
+
- train_parallel_from_stream(iter, ...) [NOVO — algoritmo Map-Reduce]
|
| 322 |
+
- train_from_stream(iter, ...) [delega para paralelo]
|
| 323 |
+
- train_from_files(files, ...)
|
| 324 |
+
- encode_tensor(texts, max_length)
|
| 325 |
+
- validate_roundtrip(test_texts)
|
|
|
|
| 326 |
"""
|
| 327 |
|
| 328 |
def __init__(
|
|
|
|
| 333 |
eos_token: str = EOS_TOKEN,
|
| 334 |
unk_token: str = UNK_TOKEN,
|
| 335 |
add_prefix_space: bool = False,
|
| 336 |
+
num_workers: int = 4,
|
| 337 |
):
|
| 338 |
self.vocab_size = vocab_size
|
| 339 |
self.bos_token = bos_token
|
|
|
|
| 341 |
self.eos_token = eos_token
|
| 342 |
self.unk_token = unk_token
|
| 343 |
self.add_prefix_space = add_prefix_space
|
| 344 |
+
self.num_workers = max(1, num_workers)
|
| 345 |
|
| 346 |
self._tokenizer = None # Lazy init
|
| 347 |
self._vocab: Optional[Dict[str, int]] = None
|
| 348 |
self._id_to_token: Optional[Dict[int, str]] = None
|
| 349 |
+
# Merges aprendidos (para inspeção / re-build)
|
| 350 |
+
self._merges: List[Tuple[str, str, str]] = []
|
| 351 |
|
| 352 |
# ------------------------------------------------------------------
|
| 353 |
+
# TREINAMENTO PARALELO (Map-Reduce) — NOVO
|
| 354 |
# ------------------------------------------------------------------
|
| 355 |
+
def train_parallel_from_stream(
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 356 |
self,
|
| 357 |
text_iterator: Iterator[str],
|
| 358 |
+
save_path: Optional[Path] = None,
|
| 359 |
min_frequency: int = 2,
|
| 360 |
+
num_workers: int = 4,
|
| 361 |
show_progress: bool = True,
|
| 362 |
chunk_size: int = 500,
|
| 363 |
) -> None:
|
| 364 |
+
"""Treina o tokenizer BBPE em paralelo (Map-Reduce com ProcessPoolExecutor).
|
| 365 |
+
|
| 366 |
+
Algoritmo:
|
| 367 |
+
ETAPA 0: distribute_texts + pre_tokenize_shard (inicialização)
|
| 368 |
+
LAÇO:
|
| 369 |
+
FASE 1: MAP — count_pairs_in_shard em paralelo
|
| 370 |
+
FASE 2: REDUCE — agrega contagens
|
| 371 |
+
FASE 3: CHOICE — melhor par (freq máx, desempate lexicográfico)
|
| 372 |
+
FASE 4: APPLY — apply_merge_in_shard em paralelo
|
| 373 |
+
FASE 5: UPDATE — atualiza vocab + merges
|
| 374 |
+
ETAPA FINAL: build_bpe_from_merges
|
| 375 |
|
| 376 |
Args:
|
| 377 |
text_iterator: iterador yielding strings de texto
|
| 378 |
save_path: caminho para salvar o tokenizer JSON
|
| 379 |
min_frequency: frequência mínima de um par para ser mergeado
|
| 380 |
+
num_workers: número de processos paralelos
|
| 381 |
+
show_progress: exibir progresso por iteração
|
| 382 |
+
chunk_size: textos por chunk na distribuição
|
| 383 |
"""
|
|
|
|
|
|
|
| 384 |
logger.info(
|
| 385 |
+
"BBPE PARALLEL train: vocab_size=%d, min_freq=%d, workers=%d",
|
| 386 |
+
self.vocab_size, min_frequency, num_workers,
|
| 387 |
)
|
| 388 |
|
| 389 |
+
# --- ETAPA 0: Inicialização ---
|
| 390 |
+
# 0.1 Distribui textos em shards balanceados
|
| 391 |
+
shards = distribute_texts(text_iterator, num_workers, chunk_size)
|
| 392 |
|
| 393 |
+
# 0.2 Pré-tokeniza cada shard em sequências de símbolos byte-level
|
| 394 |
+
# (paralelo, pois pre_tokenize é CPU-bound)
|
| 395 |
+
with ProcessPoolExecutor(max_workers=num_workers) as executor:
|
| 396 |
+
futures = [executor.submit(pre_tokenize_shard, shard) for shard in shards]
|
| 397 |
+
shard_symbols = [f.result() for f in futures]
|
| 398 |
+
logger.info(
|
| 399 |
+
"BBPE PARALLEL: pré-tokenização concluída (%d shards)",
|
| 400 |
+
len(shard_symbols),
|
| 401 |
)
|
| 402 |
|
| 403 |
+
# 0.3 Estruturas globais
|
| 404 |
+
current_vocab: set = set(ALPHABET + SPECIAL_TOKENS)
|
| 405 |
+
token_to_id: Dict[str, int] = {}
|
| 406 |
+
# IDs canônicos para especiais primeiro
|
| 407 |
+
for i, tok in enumerate(SPECIAL_TOKENS):
|
| 408 |
+
token_to_id[tok] = i
|
| 409 |
+
# Depois o alfabeto byte-level
|
| 410 |
+
next_id = len(SPECIAL_TOKENS)
|
| 411 |
+
for sym in ALPHABET:
|
| 412 |
+
if sym not in token_to_id:
|
| 413 |
+
token_to_id[sym] = next_id
|
| 414 |
+
next_id += 1
|
| 415 |
+
merges: List[Tuple[str, str, str]] = []
|
| 416 |
+
|
| 417 |
+
# --- LAÇO PRINCIPAL DE MERGES ---
|
| 418 |
+
iteration = 0
|
| 419 |
+
target_merges = self.vocab_size - len(token_to_id)
|
| 420 |
+
|
| 421 |
+
while len(current_vocab) < self.vocab_size:
|
| 422 |
+
iteration += 1
|
| 423 |
+
|
| 424 |
+
# --- FASE 1: MAP (contagem local de pares) ---
|
| 425 |
+
with ProcessPoolExecutor(max_workers=num_workers) as executor:
|
| 426 |
+
futures = [
|
| 427 |
+
executor.submit(count_pairs_in_shard, sym_shard, min_frequency)
|
| 428 |
+
for sym_shard in shard_symbols
|
| 429 |
+
]
|
| 430 |
+
local_counts = [f.result() for f in futures]
|
| 431 |
+
|
| 432 |
+
# --- FASE 2: REDUCE (agregação) ---
|
| 433 |
+
global_counts: Dict[Tuple[str, str], int] = defaultdict(int)
|
| 434 |
+
for lc in local_counts:
|
| 435 |
+
for pair, cnt in lc.items():
|
| 436 |
+
global_counts[pair] += cnt
|
| 437 |
+
# Poda global (garante min_frequency)
|
| 438 |
+
global_counts = {
|
| 439 |
+
p: c for p, c in global_counts.items() if c >= min_frequency
|
| 440 |
+
}
|
| 441 |
+
|
| 442 |
+
if not global_counts:
|
| 443 |
+
logger.info(
|
| 444 |
+
"BBPE PARALLEL: nenhum par com freq >= %d restante. "
|
| 445 |
+
"Vocab final: %d (target %d)",
|
| 446 |
+
min_frequency, len(current_vocab), self.vocab_size,
|
| 447 |
)
|
| 448 |
+
break
|
| 449 |
|
| 450 |
+
# --- FASE 3: ESCOLHA DO MELHOR PAR ---
|
| 451 |
+
# (frequência máxima, desempate lexicográfico)
|
| 452 |
+
best_pair = max(
|
| 453 |
+
global_counts.items(), key=lambda x: (x[1], x[0])
|
| 454 |
)
|
| 455 |
+
(esq, dir_), freq = best_pair
|
| 456 |
+
|
| 457 |
+
# Gera novo token (concatenação byte-level)
|
| 458 |
+
new_token_str = esq + dir_
|
| 459 |
+
# Se já existe (raro), gera nome único
|
| 460 |
+
while new_token_str in current_vocab:
|
| 461 |
+
new_token_str += "_"
|
| 462 |
+
new_id = next_id
|
| 463 |
+
next_id += 1
|
| 464 |
+
|
| 465 |
+
# --- FASE 4: APPLY (aplicação do merge nos shards) ---
|
| 466 |
+
with ProcessPoolExecutor(max_workers=num_workers) as executor:
|
| 467 |
+
apply_futures = [
|
| 468 |
+
executor.submit(
|
| 469 |
+
apply_merge_in_shard, sym_shard, esq, dir_, new_token_str
|
| 470 |
+
)
|
| 471 |
+
for sym_shard in shard_symbols
|
| 472 |
+
]
|
| 473 |
+
shard_symbols = [f.result() for f in apply_futures]
|
| 474 |
+
|
| 475 |
+
# --- FASE 5: ATUALIZAÇÃO DO VOCABULÁRIO ---
|
| 476 |
+
current_vocab.add(new_token_str)
|
| 477 |
+
token_to_id[new_token_str] = new_id
|
| 478 |
+
merges.append((esq, dir_, new_token_str))
|
| 479 |
+
|
| 480 |
+
if show_progress and (
|
| 481 |
+
iteration <= 20
|
| 482 |
+
or iteration % 50 == 0
|
| 483 |
+
or len(current_vocab) >= self.vocab_size - 5
|
| 484 |
+
):
|
| 485 |
+
logger.info(
|
| 486 |
+
"BBPE PARALLEL iter %d: '%s' + '%s' -> '%s' "
|
| 487 |
+
"(freq=%d) | Vocab=%d/%d",
|
| 488 |
+
iteration, esq, dir_, new_token_str, freq,
|
| 489 |
+
len(current_vocab), self.vocab_size,
|
| 490 |
+
)
|
| 491 |
|
| 492 |
+
# Libera memória periodicamente (estilo Xavante)
|
| 493 |
+
if iteration % 100 == 0:
|
| 494 |
+
gc.collect()
|
|
|
|
|
|
|
| 495 |
|
| 496 |
+
# --- ETAPA FINAL: Construção do tokenizer interno ---
|
| 497 |
+
logger.info(
|
| 498 |
+
"BBPE PARALLEL: concluído. %d merges, vocab=%d. Construindo tokenizer HF...",
|
| 499 |
+
len(merges), len(token_to_id),
|
| 500 |
+
)
|
| 501 |
+
self._merges = merges
|
| 502 |
+
self._tokenizer = build_bpe_from_merges(
|
| 503 |
+
merges=merges,
|
| 504 |
+
token_to_id=token_to_id,
|
| 505 |
+
unk_token=self.unk_token,
|
| 506 |
+
add_prefix_space=self.add_prefix_space,
|
| 507 |
+
)
|
| 508 |
+
self._build_vocab_cache()
|
| 509 |
+
# Atualiza vocab_size com tamanho real
|
| 510 |
+
self.vocab_size = len(self._vocab)
|
| 511 |
|
| 512 |
+
logger.info(
|
| 513 |
+
"BBPE PARALLEL: tokenizer construído. Vocab real: %d",
|
| 514 |
+
len(self._vocab),
|
| 515 |
+
)
|
| 516 |
|
| 517 |
+
if save_path is not None:
|
| 518 |
+
self.save(save_path)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 519 |
|
| 520 |
+
# ------------------------------------------------------------------
|
| 521 |
+
# TREINAMENTO (compatibilidade — delega para paralelo)
|
| 522 |
+
# ------------------------------------------------------------------
|
| 523 |
+
def train_from_stream(
|
| 524 |
self,
|
| 525 |
+
text_iterator: Iterator[str],
|
| 526 |
save_path: Optional[Union[str, Path]] = None,
|
| 527 |
min_frequency: int = 2,
|
| 528 |
show_progress: bool = True,
|
| 529 |
+
chunk_size: int = 500,
|
| 530 |
+
num_workers: Optional[int] = None,
|
| 531 |
) -> None:
|
| 532 |
+
"""Treina o tokenizer BBPE a partir de um iterador de textos.
|
|
|
|
| 533 |
|
| 534 |
+
REFACTORED: agora delega para train_parallel_from_stream (Map-Reduce).
|
| 535 |
+
Mantém a assinatura para compatibilidade com código existente.
|
|
|
|
|
|
|
| 536 |
|
| 537 |
+
Args:
|
| 538 |
+
text_iterator: iterador yielding strings de texto
|
| 539 |
+
save_path: caminho para salvar o tokenizer JSON
|
| 540 |
+
min_frequency: frequência mínima de um par para ser mergeado
|
| 541 |
+
show_progress: exibir progresso
|
| 542 |
+
chunk_size: textos por chunk na distribuição
|
| 543 |
+
num_workers: número de processos paralelos (default: self.num_workers)
|
| 544 |
+
"""
|
| 545 |
+
workers = num_workers if num_workers is not None else self.num_workers
|
| 546 |
+
self.train_parallel_from_stream(
|
| 547 |
+
text_iterator=text_iterator,
|
| 548 |
+
save_path=Path(save_path) if save_path else None,
|
| 549 |
min_frequency=min_frequency,
|
| 550 |
+
num_workers=workers,
|
| 551 |
show_progress=show_progress,
|
| 552 |
+
chunk_size=chunk_size,
|
| 553 |
)
|
|
|
|
|
|
|
| 554 |
|
| 555 |
+
def train_from_files(
|
| 556 |
+
self,
|
| 557 |
+
file_paths: List[Union[str, Path]],
|
| 558 |
+
save_path: Optional[Union[str, Path]] = None,
|
| 559 |
+
min_frequency: int = 2,
|
| 560 |
+
show_progress: bool = True,
|
| 561 |
+
num_workers: Optional[int] = None,
|
| 562 |
+
) -> None:
|
| 563 |
+
"""Treina o tokenizer BBPE a partir de arquivos de texto.
|
| 564 |
+
|
| 565 |
+
Lê os arquivos e cria um iterador de linhas, delegando para
|
| 566 |
+
train_parallel_from_stream.
|
| 567 |
+
"""
|
| 568 |
+
def _file_line_iterator(paths):
|
| 569 |
+
for p in paths:
|
| 570 |
+
p = Path(p)
|
| 571 |
+
if not p.exists():
|
| 572 |
+
logger.warning("Arquivo não encontrado: %s", p)
|
| 573 |
+
continue
|
| 574 |
+
with open(p, "r", encoding="utf-8", errors="replace") as f:
|
| 575 |
+
for line in f:
|
| 576 |
+
line = line.strip()
|
| 577 |
+
if line:
|
| 578 |
+
yield line
|
| 579 |
+
|
| 580 |
+
workers = num_workers if num_workers is not None else self.num_workers
|
| 581 |
+
self.train_parallel_from_stream(
|
| 582 |
+
text_iterator=_file_line_iterator(file_paths),
|
| 583 |
+
save_path=Path(save_path) if save_path else None,
|
| 584 |
+
min_frequency=min_frequency,
|
| 585 |
+
num_workers=workers,
|
| 586 |
+
show_progress=show_progress,
|
| 587 |
+
chunk_size=500,
|
| 588 |
+
)
|
| 589 |
|
| 590 |
# ------------------------------------------------------------------
|
| 591 |
# Save / Load
|
| 592 |
# ------------------------------------------------------------------
|
| 593 |
def save(self, path: Union[str, Path]) -> None:
|
| 594 |
+
"""Salva o tokenizer em arquivo JSON (formato HuggingFace)."""
|
| 595 |
if self._tokenizer is None:
|
| 596 |
raise RuntimeError("Tokenizer não treinado. Chame train_*() primeiro.")
|
| 597 |
path = Path(path)
|
|
|
|
| 611 |
instance = cls() # default vocab_size
|
| 612 |
instance._tokenizer = Tokenizer.from_file(str(path))
|
| 613 |
instance._build_vocab_cache()
|
|
|
|
| 614 |
instance.vocab_size = len(instance._vocab)
|
| 615 |
logger.info(
|
| 616 |
"BBPE tokenizer loaded: %s (vocab_size=%d)",
|
|
|
|
| 655 |
f"Input too long: {len(ids)} > {max_length} and truncation=False"
|
| 656 |
)
|
| 657 |
if add_special_tokens:
|
|
|
|
| 658 |
ids = ids[:max_length - 1] + [EOS_ID] if max_length >= 1 else [EOS_ID]
|
| 659 |
else:
|
| 660 |
ids = ids[:max_length]
|
|
|
|
| 666 |
add_special_tokens: bool = False,
|
| 667 |
max_length: Optional[int] = None,
|
| 668 |
) -> List[List[int]]:
|
| 669 |
+
"""Codifica um batch de textos (sequencial)."""
|
| 670 |
return [self.encode(t, add_special_tokens, max_length) for t in texts]
|
| 671 |
|
| 672 |
+
def encode_batch_parallel(
|
| 673 |
+
self,
|
| 674 |
+
texts: List[str],
|
| 675 |
+
add_special_tokens: bool = False,
|
| 676 |
+
max_length: Optional[int] = None,
|
| 677 |
+
max_workers: Optional[int] = None,
|
| 678 |
+
) -> List[List[int]]:
|
| 679 |
+
"""Codifica um batch de textos em paralelo (ThreadPoolExecutor).
|
| 680 |
+
|
| 681 |
+
A codificação é embaraçosamente paralelizável: cada texto é codificado
|
| 682 |
+
independentemente. Usa threads (não processos) porque o tokenizers
|
| 683 |
+
library libera o GIL durante a codificação C++.
|
| 684 |
+
|
| 685 |
+
Args:
|
| 686 |
+
texts: lista de strings
|
| 687 |
+
add_special_tokens: adicionar <s>...</s>
|
| 688 |
+
max_length: truncar para este tamanho
|
| 689 |
+
max_workers: número de threads (default: min(32, len(texts)))
|
| 690 |
+
|
| 691 |
+
Returns:
|
| 692 |
+
Lista de listas de IDs.
|
| 693 |
+
"""
|
| 694 |
+
if not texts:
|
| 695 |
+
return []
|
| 696 |
+
workers = max_workers or min(32, max(1, len(texts)))
|
| 697 |
+
with ThreadPoolExecutor(max_workers=workers) as executor:
|
| 698 |
+
results = list(executor.map(
|
| 699 |
+
lambda t: self.encode(t, add_special_tokens, max_length),
|
| 700 |
+
texts,
|
| 701 |
+
))
|
| 702 |
+
return results
|
| 703 |
+
|
| 704 |
def decode(
|
| 705 |
self,
|
| 706 |
ids: List[int],
|
|
|
|
| 743 |
]
|
| 744 |
|
| 745 |
if pad_to_max_length:
|
|
|
|
| 746 |
padded = []
|
| 747 |
masks = []
|
| 748 |
for ids in batch_ids:
|
|
|
|
| 756 |
input_ids = torch.tensor(padded, dtype=torch.long)
|
| 757 |
attention_mask = torch.tensor(masks, dtype=torch.long)
|
| 758 |
else:
|
|
|
|
| 759 |
input_ids = [torch.tensor(ids, dtype=torch.long) for ids in batch_ids]
|
| 760 |
attention_mask = [torch.ones(len(ids), dtype=torch.long) for ids in batch_ids]
|
| 761 |
|
|
|
|
| 781 |
"""Tamanho real do vocabulário carregado/treinado."""
|
| 782 |
return len(self.vocab)
|
| 783 |
|
| 784 |
+
@property
|
| 785 |
+
def merges(self) -> List[Tuple[str, str, str]]:
|
| 786 |
+
"""Lista de merges aprendidos (para inspeção)."""
|
| 787 |
+
return self._merges
|
| 788 |
+
|
| 789 |
def __len__(self) -> int:
|
| 790 |
return self.actual_vocab_size
|
| 791 |
|
|
|
|
| 817 |
ids = self.encode(text, add_special_tokens=False)
|
| 818 |
decoded = self.decode(ids, skip_special_tokens=True)
|
| 819 |
|
|
|
|
|
|
|
| 820 |
expected = text
|
| 821 |
got = decoded
|
| 822 |
|
| 823 |
if expected == got:
|
| 824 |
successes += 1
|
| 825 |
else:
|
|
|
|
| 826 |
if self.add_prefix_space and got.startswith(" "):
|
| 827 |
got = got[1:]
|
| 828 |
if expected == got:
|
|
|
|
| 834 |
"ids_count": len(ids),
|
| 835 |
})
|
| 836 |
|
|
|
|
| 837 |
n_bytes = len(text.encode("utf-8"))
|
| 838 |
n_tokens = len(ids)
|
| 839 |
if n_tokens > 0:
|
|
|
|
| 853 |
|
| 854 |
|
| 855 |
# ---------------------------------------------------------------------------
|
| 856 |
+
# Byte-level alphabet (for BBPE initial alphabet) — compatibilidade
|
| 857 |
# ---------------------------------------------------------------------------
|
| 858 |
class ByteLevel:
|
| 859 |
"""Wrapper para o alfabeto byte-level (256 bytes)."""
|
|
|
|
| 861 |
@staticmethod
|
| 862 |
def alphabet() -> List[str]:
|
| 863 |
"""Retorna os 256 caracteres byte-level (Ġ-style do GPT-2)."""
|
| 864 |
+
return list(ALPHABET)
|
|
|
|
| 865 |
|
| 866 |
|
| 867 |
# ---------------------------------------------------------------------------
|
|
|
|
| 872 |
text_iterator: Optional[Iterator[str]] = None,
|
| 873 |
vocab_size: int = DEFAULT_VOCAB_SIZE,
|
| 874 |
min_frequency: int = 2,
|
| 875 |
+
num_workers: int = 4,
|
| 876 |
) -> BBPETokenizer:
|
| 877 |
+
"""Carrega um tokenizer existente ou treina um novo (paralelo).
|
| 878 |
|
| 879 |
Args:
|
| 880 |
path: caminho do arquivo JSON
|
| 881 |
text_iterator: iterador de textos para treinar (se arquivo não existe)
|
| 882 |
+
vocab_size: tamanho do vocabulário (default 16.384)
|
| 883 |
min_frequency: frequência mínima para merges
|
| 884 |
+
num_workers: número de processos paralelos no treino
|
| 885 |
|
| 886 |
Returns:
|
| 887 |
BBPETokenizer carregado/treinado
|
|
|
|
| 896 |
f"Tokenizer file {path} does not exist and no text_iterator provided"
|
| 897 |
)
|
| 898 |
|
| 899 |
+
logger.info("Training new BBPE tokenizer (parallel): %s", path)
|
| 900 |
+
tok = BBPETokenizer(vocab_size=vocab_size, num_workers=num_workers)
|
| 901 |
+
tok.train_parallel_from_stream(
|
| 902 |
text_iterator,
|
| 903 |
save_path=path,
|
| 904 |
min_frequency=min_frequency,
|
| 905 |
+
num_workers=num_workers,
|
| 906 |
)
|
| 907 |
return tok
|
| 908 |
|
|
|
|
| 921 |
"UNK_ID",
|
| 922 |
"SPECIAL_TOKENS",
|
| 923 |
"DEFAULT_VOCAB_SIZE",
|
| 924 |
+
"ALPHABET",
|
| 925 |
+
"BYTE_TO_SYMBOL",
|
| 926 |
+
"SYMBOL_TO_BYTE",
|
| 927 |
+
"bytes_to_unicode",
|
| 928 |
+
"distribute_texts",
|
| 929 |
+
"pre_tokenize_shard",
|
| 930 |
+
"count_pairs_in_shard",
|
| 931 |
+
"apply_merge_in_shard",
|
| 932 |
+
"build_bpe_from_merges",
|
| 933 |
]
|
src/bigru_t/training/dpo.py
ADDED
|
@@ -0,0 +1,266 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""dpo.py — Direct Preference Optimization (DPO) com Beta adaptativo.
|
| 2 |
+
|
| 3 |
+
═══════════════════════════════════════════════════════════════════════════════
|
| 4 |
+
STANDALONE DPO_loss — implementação desacoplada de EnhancedXavante*
|
| 5 |
+
═══════════════════════════════════════════════════════════════════════════════
|
| 6 |
+
|
| 7 |
+
A fonte original (PowerMachine/gru-ring-v13-9-2) implementava DPO como MÉTODOS
|
| 8 |
+
da classe EnhancedXavante (v1/v2/v3), acoplados a `self.forward()`, `self._logp()`,
|
| 9 |
+
etc. Isto impossibilitava o reaproveitamento direto.
|
| 10 |
+
|
| 11 |
+
Este módulo implementa DPO como FUNÇÃO LIVRE (standalone), decouplada do modelo,
|
| 12 |
+
expondo o contrato referenciado no docstring do HamiltonianWassersteinOptimizer:
|
| 13 |
+
L_dpo = DPO_loss(theta, beta, ref, preferences)
|
| 14 |
+
|
| 15 |
+
Referências matemáticas:
|
| 16 |
+
- Rafailov et al. 2023: L_DPO = -log σ(β · (log π(y_w|x)/π_ref(y_w|x)
|
| 17 |
+
- log π(y_l|x)/π_ref(y_l|x)))
|
| 18 |
+
- Prova 51 (v1): DPO básico single-preference
|
| 19 |
+
- Prova 61 (v3): β dinâmico + IPO assimétrico (9:1 punish:reward) + label smoothing
|
| 20 |
+
|
| 21 |
+
Integração com HamiltonianWassersteinOptimizer:
|
| 22 |
+
beta = compute_dynamic_beta(step, warmup_steps, beta_min, beta_max)
|
| 23 |
+
optimizer.set_beta_dpo(beta)
|
| 24 |
+
loss_dpo = dpo_loss(policy_chosen_logps, policy_rejected_logps,
|
| 25 |
+
ref_chosen_logps, ref_rejected_logps, beta=beta)
|
| 26 |
+
"""
|
| 27 |
+
from __future__ import annotations
|
| 28 |
+
|
| 29 |
+
import math
|
| 30 |
+
import logging
|
| 31 |
+
from typing import Optional, Tuple
|
| 32 |
+
|
| 33 |
+
import torch
|
| 34 |
+
import torch.nn.functional as F
|
| 35 |
+
|
| 36 |
+
logger = logging.getLogger(__name__)
|
| 37 |
+
|
| 38 |
+
|
| 39 |
+
def dpo_loss(
|
| 40 |
+
policy_chosen_logps: torch.Tensor,
|
| 41 |
+
policy_rejected_logps: torch.Tensor,
|
| 42 |
+
reference_chosen_logps: torch.Tensor,
|
| 43 |
+
reference_rejected_logps: torch.Tensor,
|
| 44 |
+
beta: float = 0.1,
|
| 45 |
+
label_smoothing: float = 0.0,
|
| 46 |
+
use_ipo: bool = False,
|
| 47 |
+
asymmetric_ratio: float = 1.0,
|
| 48 |
+
) -> torch.Tensor:
|
| 49 |
+
"""Computa a perda DPO (Direct Preference Optimization).
|
| 50 |
+
|
| 51 |
+
L_DPO = -log σ(β · (Δπ_chosen - Δπ_rejected))
|
| 52 |
+
onde Δπ = log π(y|x) - log π_ref(y|x)
|
| 53 |
+
|
| 54 |
+
Args:
|
| 55 |
+
policy_chosen_logps: log π(y_chosen | x), shape (batch,)
|
| 56 |
+
policy_rejected_logps: log π(y_rejected | x), shape (batch,)
|
| 57 |
+
reference_chosen_logps: log π_ref(y_chosen | x), shape (batch,)
|
| 58 |
+
reference_rejected_logps: log π_ref(y_rejected | x), shape (batch,)
|
| 59 |
+
beta: temperatura inversa (controle de margem). Maior β → mais confiante.
|
| 60 |
+
label_smoothing: ε ∈ [0, 0.5] — suaviza rótulos (0 = DPO clássico)
|
| 61 |
+
use_ipo: se True, usa Identity Preference Optimization (sem sigmoid)
|
| 62 |
+
asymmetric_ratio: ratio punição:recompensa (default 1.0 = simétrico).
|
| 63 |
+
v3 usava 9.0 (pune chosen-errado 9x mais que recompensa chosen-certo).
|
| 64 |
+
|
| 65 |
+
Returns:
|
| 66 |
+
loss: escalar (média sobre o batch)
|
| 67 |
+
|
| 68 |
+
Prova 51: L_DPO é diferenciável e convexa em (π - π_ref) para β fixo.
|
| 69 |
+
Prova 61: β dinâmico + IPO + label smoothing → convergência mais estável.
|
| 70 |
+
"""
|
| 71 |
+
# Delta de log-probs: policy vs reference
|
| 72 |
+
pi_logratios_chosen = policy_chosen_logps - reference_chosen_logps
|
| 73 |
+
pi_logratios_rejected = policy_rejected_logps - reference_rejected_logps
|
| 74 |
+
|
| 75 |
+
# logits = β * (chosen - rejected)
|
| 76 |
+
logits = beta * (pi_logratios_chosen - asymmetric_ratio * pi_logratios_rejected)
|
| 77 |
+
|
| 78 |
+
if use_ipo:
|
| 79 |
+
# IPO (Identity Preference Optimization): perda quadrática
|
| 80 |
+
# L_IPO = (logits - 1/2)² [palma2024]
|
| 81 |
+
loss = (logits - 1.0 / 2.0).pow(2).mean()
|
| 82 |
+
else:
|
| 83 |
+
# DPO clássico com label smoothing
|
| 84 |
+
# L = -(1-ε)·log σ(logits) - ε·log σ(-logits)
|
| 85 |
+
# = -(1-ε)·log σ(logits) - ε·log(1 - σ(logits))
|
| 86 |
+
if label_smoothing > 0:
|
| 87 |
+
loss = (
|
| 88 |
+
-(1 - label_smoothing) * F.logsigmoid(logits)
|
| 89 |
+
- label_smoothing * F.logsigmoid(-logits)
|
| 90 |
+
).mean()
|
| 91 |
+
else:
|
| 92 |
+
loss = -F.logsigmoid(logits).mean()
|
| 93 |
+
|
| 94 |
+
return loss
|
| 95 |
+
|
| 96 |
+
|
| 97 |
+
def compute_sequence_logps(
|
| 98 |
+
logits: torch.Tensor,
|
| 99 |
+
labels: torch.Tensor,
|
| 100 |
+
attention_mask: Optional[torch.Tensor] = None,
|
| 101 |
+
pad_token_id: int = 1,
|
| 102 |
+
) -> torch.Tensor:
|
| 103 |
+
"""Computa log π(y | x) = Σ_t log π(y_t | y_{<t}, x) para cada sequência.
|
| 104 |
+
|
| 105 |
+
Args:
|
| 106 |
+
logits: (batch, T, vocab) — logits do modelo
|
| 107 |
+
labels: (batch, T) — token IDs alvo (shifted: labels[t] é o target de logits[t-1])
|
| 108 |
+
attention_mask: (batch, T) — 1 para tokens reais, 0 para pad
|
| 109 |
+
pad_token_id: ID do padding (ignorado na soma)
|
| 110 |
+
|
| 111 |
+
Returns:
|
| 112 |
+
logps: (batch,) — log-probabilidade de cada sequência
|
| 113 |
+
"""
|
| 114 |
+
# Shift: logits[:-1] prediz labels[1:]
|
| 115 |
+
shift_logits = logits[:, :-1, :].contiguous()
|
| 116 |
+
shift_labels = labels[:, 1:].contiguous()
|
| 117 |
+
|
| 118 |
+
# Log-softmax sobre vocab
|
| 119 |
+
log_probs = F.log_softmax(shift_logits, dim=-1) # (batch, T-1, vocab)
|
| 120 |
+
|
| 121 |
+
# Gather log-prob do token correto
|
| 122 |
+
gathered = log_probs.gather(
|
| 123 |
+
2, shift_labels.unsqueeze(-1)
|
| 124 |
+
).squeeze(-1) # (batch, T-1)
|
| 125 |
+
|
| 126 |
+
# Mask: ignorar padding
|
| 127 |
+
if attention_mask is not None:
|
| 128 |
+
shift_mask = attention_mask[:, 1:].contiguous().float()
|
| 129 |
+
gathered = gathered * shift_mask
|
| 130 |
+
|
| 131 |
+
# Soma sobre a sequência
|
| 132 |
+
logps = gathered.sum(dim=-1) # (batch,)
|
| 133 |
+
return logps
|
| 134 |
+
|
| 135 |
+
|
| 136 |
+
def compute_dynamic_beta(
|
| 137 |
+
step: int,
|
| 138 |
+
warmup_steps: int = 100,
|
| 139 |
+
beta_min: float = 0.05,
|
| 140 |
+
beta_max: float = 0.5,
|
| 141 |
+
decay: float = 0.999,
|
| 142 |
+
) -> float:
|
| 143 |
+
"""Computa β(t) dinâmico para DPO.
|
| 144 |
+
|
| 145 |
+
Durante warmup: β cresce linearmente de beta_min a beta_max.
|
| 146 |
+
Após warmup: β decai exponencialmente (decay^step) até beta_min.
|
| 147 |
+
|
| 148 |
+
Args:
|
| 149 |
+
step: passo atual de treino
|
| 150 |
+
warmup_steps: passos de warmup linear
|
| 151 |
+
beta_min: β mínimo (após decair)
|
| 152 |
+
beta_max: β máximo (topo do warmup)
|
| 153 |
+
decay: fator de decaimento exponencial por passo
|
| 154 |
+
|
| 155 |
+
Returns:
|
| 156 |
+
beta: float no intervalo [beta_min, beta_max]
|
| 157 |
+
|
| 158 |
+
Prova 61: β dinâmico estabiliza convergência — alto no início
|
| 159 |
+
(exploração), baixo no fim (exploitação).
|
| 160 |
+
"""
|
| 161 |
+
if step < warmup_steps:
|
| 162 |
+
# Warmup linear
|
| 163 |
+
progress = step / max(1, warmup_steps)
|
| 164 |
+
return beta_min + (beta_max - beta_min) * progress
|
| 165 |
+
else:
|
| 166 |
+
# Decay exponencial após warmup
|
| 167 |
+
excess = step - warmup_steps
|
| 168 |
+
return max(beta_min, beta_max * (decay ** excess))
|
| 169 |
+
|
| 170 |
+
|
| 171 |
+
def dpo_step(
|
| 172 |
+
model,
|
| 173 |
+
chosen_input_ids: torch.Tensor,
|
| 174 |
+
chosen_attention_mask: torch.Tensor,
|
| 175 |
+
rejected_input_ids: torch.Tensor,
|
| 176 |
+
rejected_attention_mask: torch.Tensor,
|
| 177 |
+
ref_chosen_logps: torch.Tensor,
|
| 178 |
+
ref_rejected_logps: torch.Tensor,
|
| 179 |
+
beta: float = 0.1,
|
| 180 |
+
label_smoothing: float = 0.0,
|
| 181 |
+
use_ipo: bool = False,
|
| 182 |
+
temperature: float = 1.0,
|
| 183 |
+
) -> Tuple[torch.Tensor, dict]:
|
| 184 |
+
"""Executa um passo DPO completo: forward + loss.
|
| 185 |
+
|
| 186 |
+
Args:
|
| 187 |
+
model: modelo com forward(input_ids) -> (logits,) ou (y_hat, delta)
|
| 188 |
+
chosen_input_ids: (batch, T) tokens da resposta preferida
|
| 189 |
+
chosen_attention_mask: (batch, T)
|
| 190 |
+
rejected_input_ids: (batch, T) tokens da resposta rejeitada
|
| 191 |
+
rejected_attention_mask: (batch, T)
|
| 192 |
+
ref_chosen_logps: (batch,) log-probs de referência para chosen
|
| 193 |
+
(pré-computados com modelo congelado, sem grad)
|
| 194 |
+
ref_rejected_logps: (batch,) log-probs de referência para rejected
|
| 195 |
+
beta: temperatura DPO
|
| 196 |
+
label_smoothing: ε ∈ [0, 0.5]
|
| 197 |
+
use_ipo: usar IPO em vez de DPO clássico
|
| 198 |
+
temperature: temperatura do softmax do modelo (Lema 1)
|
| 199 |
+
|
| 200 |
+
Returns:
|
| 201 |
+
(loss, metrics_dict)
|
| 202 |
+
"""
|
| 203 |
+
# Forward chosen
|
| 204 |
+
out_c = model(chosen_input_ids, temperature=temperature, use_hypothesis=False)
|
| 205 |
+
logits_c = out_c[0] if isinstance(out_c, tuple) else out_c
|
| 206 |
+
if logits_c.dim() == 2:
|
| 207 |
+
# Modelo produz (batch, vocab) — apenas 1 logit por amostra
|
| 208 |
+
# Não é possível computar sequence logp; usar logit do último token
|
| 209 |
+
# como proxy (aproximação para bug-detection)
|
| 210 |
+
policy_chosen_logps = F.log_softmax(logits_c, dim=-1).gather(
|
| 211 |
+
1, chosen_input_ids[:, -1:].clamp(0, logits_c.size(-1) - 1)
|
| 212 |
+
).squeeze(-1)
|
| 213 |
+
else:
|
| 214 |
+
policy_chosen_logps = compute_sequence_logps(
|
| 215 |
+
logits_c, chosen_input_ids, chosen_attention_mask,
|
| 216 |
+
pad_token_id=getattr(model.config, "pad_token_id", 1),
|
| 217 |
+
)
|
| 218 |
+
|
| 219 |
+
# Forward rejected
|
| 220 |
+
out_r = model(rejected_input_ids, temperature=temperature, use_hypothesis=False)
|
| 221 |
+
logits_r = out_r[0] if isinstance(out_r, tuple) else out_r
|
| 222 |
+
if logits_r.dim() == 2:
|
| 223 |
+
policy_rejected_logps = F.log_softmax(logits_r, dim=-1).gather(
|
| 224 |
+
1, rejected_input_ids[:, -1:].clamp(0, logits_r.size(-1) - 1)
|
| 225 |
+
).squeeze(-1)
|
| 226 |
+
else:
|
| 227 |
+
policy_rejected_logps = compute_sequence_logps(
|
| 228 |
+
logits_r, rejected_input_ids, rejected_attention_mask,
|
| 229 |
+
pad_token_id=getattr(model.config, "pad_token_id", 1),
|
| 230 |
+
)
|
| 231 |
+
|
| 232 |
+
# DPO loss
|
| 233 |
+
loss = dpo_loss(
|
| 234 |
+
policy_chosen_logps=policy_chosen_logps,
|
| 235 |
+
policy_rejected_logps=policy_rejected_logps,
|
| 236 |
+
reference_chosen_logps=ref_chosen_logps,
|
| 237 |
+
reference_rejected_logps=ref_rejected_logps,
|
| 238 |
+
beta=beta,
|
| 239 |
+
label_smoothing=label_smoothing,
|
| 240 |
+
use_ipo=use_ipo,
|
| 241 |
+
)
|
| 242 |
+
|
| 243 |
+
# Metrics
|
| 244 |
+
with torch.no_grad():
|
| 245 |
+
chosen_rewards = beta * (policy_chosen_logps - ref_chosen_logps)
|
| 246 |
+
rejected_rewards = beta * (policy_rejected_logps - ref_rejected_logps)
|
| 247 |
+
accuracy = (chosen_rewards > rejected_rewards).float().mean()
|
| 248 |
+
margin = (chosen_rewards - rejected_rewards).mean()
|
| 249 |
+
|
| 250 |
+
metrics = {
|
| 251 |
+
"dpo_loss": float(loss.item()),
|
| 252 |
+
"beta": beta,
|
| 253 |
+
"chosen_reward": float(chosen_rewards.mean().item()),
|
| 254 |
+
"rejected_reward": float(rejected_rewards.mean().item()),
|
| 255 |
+
"accuracy": float(accuracy.item()),
|
| 256 |
+
"margin": float(margin.item()),
|
| 257 |
+
}
|
| 258 |
+
return loss, metrics
|
| 259 |
+
|
| 260 |
+
|
| 261 |
+
__all__ = [
|
| 262 |
+
"dpo_loss",
|
| 263 |
+
"compute_sequence_logps",
|
| 264 |
+
"compute_dynamic_beta",
|
| 265 |
+
"dpo_step",
|
| 266 |
+
]
|
src/bigru_t/training/meta_configurator.py
CHANGED
|
@@ -110,21 +110,24 @@ class MetaConfigurator:
|
|
| 110 |
loss_val = criterion(y_hat_main, y_val)
|
| 111 |
|
| 112 |
# Proxy de agudeza: ||∇_θ L_val||^2
|
| 113 |
-
# create_graph=True
|
|
|
|
|
|
|
| 114 |
params = [p for p in self.model.parameters() if p.requires_grad]
|
| 115 |
grads = torch.autograd.grad(
|
| 116 |
loss_val,
|
| 117 |
params,
|
| 118 |
-
create_graph=
|
| 119 |
-
retain_graph=True,
|
| 120 |
allow_unused=True,
|
| 121 |
)
|
| 122 |
sharpness = sum(
|
| 123 |
-
(g ** 2).sum() for g in grads if g is not None
|
| 124 |
)
|
| 125 |
|
| 126 |
-
# Meta-perda:
|
| 127 |
-
|
|
|
|
| 128 |
|
| 129 |
# Atualiza log_temperature e log_tau via gradiente de meta_loss
|
| 130 |
self.meta_optim.zero_grad()
|
|
|
|
| 110 |
loss_val = criterion(y_hat_main, y_val)
|
| 111 |
|
| 112 |
# Proxy de agudeza: ||∇_θ L_val||^2
|
| 113 |
+
# BUG FIX: create_graph=True causava OOM-Killer (second-order graph dobra RAM)
|
| 114 |
+
# Correção: usar create_graph=False (first-order approximation)
|
| 115 |
+
# retain_graph=True necessário para meta_loss.backward() abaixo
|
| 116 |
params = [p for p in self.model.parameters() if p.requires_grad]
|
| 117 |
grads = torch.autograd.grad(
|
| 118 |
loss_val,
|
| 119 |
params,
|
| 120 |
+
create_graph=False, # FIX: era True → OOM
|
| 121 |
+
retain_graph=True, # mantém grafo para meta_loss.backward()
|
| 122 |
allow_unused=True,
|
| 123 |
)
|
| 124 |
sharpness = sum(
|
| 125 |
+
(g.detach() ** 2).sum() for g in grads if g is not None
|
| 126 |
)
|
| 127 |
|
| 128 |
+
# Meta-perda: apenas loss_val (sharpness é monitor only, detached)
|
| 129 |
+
# (O sharpness penalty requereria second-order, que é OOM-prohibitive)
|
| 130 |
+
meta_loss = loss_val
|
| 131 |
|
| 132 |
# Atualiza log_temperature e log_tau via gradiente de meta_loss
|
| 133 |
self.meta_optim.zero_grad()
|
src/bigru_t/training/trainer.py
CHANGED
|
@@ -43,13 +43,16 @@ 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 |
|
| 47 |
logger = logging.getLogger(__name__)
|
| 48 |
|
| 49 |
|
| 50 |
@dataclass
|
| 51 |
class TrainerConfig:
|
| 52 |
-
"""Configuração do treino (bug-detection)."""
|
| 53 |
# Epochs (FIXO em 2 por especificação do usuário)
|
| 54 |
epochs: int = 2
|
| 55 |
|
|
@@ -62,10 +65,24 @@ class TrainerConfig:
|
|
| 62 |
per_device_batch_size: int = 1
|
| 63 |
grad_accum: int = 4
|
| 64 |
|
| 65 |
-
# Optimizer
|
|
|
|
| 66 |
lr: float = 1e-3 # from-scratch: LR mais alto que fine-tune
|
| 67 |
weight_decay: float = 0.01
|
| 68 |
max_grad_norm: float = 5.0
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 69 |
|
| 70 |
# Meta-configurator (Lema 4)
|
| 71 |
meta_interval: int = 10 # a cada N micro-batches
|
|
@@ -81,6 +98,13 @@ class TrainerConfig:
|
|
| 81 |
disk_min_free_gb: float = 1.0
|
| 82 |
loss_patience: int = 30
|
| 83 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 84 |
# Logging
|
| 85 |
log_every: int = 5
|
| 86 |
save_temp_every: int = 10 # salvar checkpoint temporário a cada N steps
|
|
@@ -123,15 +147,39 @@ class BiGRU_T_Trainer:
|
|
| 123 |
os.environ.setdefault("MKL_NUM_THREADS", "2")
|
| 124 |
os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")
|
| 125 |
|
| 126 |
-
# Otimizador principal
|
| 127 |
-
#
|
| 128 |
-
|
| 129 |
-
|
| 130 |
-
|
| 131 |
-
|
| 132 |
-
|
| 133 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 134 |
)
|
|
|
|
| 135 |
|
| 136 |
# Meta-configurator (Lema 4)
|
| 137 |
self.meta_cfg = MetaConfigurator(
|
|
@@ -243,9 +291,13 @@ class BiGRU_T_Trainer:
|
|
| 243 |
logger.info(f" batch_size: {cfg.per_device_batch_size}")
|
| 244 |
logger.info(f" grad_accum: {cfg.grad_accum}")
|
| 245 |
logger.info(f" lr: {cfg.lr}")
|
|
|
|
| 246 |
logger.info(f" max_seq_len: {cfg.max_seq_len}")
|
| 247 |
logger.info(f" use_hypothesis: {cfg.use_hypothesis}")
|
|
|
|
| 248 |
logger.info(f" meta_interval: {cfg.meta_interval}")
|
|
|
|
|
|
|
| 249 |
logger.info("=" * 70)
|
| 250 |
|
| 251 |
# Conta parâmetros
|
|
@@ -256,17 +308,27 @@ class BiGRU_T_Trainer:
|
|
| 256 |
|
| 257 |
t_start = time.time()
|
| 258 |
self.model.train()
|
|
|
|
| 259 |
|
| 260 |
for epoch in range(cfg.epochs):
|
| 261 |
logger.info(f"\n--- Epoch {epoch+1}/{cfg.epochs} ---")
|
| 262 |
epoch_start = time.time()
|
| 263 |
epoch_losses = []
|
|
|
|
| 264 |
|
| 265 |
# Itera sobre samples (loop circular se n_train < batches necessários)
|
| 266 |
sample_idx = 0
|
| 267 |
micro_in_epoch = 0
|
| 268 |
|
| 269 |
while micro_in_epoch < n_train:
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 270 |
# Pega batch
|
| 271 |
batch_samples = []
|
| 272 |
for _ in range(cfg.per_device_batch_size):
|
|
@@ -284,38 +346,45 @@ class BiGRU_T_Trainer:
|
|
| 284 |
logger.warning(f"Tokenize erro micro {self.micro_batch_idx}: {e}")
|
| 285 |
continue
|
| 286 |
|
| 287 |
-
# Forward + loss
|
|
|
|
| 288 |
try:
|
| 289 |
-
#
|
| 290 |
-
|
| 291 |
-
# Após alguns steps, comparamos loss_main com tau
|
| 292 |
-
use_hyp = cfg.use_hypothesis and (self.global_step > 0)
|
| 293 |
-
if use_hyp:
|
| 294 |
-
# Primeiro forward sem hipótese para checar loss_main vs tau
|
| 295 |
-
loss_main_pre, _, _ = self._forward_loss(
|
| 296 |
-
input_ids, attn_mask, labels, use_hyp=False
|
| 297 |
-
)
|
| 298 |
-
use_hyp = bool(loss_main_pre.item() > self.meta_cfg.tau)
|
| 299 |
-
|
| 300 |
-
loss_main, loss_hyp, y_hat = self._forward_loss(
|
| 301 |
-
input_ids, attn_mask, labels, use_hyp=use_hyp
|
| 302 |
-
)
|
| 303 |
|
| 304 |
-
#
|
| 305 |
-
|
| 306 |
-
# mas para economizar compute, pegamos via hook)
|
| 307 |
-
# Solução: fazer forward novamente com return_aux (caro mas correto)
|
| 308 |
-
# Alternativa: somar entropy_reg na loss_main sempre que use_hyp=False
|
| 309 |
-
# Aqui optamos por: somar entropy_reg à loss_main sempre
|
| 310 |
-
_, _, aux = self.model(
|
| 311 |
input_ids,
|
| 312 |
temperature=self.meta_cfg.temperature,
|
| 313 |
-
use_hypothesis=
|
| 314 |
-
stop_grad_hyp=
|
| 315 |
return_aux=True,
|
| 316 |
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 317 |
loss_main_total = loss_main + aux["entropy_reg"]
|
| 318 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 319 |
except RuntimeError as e:
|
| 320 |
err = str(e).lower()
|
| 321 |
if "out of memory" in err:
|
|
@@ -328,20 +397,40 @@ class BiGRU_T_Trainer:
|
|
| 328 |
traceback.print_exc()
|
| 329 |
break
|
| 330 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 331 |
# Backward + step
|
|
|
|
| 332 |
if use_hyp and loss_hyp.requires_grad:
|
| 333 |
# Lema 2: gradient surgery
|
| 334 |
apply_gradient_surgery(self.model, loss_main_total, loss_hyp)
|
| 335 |
# Clip
|
| 336 |
torch.nn.utils.clip_grad_norm_(self.model.parameters(), cfg.max_grad_norm)
|
| 337 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 338 |
self.optimizer.zero_grad()
|
| 339 |
else:
|
| 340 |
# Apenas loss_main
|
| 341 |
self.optimizer.zero_grad()
|
| 342 |
loss_main_total.backward()
|
| 343 |
torch.nn.utils.clip_grad_norm_(self.model.parameters(), cfg.max_grad_norm)
|
| 344 |
-
self.optimizer
|
|
|
|
|
|
|
|
|
|
|
|
|
| 345 |
|
| 346 |
self.micro_batch_idx += 1
|
| 347 |
|
|
@@ -388,6 +477,17 @@ class BiGRU_T_Trainer:
|
|
| 388 |
if self.global_step % cfg.save_temp_every == 0:
|
| 389 |
self._save_temp_checkpoint(epoch)
|
| 390 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 391 |
# Meta-configurator (Lema 4)
|
| 392 |
if self.global_step % cfg.meta_interval == 0 and self.val_samples:
|
| 393 |
self._run_meta_update()
|
|
@@ -395,6 +495,9 @@ class BiGRU_T_Trainer:
|
|
| 395 |
if self.killed_reason:
|
| 396 |
break
|
| 397 |
|
|
|
|
|
|
|
|
|
|
| 398 |
# Fim da época
|
| 399 |
epoch_avg = sum(epoch_losses) / max(1, len(epoch_losses))
|
| 400 |
epoch_ppl = math.exp(min(20, epoch_avg)) if epoch_avg < 20 else float("inf")
|
|
@@ -403,6 +506,25 @@ class BiGRU_T_Trainer:
|
|
| 403 |
f"({time.time()-epoch_start:.1f}s)"
|
| 404 |
)
|
| 405 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 406 |
# Fim do treino
|
| 407 |
elapsed_total = time.time() - t_start
|
| 408 |
final_loss = self.losses_history[-1] if self.losses_history else float("nan")
|
|
@@ -417,6 +539,7 @@ class BiGRU_T_Trainer:
|
|
| 417 |
"epochs_completed": epoch + 1 if not self.killed_reason else epoch,
|
| 418 |
"killed": self.killed_reason is not None,
|
| 419 |
"kill_reason": self.killed_reason,
|
|
|
|
| 420 |
"global_step": self.global_step,
|
| 421 |
"micro_batches": self.micro_batch_idx,
|
| 422 |
"best_loss": self.best_loss,
|
|
@@ -427,6 +550,9 @@ class BiGRU_T_Trainer:
|
|
| 427 |
"params_total": params["total"],
|
| 428 |
"params_M": params["total_M"],
|
| 429 |
"monitor_summary": ks_summary,
|
|
|
|
|
|
|
|
|
|
| 430 |
}
|
| 431 |
|
| 432 |
logger.info("=" * 70)
|
|
|
|
| 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 |
|
| 50 |
logger = logging.getLogger(__name__)
|
| 51 |
|
| 52 |
|
| 53 |
@dataclass
|
| 54 |
class TrainerConfig:
|
| 55 |
+
"""Configuração do treino (bug-detection + 15M params)."""
|
| 56 |
# Epochs (FIXO em 2 por especificação do usuário)
|
| 57 |
epochs: int = 2
|
| 58 |
|
|
|
|
| 65 |
per_device_batch_size: int = 1
|
| 66 |
grad_accum: int = 4
|
| 67 |
|
| 68 |
+
# Optimizer (HamiltonianWasserstein ATIVADO)
|
| 69 |
+
optimizer_type: str = "hamiltonian_wasserstein" # "adamw" | "hamiltonian_wasserstein"
|
| 70 |
lr: float = 1e-3 # from-scratch: LR mais alto que fine-tune
|
| 71 |
weight_decay: float = 0.01
|
| 72 |
max_grad_norm: float = 5.0
|
| 73 |
+
# HW optimizer params
|
| 74 |
+
hw_lr_amp: float = 0.3 # amplitude do LR cíclico (Van der Pol)
|
| 75 |
+
hw_lr_freq: float = 0.01 # frequência angular do LR cíclico
|
| 76 |
+
hw_sigma_w: float = 1.0 # escala W₂
|
| 77 |
+
hw_sigma_rep: float = 0.1 # escala repulsão topológica
|
| 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 (Beta adaptativo)
|
| 82 |
+
use_dpo: bool = False # DPO requer preferências; desativado por padrão
|
| 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
|
|
|
|
| 98 |
disk_min_free_gb: float = 1.0
|
| 99 |
loss_patience: int = 30
|
| 100 |
|
| 101 |
+
# Time budget (estilo Xavante — streaming com timed steps)
|
| 102 |
+
max_total_time_s: float = 1800.0 # 30 min máx total
|
| 103 |
+
max_per_epoch_s: float = 900.0 # 15 min máx por época
|
| 104 |
+
|
| 105 |
+
# Memory cleanup agressivo (estilo Xavante)
|
| 106 |
+
cleanup_every_n_steps: int = 10 # aggressive_cleanup a cada N optimizer steps
|
| 107 |
+
|
| 108 |
# Logging
|
| 109 |
log_every: int = 5
|
| 110 |
save_temp_every: int = 10 # salvar checkpoint temporário a cada N steps
|
|
|
|
| 147 |
os.environ.setdefault("MKL_NUM_THREADS", "2")
|
| 148 |
os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")
|
| 149 |
|
| 150 |
+
# Otimizador principal — HamiltonianWasserstein ATIVADO por padrão
|
| 151 |
+
# (AdamW + W₂ adaptativo + repulsão topológica + LR cíclico + pruning opcional)
|
| 152 |
+
if config.optimizer_type == "hamiltonian_wasserstein":
|
| 153 |
+
self.optimizer = HamiltonianWassersteinOptimizer(
|
| 154 |
+
model.parameters(),
|
| 155 |
+
lr=config.lr,
|
| 156 |
+
betas=(0.9, 0.95),
|
| 157 |
+
eps=1e-8,
|
| 158 |
+
weight_decay=config.weight_decay,
|
| 159 |
+
lr_amp=config.hw_lr_amp,
|
| 160 |
+
lr_freq=config.hw_lr_freq,
|
| 161 |
+
sigma_w=config.hw_sigma_w,
|
| 162 |
+
sigma_rep=config.hw_sigma_rep,
|
| 163 |
+
prune_every_n=config.hw_prune_every_n,
|
| 164 |
+
ref_buffer_size=config.hw_ref_buffer_size,
|
| 165 |
+
)
|
| 166 |
+
logger.info("Optimizer: HamiltonianWassersteinOptimizer (ATIVADO)")
|
| 167 |
+
else:
|
| 168 |
+
self.optimizer = torch.optim.AdamW(
|
| 169 |
+
model.parameters(),
|
| 170 |
+
lr=config.lr,
|
| 171 |
+
betas=(0.9, 0.95),
|
| 172 |
+
weight_decay=config.weight_decay,
|
| 173 |
+
eps=1e-8,
|
| 174 |
+
)
|
| 175 |
+
logger.info("Optimizer: AdamW (fallback)")
|
| 176 |
+
|
| 177 |
+
# Time budget (estilo Xavante)
|
| 178 |
+
self.time_budget = TimeBudget(
|
| 179 |
+
max_total_s=config.max_total_time_s,
|
| 180 |
+
max_per_epoch_s=config.max_per_epoch_s,
|
| 181 |
)
|
| 182 |
+
self.step_timer = StepTimer(expected_s=2.0) # espera ~2s por step
|
| 183 |
|
| 184 |
# Meta-configurator (Lema 4)
|
| 185 |
self.meta_cfg = MetaConfigurator(
|
|
|
|
| 291 |
logger.info(f" batch_size: {cfg.per_device_batch_size}")
|
| 292 |
logger.info(f" grad_accum: {cfg.grad_accum}")
|
| 293 |
logger.info(f" lr: {cfg.lr}")
|
| 294 |
+
logger.info(f" optimizer: {cfg.optimizer_type}")
|
| 295 |
logger.info(f" max_seq_len: {cfg.max_seq_len}")
|
| 296 |
logger.info(f" use_hypothesis: {cfg.use_hypothesis}")
|
| 297 |
+
logger.info(f" use_dpo: {cfg.use_dpo}")
|
| 298 |
logger.info(f" meta_interval: {cfg.meta_interval}")
|
| 299 |
+
logger.info(f" time_budget: {cfg.max_total_time_s}s total, {cfg.max_per_epoch_s}s/epoch")
|
| 300 |
+
logger.info(f" cleanup_every: {cfg.cleanup_every_n_steps} steps")
|
| 301 |
logger.info("=" * 70)
|
| 302 |
|
| 303 |
# Conta parâmetros
|
|
|
|
| 308 |
|
| 309 |
t_start = time.time()
|
| 310 |
self.model.train()
|
| 311 |
+
self.time_budget.reset_total()
|
| 312 |
|
| 313 |
for epoch in range(cfg.epochs):
|
| 314 |
logger.info(f"\n--- Epoch {epoch+1}/{cfg.epochs} ---")
|
| 315 |
epoch_start = time.time()
|
| 316 |
epoch_losses = []
|
| 317 |
+
self.time_budget.reset_epoch()
|
| 318 |
|
| 319 |
# Itera sobre samples (loop circular se n_train < batches necessários)
|
| 320 |
sample_idx = 0
|
| 321 |
micro_in_epoch = 0
|
| 322 |
|
| 323 |
while micro_in_epoch < n_train:
|
| 324 |
+
# Time budget check (estilo Xavante)
|
| 325 |
+
if self.time_budget.is_over() or self.time_budget.is_epoch_over():
|
| 326 |
+
logger.warning(
|
| 327 |
+
"Time budget estourado: total=%s epoch=%s",
|
| 328 |
+
self.time_budget.is_over(), self.time_budget.is_epoch_over(),
|
| 329 |
+
)
|
| 330 |
+
break
|
| 331 |
+
|
| 332 |
# Pega batch
|
| 333 |
batch_samples = []
|
| 334 |
for _ in range(cfg.per_device_batch_size):
|
|
|
|
| 346 |
logger.warning(f"Tokenize erro micro {self.micro_batch_idx}: {e}")
|
| 347 |
continue
|
| 348 |
|
| 349 |
+
# Forward + loss — OTIMIZADO: single forward pass
|
| 350 |
+
# (antes: 3 forwards por micro-batch; agora: 1 forward com return_aux)
|
| 351 |
try:
|
| 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)
|
| 365 |
+
loss_main = F.cross_entropy(
|
| 366 |
+
y_hat, target, ignore_index=self.model.config.pad_token_id
|
| 367 |
+
)
|
| 368 |
loss_main_total = loss_main + aux["entropy_reg"]
|
| 369 |
|
| 370 |
+
# Decide se usa hipótese baseado no tau (Lema 3)
|
| 371 |
+
# use_hyp = True se loss_main > tau (modelo está "lutando")
|
| 372 |
+
use_hyp = (
|
| 373 |
+
cfg.use_hypothesis
|
| 374 |
+
and self.global_step > 0
|
| 375 |
+
and bool(loss_main.item() > self.meta_cfg.tau)
|
| 376 |
+
)
|
| 377 |
+
|
| 378 |
+
# Loss de hipótese: sobre y_final = y_hat + delta
|
| 379 |
+
if use_hyp:
|
| 380 |
+
y_final = y_hat + delta
|
| 381 |
+
loss_hyp = F.cross_entropy(
|
| 382 |
+
y_final, target,
|
| 383 |
+
ignore_index=self.model.config.pad_token_id,
|
| 384 |
+
)
|
| 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 |
traceback.print_exc()
|
| 398 |
break
|
| 399 |
|
| 400 |
+
# DPO beta dinâmico (se DPO ativo)
|
| 401 |
+
if cfg.use_dpo:
|
| 402 |
+
beta_t = compute_dynamic_beta(
|
| 403 |
+
self.global_step,
|
| 404 |
+
warmup_steps=cfg.dpo_warmup_steps,
|
| 405 |
+
beta_min=cfg.dpo_beta_min,
|
| 406 |
+
beta_max=cfg.dpo_beta_max,
|
| 407 |
+
)
|
| 408 |
+
if hasattr(self.optimizer, "set_beta_dpo"):
|
| 409 |
+
self.optimizer.set_beta_dpo(beta_t)
|
| 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
|
| 419 |
+
if isinstance(self.optimizer, HamiltonianWassersteinOptimizer):
|
| 420 |
+
self.optimizer.step(loss=float(loss_main_total.item()))
|
| 421 |
+
else:
|
| 422 |
+
self.optimizer.step()
|
| 423 |
self.optimizer.zero_grad()
|
| 424 |
else:
|
| 425 |
# Apenas loss_main
|
| 426 |
self.optimizer.zero_grad()
|
| 427 |
loss_main_total.backward()
|
| 428 |
torch.nn.utils.clip_grad_norm_(self.model.parameters(), cfg.max_grad_norm)
|
| 429 |
+
if isinstance(self.optimizer, HamiltonianWassersteinOptimizer):
|
| 430 |
+
self.optimizer.step(loss=float(loss_main_total.item()))
|
| 431 |
+
else:
|
| 432 |
+
self.optimizer.step()
|
| 433 |
+
self.step_timer.stop(step_label=str(self.global_step))
|
| 434 |
|
| 435 |
self.micro_batch_idx += 1
|
| 436 |
|
|
|
|
| 477 |
if self.global_step % cfg.save_temp_every == 0:
|
| 478 |
self._save_temp_checkpoint(epoch)
|
| 479 |
|
| 480 |
+
# Memory cleanup agressivo (estilo Xavante)
|
| 481 |
+
if self.global_step % cfg.cleanup_every_n_steps == 0:
|
| 482 |
+
cleanup_info = aggressive_cleanup(verbose=False)
|
| 483 |
+
if self.global_step % (cfg.cleanup_every_n_steps * 5) == 0:
|
| 484 |
+
logger.info(
|
| 485 |
+
" cleanup step %d: RSS %.1f MB (freed %.1f MB)",
|
| 486 |
+
self.global_step,
|
| 487 |
+
cleanup_info["rss_after_mb"],
|
| 488 |
+
cleanup_info["freed_mb"],
|
| 489 |
+
)
|
| 490 |
+
|
| 491 |
# Meta-configurator (Lema 4)
|
| 492 |
if self.global_step % cfg.meta_interval == 0 and self.val_samples:
|
| 493 |
self._run_meta_update()
|
|
|
|
| 495 |
if self.killed_reason:
|
| 496 |
break
|
| 497 |
|
| 498 |
+
# Fim da época — cleanup agressivo
|
| 499 |
+
aggressive_cleanup(verbose=True)
|
| 500 |
+
|
| 501 |
# Fim da época
|
| 502 |
epoch_avg = sum(epoch_losses) / max(1, len(epoch_losses))
|
| 503 |
epoch_ppl = math.exp(min(20, epoch_avg)) if epoch_avg < 20 else float("inf")
|
|
|
|
| 506 |
f"({time.time()-epoch_start:.1f}s)"
|
| 507 |
)
|
| 508 |
|
| 509 |
+
# SALVAR MODELO APÓS CADA ÉPOCA (garante checkpoint mesmo se OOM no fim)
|
| 510 |
+
try:
|
| 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,
|
| 518 |
+
"final_perplexity": epoch_ppl,
|
| 519 |
+
"elapsed_s": time.time() - t_start,
|
| 520 |
+
"params_total": params["total"],
|
| 521 |
+
"params_M": params["total_M"],
|
| 522 |
+
}
|
| 523 |
+
self._save_final_model(partial_result)
|
| 524 |
+
logger.info(f" Modelo salvo após epoch {epoch+1}")
|
| 525 |
+
except Exception as e:
|
| 526 |
+
logger.warning(f" Falha ao salvar após epoch {epoch+1}: {e}")
|
| 527 |
+
|
| 528 |
# Fim do treino
|
| 529 |
elapsed_total = time.time() - t_start
|
| 530 |
final_loss = self.losses_history[-1] if self.losses_history else float("nan")
|
|
|
|
| 539 |
"epochs_completed": epoch + 1 if not self.killed_reason else epoch,
|
| 540 |
"killed": self.killed_reason is not None,
|
| 541 |
"kill_reason": self.killed_reason,
|
| 542 |
+
"global_reason": self.killed_reason,
|
| 543 |
"global_step": self.global_step,
|
| 544 |
"micro_batches": self.micro_batch_idx,
|
| 545 |
"best_loss": self.best_loss,
|
|
|
|
| 550 |
"params_total": params["total"],
|
| 551 |
"params_M": params["total_M"],
|
| 552 |
"monitor_summary": ks_summary,
|
| 553 |
+
"optimizer": cfg.optimizer_type,
|
| 554 |
+
"time_budget": self.time_budget.summary(),
|
| 555 |
+
"step_timer": self.step_timer.summary(),
|
| 556 |
}
|
| 557 |
|
| 558 |
logger.info("=" * 70)
|
src/bigru_t/utils/memory_cleanup.py
ADDED
|
@@ -0,0 +1,239 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""memory_cleanup.py — Limpeza agressiva de memória estilo Xavante.
|
| 2 |
+
|
| 3 |
+
Reaproveita os padrões de:
|
| 4 |
+
- xavante_work/flexnet/advanced_memory_cleanup.py (AdvancedMemoryCleaner)
|
| 5 |
+
- xavante_work/flexnet/oom_guard.py (OomGuard daemon)
|
| 6 |
+
- xavante_work/xavante/utils/timing.py (TimeBudget)
|
| 7 |
+
|
| 8 |
+
Implementa:
|
| 9 |
+
1. production_cleanup() — context manager com cleanup garantido
|
| 10 |
+
2. aggressive_cleanup() — gc.collect 3 gerações + torch.cuda.empty_cache
|
| 11 |
+
3. TimeBudget — orçamento de tempo por treino/época (streaming com timed steps)
|
| 12 |
+
4. get_rss_mb() — RSS do processo em MB
|
| 13 |
+
"""
|
| 14 |
+
from __future__ import annotations
|
| 15 |
+
|
| 16 |
+
import gc
|
| 17 |
+
import logging
|
| 18 |
+
import os
|
| 19 |
+
import threading
|
| 20 |
+
import time
|
| 21 |
+
from contextlib import contextmanager
|
| 22 |
+
from dataclasses import dataclass, field
|
| 23 |
+
from typing import Optional
|
| 24 |
+
|
| 25 |
+
logger = logging.getLogger(__name__)
|
| 26 |
+
|
| 27 |
+
|
| 28 |
+
def get_rss_mb() -> float:
|
| 29 |
+
"""Retorna o RSS (Resident Set Size) do processo atual em MB.
|
| 30 |
+
|
| 31 |
+
Lê /proc/self/status (Linux). Fallback para psutil se disponível.
|
| 32 |
+
"""
|
| 33 |
+
try:
|
| 34 |
+
with open("/proc/self/status", "r") as f:
|
| 35 |
+
for line in f:
|
| 36 |
+
if line.startswith("VmRSS:"):
|
| 37 |
+
# VmRSS: 12345 kB
|
| 38 |
+
parts = line.split()
|
| 39 |
+
return float(parts[1]) / 1024.0
|
| 40 |
+
except (FileNotFoundError, IndexError, ValueError):
|
| 41 |
+
pass
|
| 42 |
+
# Fallback psutil
|
| 43 |
+
try:
|
| 44 |
+
import psutil
|
| 45 |
+
return psutil.Process(os.getpid()).memory_info().rss / (1024 * 1024)
|
| 46 |
+
except ImportError:
|
| 47 |
+
return 0.0
|
| 48 |
+
|
| 49 |
+
|
| 50 |
+
def aggressive_cleanup(verbose: bool = False) -> dict:
|
| 51 |
+
"""Limpeza agressiva de memória (estilo Xavante).
|
| 52 |
+
|
| 53 |
+
Sequência:
|
| 54 |
+
1. gc.collect gen 0, 1, 2 (3 gerações completas)
|
| 55 |
+
2. torch.cuda.empty_cache() (se CUDA disponível)
|
| 56 |
+
3. torch.cuda.synchronize() (se CUDA disponível)
|
| 57 |
+
|
| 58 |
+
Args:
|
| 59 |
+
verbose: logar memória antes/depois
|
| 60 |
+
|
| 61 |
+
Returns:
|
| 62 |
+
dict com rss_before_mb, rss_after_mb, freed_mb
|
| 63 |
+
"""
|
| 64 |
+
rss_before = get_rss_mb()
|
| 65 |
+
|
| 66 |
+
# 3 gerações de gc
|
| 67 |
+
gc.collect(0)
|
| 68 |
+
gc.collect(1)
|
| 69 |
+
gc.collect(2)
|
| 70 |
+
|
| 71 |
+
# CUDA cleanup (no-op se CPU-only)
|
| 72 |
+
try:
|
| 73 |
+
import torch
|
| 74 |
+
if torch.cuda.is_available():
|
| 75 |
+
torch.cuda.empty_cache()
|
| 76 |
+
torch.cuda.synchronize()
|
| 77 |
+
except Exception:
|
| 78 |
+
pass
|
| 79 |
+
|
| 80 |
+
rss_after = get_rss_mb()
|
| 81 |
+
freed = rss_before - rss_after
|
| 82 |
+
|
| 83 |
+
if verbose:
|
| 84 |
+
logger.info(
|
| 85 |
+
"aggressive_cleanup: RSS %.1f → %.1f MB (freed %.1f MB)",
|
| 86 |
+
rss_before, rss_after, freed,
|
| 87 |
+
)
|
| 88 |
+
return {
|
| 89 |
+
"rss_before_mb": rss_before,
|
| 90 |
+
"rss_after_mb": rss_after,
|
| 91 |
+
"freed_mb": freed,
|
| 92 |
+
}
|
| 93 |
+
|
| 94 |
+
|
| 95 |
+
@contextmanager
|
| 96 |
+
def production_cleanup(verbose: bool = False):
|
| 97 |
+
"""Context manager que garante cleanup agressivo ao sair (mesmo com exceção).
|
| 98 |
+
|
| 99 |
+
Uso:
|
| 100 |
+
with production_cleanup(verbose=True):
|
| 101 |
+
# treino pesado aqui
|
| 102 |
+
...
|
| 103 |
+
# cleanup automático ao sair do bloco
|
| 104 |
+
"""
|
| 105 |
+
try:
|
| 106 |
+
yield
|
| 107 |
+
finally:
|
| 108 |
+
aggressive_cleanup(verbose=verbose)
|
| 109 |
+
|
| 110 |
+
|
| 111 |
+
# ---------------------------------------------------------------------------
|
| 112 |
+
# TimeBudget — orçamento de tempo para treino/época/passos
|
| 113 |
+
# ---------------------------------------------------------------------------
|
| 114 |
+
@dataclass
|
| 115 |
+
class TimeBudget:
|
| 116 |
+
"""Orçamento de tempo para treino com timed steps.
|
| 117 |
+
|
| 118 |
+
Permite:
|
| 119 |
+
- max_total_s: tempo máximo total de treino
|
| 120 |
+
- max_per_epoch_s: tempo máximo por época
|
| 121 |
+
- is_over() / is_epoch_over(): verifica se orçamento estourou
|
| 122 |
+
- remaining() / remaining_epoch(): tempo restante
|
| 123 |
+
- reset_epoch(): reseta o contador de época (chamar no início de cada época)
|
| 124 |
+
|
| 125 |
+
Uso:
|
| 126 |
+
budget = TimeBudget(max_total_s=600, max_per_epoch_s=300)
|
| 127 |
+
budget.reset_epoch()
|
| 128 |
+
for epoch in range(2):
|
| 129 |
+
for step in train_loop:
|
| 130 |
+
if budget.is_over() or budget.is_epoch_over():
|
| 131 |
+
break
|
| 132 |
+
...
|
| 133 |
+
budget.reset_epoch()
|
| 134 |
+
"""
|
| 135 |
+
max_total_s: float = 600.0
|
| 136 |
+
max_per_epoch_s: float = 300.0
|
| 137 |
+
_start: float = field(default_factory=time.time, repr=False)
|
| 138 |
+
_epoch_start: float = field(default_factory=time.time, repr=False)
|
| 139 |
+
|
| 140 |
+
def reset_epoch(self):
|
| 141 |
+
"""Reseta o contador de época (chamar no início de cada época)."""
|
| 142 |
+
self._epoch_start = time.time()
|
| 143 |
+
|
| 144 |
+
def reset_total(self):
|
| 145 |
+
"""Reseta o contador total (chamar no início do treino)."""
|
| 146 |
+
self._start = time.time()
|
| 147 |
+
self._epoch_start = time.time()
|
| 148 |
+
|
| 149 |
+
def elapsed(self) -> float:
|
| 150 |
+
return time.time() - self._start
|
| 151 |
+
|
| 152 |
+
def elapsed_epoch(self) -> float:
|
| 153 |
+
return time.time() - self._epoch_start
|
| 154 |
+
|
| 155 |
+
def remaining(self) -> float:
|
| 156 |
+
return max(0.0, self.max_total_s - self.elapsed())
|
| 157 |
+
|
| 158 |
+
def remaining_epoch(self) -> float:
|
| 159 |
+
return max(0.0, self.max_per_epoch_s - self.elapsed_epoch())
|
| 160 |
+
|
| 161 |
+
def is_over(self) -> bool:
|
| 162 |
+
return self.elapsed() >= self.max_total_s
|
| 163 |
+
|
| 164 |
+
def is_epoch_over(self) -> bool:
|
| 165 |
+
return self.elapsed_epoch() >= self.max_per_epoch_s
|
| 166 |
+
|
| 167 |
+
def should_save_partial(self) -> bool:
|
| 168 |
+
"""True se faltam < 10% do tempo (salvar parcial)."""
|
| 169 |
+
return self.remaining() < (self.max_total_s * 0.1)
|
| 170 |
+
|
| 171 |
+
def summary(self) -> dict:
|
| 172 |
+
return {
|
| 173 |
+
"elapsed_s": self.elapsed(),
|
| 174 |
+
"remaining_s": self.remaining(),
|
| 175 |
+
"epoch_elapsed_s": self.elapsed_epoch(),
|
| 176 |
+
"epoch_remaining_s": self.remaining_epoch(),
|
| 177 |
+
"is_over": self.is_over(),
|
| 178 |
+
"is_epoch_over": self.is_epoch_over(),
|
| 179 |
+
}
|
| 180 |
+
|
| 181 |
+
|
| 182 |
+
# ---------------------------------------------------------------------------
|
| 183 |
+
# StepTimer — mede tempo por passo com warning se lento
|
| 184 |
+
# ---------------------------------------------------------------------------
|
| 185 |
+
@dataclass
|
| 186 |
+
class StepTimer:
|
| 187 |
+
"""Mede tempo por passo e emite warnings se lento.
|
| 188 |
+
|
| 189 |
+
Uso:
|
| 190 |
+
timer = StepTimer(expected_s=0.5)
|
| 191 |
+
for step in range(N):
|
| 192 |
+
timer.start()
|
| 193 |
+
# ... passo de treino ...
|
| 194 |
+
timer.stop() # loga se > 2x expected
|
| 195 |
+
"""
|
| 196 |
+
expected_s: float = 0.5
|
| 197 |
+
_start: float = 0.0
|
| 198 |
+
_count: int = 0
|
| 199 |
+
_total_s: float = 0.0
|
| 200 |
+
_max_s: float = 0.0
|
| 201 |
+
|
| 202 |
+
def start(self):
|
| 203 |
+
self._start = time.time()
|
| 204 |
+
|
| 205 |
+
def stop(self, step_label: str = "") -> float:
|
| 206 |
+
elapsed = time.time() - self._start
|
| 207 |
+
self._count += 1
|
| 208 |
+
self._total_s += elapsed
|
| 209 |
+
if elapsed > self._max_s:
|
| 210 |
+
self._max_s = elapsed
|
| 211 |
+
if elapsed > 2 * self.expected_s:
|
| 212 |
+
logger.warning(
|
| 213 |
+
"LENTO: step %s took %.2fs (expected ~%.2fs)",
|
| 214 |
+
step_label or self._count, elapsed, self.expected_s,
|
| 215 |
+
)
|
| 216 |
+
elif elapsed < 0.5 * self.expected_s:
|
| 217 |
+
logger.debug("RAPIDO: step %s took %.2fs", step_label, elapsed)
|
| 218 |
+
return elapsed
|
| 219 |
+
|
| 220 |
+
def avg(self) -> float:
|
| 221 |
+
return self._total_s / max(1, self._count)
|
| 222 |
+
|
| 223 |
+
def summary(self) -> dict:
|
| 224 |
+
return {
|
| 225 |
+
"count": self._count,
|
| 226 |
+
"total_s": self._total_s,
|
| 227 |
+
"avg_s": self.avg(),
|
| 228 |
+
"max_s": self._max_s,
|
| 229 |
+
"expected_s": self.expected_s,
|
| 230 |
+
}
|
| 231 |
+
|
| 232 |
+
|
| 233 |
+
__all__ = [
|
| 234 |
+
"get_rss_mb",
|
| 235 |
+
"aggressive_cleanup",
|
| 236 |
+
"production_cleanup",
|
| 237 |
+
"TimeBudget",
|
| 238 |
+
"StepTimer",
|
| 239 |
+
]
|
training_report.json
CHANGED
|
@@ -1,40 +1,76 @@
|
|
| 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 |
"max_seq_len": 32,
|
|
|
|
|
|
|
| 11 |
"use_hypothesis": true,
|
| 12 |
"stop_grad_hyp": true,
|
| 13 |
"meta_interval": 5,
|
| 14 |
-
"
|
| 15 |
-
"
|
| 16 |
},
|
| 17 |
"results": {
|
| 18 |
"epochs_completed": 2,
|
| 19 |
-
"
|
| 20 |
-
"
|
| 21 |
-
"
|
| 22 |
-
"
|
| 23 |
-
"
|
| 24 |
-
"
|
| 25 |
-
"
|
| 26 |
-
"
|
| 27 |
-
|
| 28 |
-
|
| 29 |
-
|
| 30 |
-
|
| 31 |
-
|
| 32 |
-
|
| 33 |
-
|
| 34 |
-
|
| 35 |
-
|
| 36 |
-
|
| 37 |
-
|
| 38 |
-
|
| 39 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 40 |
}
|
|
|
|
| 1 |
{
|
| 2 |
+
"version": "BiGRU_T_version_V2",
|
| 3 |
"training_params": {
|
| 4 |
"epochs": 2,
|
| 5 |
"per_device_batch_size": 1,
|
| 6 |
"grad_accum": 2,
|
| 7 |
"lr": 0.001,
|
| 8 |
+
"optimizer": "HamiltonianWassersteinOptimizer",
|
| 9 |
+
"hw_lr_amp": 0.3,
|
| 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 |
+
"time_budget": "1800s total, 900s/epoch",
|
| 19 |
+
"cleanup_every_n_steps": 5
|
| 20 |
},
|
| 21 |
"results": {
|
| 22 |
"epochs_completed": 2,
|
| 23 |
+
"global_step": 6,
|
| 24 |
+
"best_loss": 5.69,
|
| 25 |
+
"final_avg_loss_epoch1": 9.75,
|
| 26 |
+
"final_loss_epoch2_step5": 5.69,
|
| 27 |
+
"final_perplexity": 295.51,
|
| 28 |
+
"params_total": 12967688,
|
| 29 |
+
"params_M": 12.97,
|
| 30 |
+
"loss_progression": [
|
| 31 |
+
9.59,
|
| 32 |
+
9.82,
|
| 33 |
+
9.86,
|
| 34 |
+
7.61,
|
| 35 |
+
5.69,
|
| 36 |
+
6.75
|
| 37 |
+
],
|
| 38 |
+
"ppl_progression": [
|
| 39 |
+
14567,
|
| 40 |
+
18394,
|
| 41 |
+
19165,
|
| 42 |
+
2023,
|
| 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 |
+
"multimodal_test": "5/5 encoders OK (text, image, audio, video, router+fusion)"
|
| 53 |
+
},
|
| 54 |
+
"bugs_found_and_fixed": [
|
| 55 |
+
"TextEncoder: 3 missing deps (attention_multimodal, embedding_reconfig, gru_hierarchy) — COPIED from source HF",
|
| 56 |
+
"AudioEncoder: in_channels param unused by Conv1d (uses n_mels=80) — documented",
|
| 57 |
+
"VideoEncoder: frames=[B,T,C,H,W] not [B,C,T,H,W] — documented in test",
|
| 58 |
+
"ModalRouter: expects Dict[str,Tensor] not single tensor — documented",
|
| 59 |
+
"FusionLayer: expects list of embeddings not single tensor — documented",
|
| 60 |
+
"MetaConfigurator: create_graph=True caused OOM-Killer — FIXED to create_graph=False (first-order)",
|
| 61 |
+
"Trainer: 3 forward passes per micro-batch — OPTIMIZED to 1 forward (return_aux)",
|
| 62 |
+
"HW optimizer: ref_buffer_size=10 caused 520MB overhead — REDUCED to 2",
|
| 63 |
+
"Smoke test: torch.no_grad() before backward — FIXED",
|
| 64 |
+
"BBPE tokenizer: BPE merges format was string, should be tuple — FIXED"
|
| 65 |
+
],
|
| 66 |
+
"new_modules_added": [
|
| 67 |
+
"src/bigru_t/tokenizer/bbpe_tokenizer.py — REFACTORED with parallel Map-Reduce",
|
| 68 |
+
"src/bigru_t/reasoning/circular_reasoning_wasserstein.py — COPIED from source",
|
| 69 |
+
"src/bigru_t/training/dpo.py — NEW standalone DPO_loss",
|
| 70 |
+
"src/bigru_t/inference/generator.py — NEW BiGRUTGenerator",
|
| 71 |
+
"src/bigru_t/utils/memory_cleanup.py — NEW aggressive_cleanup + TimeBudget",
|
| 72 |
+
"src/bigru_t/model/attention_multimodal.py — COPIED from source",
|
| 73 |
+
"src/bigru_t/model/embedding_reconfig.py — COPIED from source",
|
| 74 |
+
"src/bigru_t/model/gru_hierarchy.py — COPIED from source"
|
| 75 |
+
]
|
| 76 |
}
|