Download cnn_bigru/utils/ewc.py from PowerMachine/CNN-BiGRU: direct link, hf CLI and curl.
- Browser
- Download file 15.2 kB
-
https://huggingface.co/PowerMachine/CNN-BiGRU/resolve/main/cnn_bigru/utils/ewc.py
- Command line
-
hf download hf://PowerMachine/CNN-BiGRU/cnn_bigru/utils/ewc.py
-
curl -L -o ewc.py https://huggingface.co/PowerMachine/CNN-BiGRU/resolve/main/cnn_bigru/utils/ewc.py
15.2 kB
| """ | |
| EWC (Elastic Weight Consolidation) — Aprendizado Contínuo sem Catastrophic Forgetting. | |
| Implementa: | |
| - Cálculo da diagonal da Matriz de Informação de Fisher (empirical Fisher) | |
| - Estado consolidado (theta_star + Fisher) para múltiplas tarefas | |
| - Online EWC (média exponencial de Fisher entre tarefas) | |
| - Penalidade diferenciável L_EWC = sum_i (lambda/2) * F_i * (theta_i - theta_star_i)^2 | |
| Referência matemática: docs/MATH_ANALYSIS.md seção 2. | |
| Autor: CNN-BiGRU Project | |
| """ | |
| from __future__ import annotations | |
| import logging | |
| from dataclasses import dataclass, field | |
| from typing import Dict, List, Optional, Tuple | |
| import torch | |
| import torch.nn as nn | |
| logger = logging.getLogger(__name__) | |
| # ============================================================================ | |
| # Configuração | |
| # ============================================================================ | |
| class EWCConfig: | |
| """Configuração do EWC (Elastic Weight Consolidation).""" | |
| enabled: bool = True | |
| # Peso da penalidade EWC na perda total | |
| lambda_ewc: float = 100.0 | |
| # Número de amostras usadas para estimar a diagonal de Fisher | |
| n_samples_fisher: int = 200 | |
| # Decaimento para Online EWC: F_agg = gamma*F_old + (1-gamma)*F_new | |
| # Se None, usa soma direta (EWC padrão — explode em memória após muitas tarefas) | |
| online_gamma: Optional[float] = 0.9 | |
| # Clip para estabilizar Fisher (evitar valores extremos) | |
| fisher_clip_max: float = 1e4 | |
| # Tipos de parâmetros a penalizar (outros como LayerNorm/bias podem ser excluídos) | |
| include_param_names: Tuple[str, ...] = ("weight",) | |
| exclude_param_names: Tuple[str, ...] = ("bias", "layernorm", "layer_norm", "bn", "batchnorm") | |
| # Device padrão | |
| device: str = "cpu" | |
| # ============================================================================ | |
| # Estado EWC | |
| # ============================================================================ | |
| class EWCState: | |
| """ | |
| Mantém o estado consolidado do EWC: lista de (theta_star, fisher) por tarefa. | |
| Para Online EWC, mantém uma única Fisher agregada e theta_star atual. | |
| """ | |
| def __init__(self, config: EWCConfig): | |
| self.config = config | |
| # Lista de tarefas: cada entrada é um dict {param_name: (theta_star, fisher)} | |
| self.tasks: List[Dict[str, Tuple[torch.Tensor, torch.Tensor]]] = [] | |
| # Online: theta_star e fisher agregados | |
| self.online_theta: Dict[str, torch.Tensor] = {} | |
| self.online_fisher: Dict[str, torch.Tensor] = {} | |
| # ---------------------------------------------------------------------- | |
| # Seleção de parâmetros | |
| # ---------------------------------------------------------------------- | |
| def _select_params(self, model: nn.Module) -> List[Tuple[str, torch.Tensor]]: | |
| """Seleciona parâmetros treináveis elegíveis para EWC.""" | |
| selected = [] | |
| for name, param in model.named_parameters(): | |
| if not param.requires_grad: | |
| continue | |
| name_lower = name.lower() | |
| # Excluir nomes na lista negra | |
| if any(ex in name_lower for ex in self.config.exclude_param_names): | |
| continue | |
| # Incluir apenas nomes na lista branca (se especificada) | |
| if self.config.include_param_names and not any( | |
| inc in name_lower for inc in self.config.include_param_names | |
| ): | |
| continue | |
| selected.append((name, param)) | |
| return selected | |
| # ---------------------------------------------------------------------- | |
| # Cálculo da Diagonal de Fisher | |
| # ---------------------------------------------------------------------- | |
| def compute_fisher( | |
| self, | |
| model: nn.Module, | |
| forward_fn, | |
| n_samples: Optional[int] = None, | |
| ) -> Dict[str, torch.Tensor]: | |
| """ | |
| Calcula a diagonal da matriz de Fisher Information (empirical Fisher). | |
| Args: | |
| model: modelo treinado (parâmetros congelados idealmente) | |
| forward_fn: callable que retorna (logits, target) ou logits. | |
| Deve aceitar nenhum argumento ou um índice. | |
| n_samples: número de amostras (default: config.n_samples_fisher) | |
| Returns: | |
| Dict {param_name: fisher_diag_tensor} (mesmo shape do parâmetro) | |
| Fórmula: | |
| F_i = (1/N) sum_n [grad_i log p(y_n* | x_n; theta*)]^2 | |
| onde y_n* = argmax_y p(y | x_n; theta*) | |
| """ | |
| if n_samples is None: | |
| n_samples = self.config.n_samples_fisher | |
| model.eval() | |
| selected = self._select_params(model) | |
| # Inicializar acumulador de grad^2 | |
| fisher: Dict[str, torch.Tensor] = { | |
| name: torch.zeros_like(param.detach(), device=param.device) | |
| for name, param in selected | |
| } | |
| count = 0 | |
| # Garantir que requires_grad está ativo para os parâmetros selecionados | |
| # durante todo o cálculo da Fisher (restaurado no finally) | |
| original_requires_grad: Dict[str, bool] = {} | |
| for name, param in selected: | |
| original_requires_grad[name] = param.requires_grad | |
| param.requires_grad_(True) | |
| try: | |
| for i in range(n_samples): | |
| try: | |
| # Resetar gradientes | |
| for _, param in selected: | |
| if param.grad is not None: | |
| param.grad.zero_() | |
| # Forward + Backward — TODO em contexto com grad habilitado | |
| # IMPORTANTE: usar torch.set_grad_enabled(True) para forçar gradientes | |
| # mesmo se o caller estiver em contexto no_grad. Deve cobrir TANTO | |
| # o forward QUANTO o backward (senão as operações pós-forward perdem grad_fn). | |
| prev_grad_mode = torch.is_grad_enabled() | |
| torch.set_grad_enabled(True) | |
| try: | |
| result = forward_fn(i) | |
| if isinstance(result, tuple) and len(result) == 2: | |
| logits, target = result | |
| else: | |
| logits = result | |
| # Usar y* = argmax logits (empirical Fisher) | |
| target = logits.argmax(dim=-1) | |
| # Garantir que logits é [B, V] ou [B, T, V] → flatten | |
| if logits.dim() == 3: | |
| B, T, V = logits.shape | |
| logits_flat = logits.reshape(-1, V) | |
| target_flat = target.reshape(-1) | |
| else: | |
| logits_flat = logits | |
| target_flat = target.reshape(-1) | |
| # Log-verossimilhança (CrossEntropy = -log p) | |
| log_prob = torch.nn.functional.log_softmax(logits_flat, dim=-1) | |
| # Pegar log_prob das classes target | |
| picked = log_prob.gather(1, target_flat.unsqueeze(-1)).sum() | |
| # O sinal negativo: queremos grad de log p (não -log p) | |
| loss_for_grad = -picked | |
| # Backward — AINDA em grad mode (não restauramos ainda) | |
| loss_for_grad.backward(retain_graph=False) | |
| finally: | |
| torch.set_grad_enabled(prev_grad_mode) | |
| # Acumular grad^2 | |
| for name, param in selected: | |
| if param.grad is not None: | |
| g2 = param.grad.detach() ** 2 | |
| fisher[name] += g2 | |
| count += 1 | |
| except Exception as e: | |
| logger.warning(f"EWC fisher sample {i} falhou: {e}") | |
| continue | |
| finally: | |
| # Restaurar requires_grad original | |
| for name, param in selected: | |
| if name in original_requires_grad: | |
| param.requires_grad_(original_requires_grad[name]) | |
| if count == 0: | |
| logger.error("EWC: nenhuma amostra válida para Fisher — usando zeros") | |
| return fisher | |
| # Média | |
| for name in fisher: | |
| fisher[name] = (fisher[name] / count).clamp(0, self.config.fisher_clip_max) | |
| return fisher | |
| # ---------------------------------------------------------------------- | |
| # Consolidar tarefa | |
| # ---------------------------------------------------------------------- | |
| def consolidate( | |
| self, | |
| model: nn.Module, | |
| fisher: Optional[Dict[str, torch.Tensor]] = None, | |
| forward_fn=None, | |
| ) -> None: | |
| """ | |
| Consolida o estado atual do modelo como uma nova tarefa. | |
| Args: | |
| model: modelo treinado | |
| fisher: Fisher diagonal pré-computada (se None, calcula via forward_fn) | |
| forward_fn: necessário se fisher=None | |
| """ | |
| if fisher is None: | |
| if forward_fn is None: | |
| raise ValueError("EWC.consolidate requer forward_fn quando fisher=None") | |
| fisher = self.compute_fisher(model, forward_fn) | |
| # Snapshot dos parâmetros atuais (theta_star) | |
| theta_star: Dict[str, torch.Tensor] = {} | |
| for name, param in model.named_parameters(): | |
| if name in fisher: | |
| theta_star[name] = param.detach().clone() | |
| task_state = { | |
| name: (theta_star[name], fisher[name]) for name in fisher | |
| } | |
| self.tasks.append(task_state) | |
| # Atualizar estado online | |
| if self.config.online_gamma is not None: | |
| gamma = self.config.online_gamma | |
| if not self.online_fisher: | |
| # Primeira tarefa | |
| for name in fisher: | |
| self.online_fisher[name] = fisher[name].clone() | |
| self.online_theta[name] = theta_star[name].clone() | |
| else: | |
| # Online EWC: F_agg = gamma * F_old + (1-gamma) * F_new | |
| # theta_agg = gamma * theta_old + (1-gamma) * theta_new | |
| for name in fisher: | |
| if name in self.online_fisher: | |
| self.online_fisher[name] = ( | |
| gamma * self.online_fisher[name] | |
| + (1 - gamma) * fisher[name] | |
| ).clamp(0, self.config.fisher_clip_max) | |
| self.online_theta[name] = ( | |
| gamma * self.online_theta[name] | |
| + (1 - gamma) * theta_star[name] | |
| ) | |
| else: | |
| self.online_fisher[name] = fisher[name].clone() | |
| self.online_theta[name] = theta_star[name].clone() | |
| else: | |
| # EWC padrão — somar Fishers de todas as tarefas | |
| if not self.online_fisher: | |
| for name in fisher: | |
| self.online_fisher[name] = fisher[name].clone() | |
| self.online_theta[name] = theta_star[name].clone() | |
| else: | |
| for name in fisher: | |
| self.online_fisher[name] = ( | |
| self.online_fisher[name] + fisher[name] | |
| ).clamp(0, self.config.fisher_clip_max) | |
| # Theta online atualiza para o mais recente | |
| self.online_theta[name] = theta_star[name].clone() | |
| logger.info( | |
| f"EWC consolidado: tarefa #{len(self.tasks)}, " | |
| f"parâmetros rastreados: {len(theta_star)}" | |
| ) | |
| # ---------------------------------------------------------------------- | |
| # Penalidade EWC | |
| # ---------------------------------------------------------------------- | |
| def penalty(self, model: nn.Module) -> torch.Tensor: | |
| """ | |
| Calcula a penalidade EWC: L_EWC = sum_i (lambda/2) * F_i * (theta_i - theta_star_i)^2 | |
| Usa o estado online (agregado) se disponível, caso contrário soma de | |
| todas as tarefas (EWC padrão). | |
| Returns: | |
| Escalar (tensor 0-dim) com a penalidade. | |
| """ | |
| if not self.config.enabled or not self.online_fisher: | |
| # Retorna zero no device do modelo | |
| try: | |
| device = next(model.parameters()).device | |
| except StopIteration: | |
| device = torch.device(self.config.device) | |
| return torch.zeros((), device=device) | |
| total = None | |
| for name, param in model.named_parameters(): | |
| if name not in self.online_fisher: | |
| continue | |
| F = self.online_fisher[name] | |
| theta_s = self.online_theta[name] | |
| # Penalidade quadrática | |
| diff = param - theta_s.to(param.device) | |
| pen = (F.to(param.device) * diff.pow(2)).sum() | |
| if total is None: | |
| total = pen | |
| else: | |
| total = total + pen | |
| if total is None: | |
| try: | |
| device = next(model.parameters()).device | |
| except StopIteration: | |
| device = torch.device(self.config.device) | |
| return torch.zeros((), device=device) | |
| return (self.config.lambda_ewc * 0.5) * total | |
| # ---------------------------------------------------------------------- | |
| # Utilitários | |
| # ---------------------------------------------------------------------- | |
| def num_tasks(self) -> int: | |
| return len(self.tasks) | |
| def is_active(self) -> bool: | |
| return self.config.enabled and bool(self.online_fisher) | |
| def state_dict(self) -> Dict: | |
| """Serializa o estado para checkpoint.""" | |
| return { | |
| "tasks": [ | |
| {name: (ts.cpu(), f.cpu()) for name, (ts, f) in task.items()} | |
| for task in self.tasks | |
| ], | |
| "online_theta": {k: v.cpu() for k, v in self.online_theta.items()}, | |
| "online_fisher": {k: v.cpu() for k, v in self.online_fisher.items()}, | |
| "config": self.config.__dict__, | |
| } | |
| def load_state_dict(self, sd: Dict) -> None: | |
| """Carrega o estado a partir de checkpoint.""" | |
| self.tasks = [ | |
| {name: (ts, f) for name, (ts, f) in task.items()} | |
| for task in sd.get("tasks", []) | |
| ] | |
| self.online_theta = sd.get("online_theta", {}) | |
| self.online_fisher = sd.get("online_fisher", {}) | |
| if "config" in sd: | |
| for k, v in sd["config"].items(): | |
| if hasattr(self.config, k): | |
| setattr(self.config, k, v) | |
| # ============================================================================ | |
| # API conveniente | |
| # ============================================================================ | |
| def apply_ewc_penalty(model: nn.Module, ewc_state: EWCState) -> torch.Tensor: | |
| """ | |
| Atalho: aplica a penalidade EWC ao modelo. | |
| Uso típico no trainer: | |
| loss = task_loss + ewc_penalty | |
| ewc_penalty = apply_ewc_penalty(model, ewc_state) | |
| """ | |
| return ewc_state.penalty(model) | |
| __all__ = ["EWCConfig", "EWCState", "apply_ewc_penalty"] | |