Download cnn_bigru/training/auto_learner.py from PowerMachine/CNN-BiGRU: direct link, hf CLI and curl.
- Browser
- Download file 11.5 kB
-
https://huggingface.co/PowerMachine/CNN-BiGRU/resolve/main/cnn_bigru/training/auto_learner.py
- Command line
-
hf download hf://PowerMachine/CNN-BiGRU/cnn_bigru/training/auto_learner.py
-
curl -L -o auto_learner.py https://huggingface.co/PowerMachine/CNN-BiGRU/resolve/main/cnn_bigru/training/auto_learner.py
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__) | |
| 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", | |
| ] | |