| """ |
| 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() |
|
|