""" 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 # ============================================================================ @dataclass 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 # ---------------------------------------------------------------------- @torch.no_grad() 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"]