""" ewc.py — Elastic Weight Consolidation for KAIZEN forgetting benchmark. Kirkpatrick et al. 2017: "Overcoming catastrophic forgetting in neural networks." L_total = L_task + λ * Σ_t Σ_i F_{t,i} * (θ_i - θ*_{t,i})² Applied to LoRAAdapter parameters only. Base model is frozen throughout. Fisher Information: empirical Fisher (mean squared gradient) computed on domain j training data immediately after convergence. """ import torch from lora import LoRAAdapter, KaizenWithLoRA from online_learner import build_update_seq class EWC: """Diagonal Laplace EWC regularizer for LoRAAdapter. Usage: ewc = EWC(lambda_ewc=1000.0) # After training domain j: ewc.update(model, adapter, domain_j_tasks, tokenizer) # During domain j+1 training, add ewc.penalty(adapter) to task loss. """ def __init__(self, lambda_ewc: float = 1000.0): self.lambda_ewc = lambda_ewc self._fishers: list[dict[str, torch.Tensor]] = [] self._optimal: list[dict[str, torch.Tensor]] = [] def update(self, model: KaizenWithLoRA, adapter: LoRAAdapter, tasks: list, tokenizer) -> float: """Compute empirical Fisher + save parameter anchor for domain tasks. Call AFTER adapter has converged on domain j, BEFORE training domain j+1. Returns mean Fisher L2 norm (diagnostic — useful for tuning lambda_ewc). """ fisher = {n: torch.zeros_like(p) for n, p in adapter.named_parameters()} n_seen = 0 for p in model.parameters(): p.requires_grad_(False) for p in adapter.parameters(): p.requires_grad_(True) for question, answer, _ in tasks: x, y = build_update_seq(tokenizer, question, answer) adapter.zero_grad() _, loss = model(x, targets=y, adapter=adapter) loss.backward() for n, p in adapter.named_parameters(): if p.grad is not None: fisher[n] = fisher[n] + p.grad.data.pow(2) n_seen += 1 scale = 1.0 / max(n_seen, 1) anchored_fisher = {n: (f * scale).clone() for n, f in fisher.items()} self._fishers.append(anchored_fisher) self._optimal.append({n: p.data.clone() for n, p in adapter.named_parameters()}) mean_fisher_norm = float( torch.stack([f.norm() for f in anchored_fisher.values()]).mean().item() ) return mean_fisher_norm def penalty(self, adapter: LoRAAdapter) -> torch.Tensor: """EWC regularization term: λ * Σ_t Σ_i F_{t,i} * (θ_i - θ*_{t,i})² Differentiable w.r.t. adapter.parameters(). Returns scalar tensor. Zero (no grad cost) when no anchor stored yet (first domain). """ if not self._fishers: return torch.tensor(0.0) penalties = [] for fisher_t, opt_t in zip(self._fishers, self._optimal): for n, p in adapter.named_parameters(): penalties.append((fisher_t[n] * (p - opt_t[n]).pow(2)).sum()) return self.lambda_ewc * torch.stack(penalties).sum()