PowerMachine's picture
v3.0: reorganiza arquivos sob cnn_bigru/ (preserva árvore de pastas)
49b8205 verified
Raw History Blame Contribute Delete
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
# ============================================================================
@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"]