"""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", ]