CNN-BiGRU / cnn_bigru /training /auto_learner.py
PowerMachine's picture
v3.0: reorganiza arquivos sob cnn_bigru/ (preserva árvore de pastas)
49b8205 verified
Raw History Blame Contribute Delete
11.5 kB
"""auto_learner.py — Auto-aprendizado: ajuste dinâmico de LR e controle de gradiente.
Implementa as seções 9-10 de dados.txt:
9. AJUSTE DINÂMICO DAS TAXAS DE APRENDIZADO:
fator_G = exp(-κ * ||grad_G||) * (1 + v_media) / 2
novo_eta_G = eta_G * fator_G
fator_V = exp(-κ * ||grad_V||) * (1 + L_V_media)^{-1}
novo_eta_V = eta_V * fator_V
10. CONTROLE DE EXPLOSÃO DE GRADIENTE E ESPECTRAL NORM:
- clip_grad_norm_(params, GRAD_CLIP)
- Spectral normalization (power iteration) para Conv1D e Linear
"""
from __future__ import annotations
import logging
import math
from dataclasses import dataclass
from typing import Dict, List, Optional, Tuple
import torch
import torch.nn as nn
logger = logging.getLogger(__name__)
@dataclass
class AutoLearnConfig:
"""Configuração do auto-aprendizado."""
kappa_curv: float = 0.01 # κ — fator de ajuste de η com norma do gradiente
grad_clip: float = 1.0 # τ — clipping de gradiente (norma L2)
spectral_radius: float = 1.0 # limite para norma espectral
lr_min: float = 1e-6 # LR mínimo
lr_max: float = 1e-2 # LR máximo
initial_lr_G: float = 1e-3 # LR inicial do gerador
initial_lr_V: float = 1e-3 # LR inicial do verificador
use_spectral_norm: bool = True
apply_after_step: bool = True
# NOVO: peso L2 usado pelo otimizador AdamW (corresponde a L2_REG do dados.txt)
l2_reg: float = 1e-5
# NOVO: potências de iteração para spectral norm (1 era muito pouco)
spectral_n_iters: int = 1
# NOVO: device padrão para buffers internos
device: str = "cpu"
class AutoLearner:
"""Controla ajuste dinâmico de LR, clipping e spectral norm.
Não é um nn.Module — é um controlador que atua sobre modelos externos.
"""
def __init__(self, config: AutoLearnConfig):
self.cfg = config
self.current_lr_G = config.initial_lr_G
self.current_lr_V = config.initial_lr_V
self.history: List[Dict] = []
def compute_total_grad_norm(self, model: nn.Module) -> float:
"""Calcula a norma L2 total dos gradientes do modelo."""
total = 0.0
for p in model.parameters():
if p.grad is not None:
total += p.grad.data.norm(2).item() ** 2
return math.sqrt(total)
def update_learning_rates(
self,
grad_G_norm: float,
grad_V_norm: float,
v_mean: float,
L_V_mean: float,
) -> Tuple[float, float]:
"""Atualiza dinamicamente as taxas de aprendizado.
Args:
grad_G_norm: norma L2 do gradiente do gerador
grad_V_norm: norma L2 do gradiente do verificador
v_mean: média das probabilidades do verificador
L_V_mean: perda média do verificador
Returns:
(novo_eta_G, novo_eta_V)
"""
# Fator do gerador: exp(-κ * ||grad||) * (1 + v_media) / 2
fator_G = math.exp(-self.cfg.kappa_curv * grad_G_norm) * ((1.0 + v_mean) / 2.0)
novo_eta_G = self.current_lr_G * fator_G
# Fator do verificador: exp(-κ * ||grad||) * (1 + L_V)^{-1}
# Evita divisão por zero
fator_V = math.exp(-self.cfg.kappa_curv * grad_V_norm) / max(1.0 + L_V_mean, 1e-6)
novo_eta_V = self.current_lr_V * fator_V
# Clamp para evitar valores absurdos
novo_eta_G = max(self.cfg.lr_min, min(self.cfg.lr_max, novo_eta_G))
novo_eta_V = max(self.cfg.lr_min, min(self.cfg.lr_max, novo_eta_V))
self.current_lr_G = novo_eta_G
self.current_lr_V = novo_eta_V
self.history.append({
"lr_G": novo_eta_G,
"lr_V": novo_eta_V,
"grad_G_norm": grad_G_norm,
"grad_V_norm": grad_V_norm,
"v_mean": v_mean,
"L_V_mean": L_V_mean,
})
return novo_eta_G, novo_eta_V
def apply_clipping(self, *models: nn.Module) -> None:
"""Aplica gradient clipping global nos modelos."""
for model in models:
if model is not None:
torch.nn.utils.clip_grad_norm_(model.parameters(), self.cfg.grad_clip)
def apply_spectral_norm(self, model: nn.Module) -> int:
"""Aplica normalização espectral via power iteration em Conv1d/Linear.
Após cada step, força ||W||_2 <= spectral_radius.
CORREÇÃO: O bug original armazenava `module._spec_u` como atributo
comum (não buffer), que NÃO migra com `.to(device)`. Agora usamos
`register_buffer` na primeira chamada para garantir movimentação
corta com o device do modelo.
"""
if not self.cfg.use_spectral_norm:
return 0
n_normalized = 0
for module in model.modules():
if isinstance(module, (nn.Linear, nn.Conv1d, nn.Conv2d)):
if module.weight is not None and module.weight.dim() >= 2:
try:
with torch.no_grad():
# Power iteration (n_iters iterações)
W = module.weight
# Reshape para 2D se necessário
W_flat = W.reshape(W.size(0), -1)
# u, v aleatórios (ou reutiliza se já existe)
buffer_name = "_spec_u"
if not hasattr(module, buffer_name):
u = torch.randn(W_flat.size(0), 1, device=W.device, dtype=W.dtype)
u = u / u.norm().clamp(min=1e-8)
# Registrar como buffer (move com .to(device))
module.register_buffer(buffer_name, u, persistent=False)
else:
u = getattr(module, buffer_name)
# Garantir que está no device correto
if u.device != W.device:
u = u.to(W.device)
setattr(module, buffer_name, u)
# n_iters iterações de power iteration
for _ in range(self.cfg.spectral_n_iters):
v = torch.matmul(W_flat.T, u)
v = v / v.norm().clamp(min=1e-8)
u = torch.matmul(W_flat, v)
u = u / u.norm().clamp(min=1e-8)
# Atualizar buffer (in-place para manter no mesmo device)
u_new = u.detach()
getattr(module, buffer_name).copy_(u_new)
sigma = (u.T @ W_flat @ v).item()
if sigma > self.cfg.spectral_radius and sigma > 0:
scale = self.cfg.spectral_radius / sigma
module.weight.data.mul_(scale)
n_normalized += 1
except Exception as e:
logger.debug("Spectral norm falhou em %s: %s", type(module).__name__, e)
return n_normalized
def set_optimizer_lr(self, optimizer: torch.optim.Optimizer, lr: float) -> None:
"""Atualiza a taxa de aprendizado de um otimizador."""
for pg in optimizer.param_groups:
pg["lr"] = lr
def orthogonal_init_model(model: nn.Module) -> int:
"""Aplica inicialização ortogonal em todas as camadas Linear/Conv/GRUCell."""
n = 0
for module in model.modules():
if isinstance(module, (nn.Linear, nn.Conv1d, nn.Conv2d)):
try:
nn.init.orthogonal_(module.weight)
if module.bias is not None:
nn.init.zeros_(module.bias)
n += 1
except Exception:
pass
elif isinstance(module, (nn.GRUCell, nn.GRU, nn.LSTM)):
for name, p in module.named_parameters():
if "weight" in name and p.dim() >= 2:
try:
nn.init.orthogonal_(p)
n += 1
except Exception:
pass
elif "bias" in name:
nn.init.zeros_(p)
return n
def estimate_curvature(
model: nn.Module,
loss_fn,
*args,
eps: float = 1e-3,
n_samples: int = 5,
) -> torch.Tensor:
"""Estima curvatura da perda via diferenças finitas.
Para uma amostra de parâmetros, calcula:
curv ≈ (L(θ + ε) - 2L(θ) + L(θ - ε)) / ε²
Retorna a média das estimativas.
CORREÇÃO: O bug original mutava `flat[i]` in-place sem try/finally,
corrompendo o modelo se uma exceção ocorresse entre `+eps` e o restore.
Agora usamos try/finally para garantir restauração sempre.
"""
curvatures = []
params = [p for p in model.parameters() if p.requires_grad and p.dim() >= 2]
if not params:
# Device-safe zero
try:
device = next(model.parameters()).device
except StopIteration:
device = torch.device("cpu")
return torch.zeros((), device=device)
# Subamostra
n_samples = min(n_samples, len(params))
indices = torch.randperm(len(params))[:n_samples]
for idx in indices:
p = params[idx]
# Pega um elemento aleatório
flat = p.data.view(-1)
if flat.numel() == 0:
continue
i = torch.randint(0, flat.numel(), (1,)).item()
orig = flat[i].item()
# TRY/FINALLY para garantir restauração sempre
try:
with torch.no_grad():
# L(θ)
try:
loss_0 = loss_fn(*args)
if isinstance(loss_0, dict):
loss_0 = loss_0["total"]
loss_0 = float(loss_0)
except Exception as e:
logger.debug("curvature L(θ) falhou: %s", e)
continue
# L(θ + ε)
flat[i] = orig + eps
try:
loss_p = loss_fn(*args)
if isinstance(loss_p, dict):
loss_p = loss_p["total"]
loss_p = float(loss_p)
except Exception as e:
logger.debug("curvature L(θ+ε) falhou: %s", e)
continue
# L(θ - ε)
flat[i] = orig - eps
try:
loss_m = loss_fn(*args)
if isinstance(loss_m, dict):
loss_m = loss_m["total"]
loss_m = float(loss_m)
except Exception as e:
logger.debug("curvature L(θ-ε) falhou: %s", e)
continue
curv = (loss_p - 2 * loss_0 + loss_m) / (eps ** 2)
curvatures.append(curv)
finally:
# SEMPRE restaurar o valor original
flat[i] = orig
if not curvatures:
try:
device = next(model.parameters()).device
except StopIteration:
device = torch.device("cpu")
return torch.zeros((), device=device)
return torch.tensor(sum(curvatures) / len(curvatures))
__all__ = [
"AutoLearnConfig",
"AutoLearner",
"orthogonal_init_model",
"estimate_curvature",
]