kaizen-42m / ewc.py
qoa's picture
Add KAIZEN inference code, benchmarks, semantic head, example memory, README, requirements
4700286 verified
Raw
History Blame Contribute Delete
3.13 kB
"""
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()