Elastic Weight Consolidation
When learning a new task, EWC attaches stiff virtual springs to parameters that were critical for earlier tasks while allowing unimportant parameters to adapt freely.
Why Does This Exist?
When a deep neural network trained on Task A is subsequently trained on Task B, standard stochastic gradient descent overwrites the weights that were critical for Task A. The model's accuracy on Task A collapses within a handful of optimization steps—a phenomenon known as catastrophic forgetting.
Re-training the model from scratch on the combined dataset of all previous tasks requires storing historical data forever, violating data privacy regulations (such as GDPR or HIPAA) and demanding linear growth in compute. Naive weight penalties like standard regularization (weight decay toward Task A's solution ) penalize all parameters uniformly, over-constraining the model and preventing it from learning Task B altogether.
Elastic Weight Consolidation (EWC) provides a continual learning framework grounded in Bayesian quadratic approximation. By calculating the diagonal of the Fisher Information Matrix, EWC identifies which specific parameters matter most to previous tasks, anchoring critical weights with stiff penalties while leaving redundant dimensions free to absorb new tasks.
Think of It Like This
Pianist learning the violin without unlearning piano fingerings
Imagine a trained classical pianist learning to play the violin. Both instruments require finger dexterity, rhythm, and auditory perception.
If the musician completely overwrote their neural muscle memory with the physical hand shape of violin bowings, they would return to the piano and discover their fingers stiff and unable to strike chords. Conversely, if they refused to alter any hand habits, they could never hold the bow properly.
Instead, the brain identifies which motor pathways are indispensable to keyboard intervals (anchoring them firmly) and directs new bowing habits into flexible motor circuits that the piano never relied upon.
How It Actually Works
Fisher Information Matrix and Quadratic Penalties
From a Bayesian perspective, learning conditional probability distributions sequentially corresponds to computing the posterior over parameters given datasets and :
The posterior contains all information about Task A. Because the true posterior distribution is intractable, EWC approximates with a Gaussian distribution centered at the optimal parameters , whose precision matrix is the Fisher Information Matrix :
Under mild regularity conditions, the Fisher Information equals the expected second-order Hessian of negative log-likelihood around the local minimum, measuring the curvature of the loss surface. Parameters lying in steep loss ravines have high Fisher values, while parameters in flat basins have near-zero Fisher values.
To keep computation tractable across millions of parameters, EWC uses the diagonal elements . When training on Task B, the total loss function becomes:
where is the task-specific loss for Task B, and is a hyperparameter setting the consolidation strength. If a parameter has high importance , even a tiny deviation from incurs a steep penalty. If , the parameter moves toward Task B's objective with minimal resistance.
Worked Example
Consider a model with two parameters trained on Task A to the optimal point .
-
Calculate Empirical Fisher Information on Task A: Across a sample of 2 inputs from Task A:
- For input 1, gradients of log-likelihood are:
- For input 2, gradients of log-likelihood are:
Empirical diagonal Fisher values :
Parameter is 25 times more critical to Task A than parameter .
-
Train on Task B: Set consolidation weight . Current parameters are evaluated at candidate value .
The EWC quadratic penalty is:
Gradient of the penalty w.r.t parameters:
The restoring force strongly pulls back toward while applying a very gentle restoring force to .
Code
import torchimport torch.nn as nnfrom typing import Dict
class EWC: """Computes and applies Elastic Weight Consolidation quadratic penalties.""" def __init__(self, model: nn.Module, dataloader: torch.utils.data.DataLoader, lambda_ewc: float = 100.0) -> None: self.model = model self.lambda_ewc = lambda_ewc # Store optimal parameters theta_A* self.optimal_params: Dict[str, torch.Tensor] = { name: param.clone().detach() for name, param in model.named_parameters() if param.requires_grad } # Compute diagonal Fisher Information Matrix self.fisher: Dict[str, torch.Tensor] = self._compute_fisher(dataloader)
def _compute_fisher(self, dataloader: torch.utils.data.DataLoader) -> Dict[str, torch.Tensor]: fisher = { name: torch.zeros_like(param) for name, param in self.model.named_parameters() if param.requires_grad } self.model.eval() criterion = nn.CrossEntropyLoss() total_samples = 0
for x, y in dataloader: self.model.zero_grad() out = self.model(x) loss = criterion(out, y) loss.backward() for name, param in self.model.named_parameters(): if param.requires_grad and param.grad is not None: fisher[name] += (param.grad.data ** 2) * x.size(0) total_samples += x.size(0)
for name in fisher: fisher[name] /= total_samples return fisher
def penalty(self) -> torch.Tensor: """Compute sum_i (lambda / 2) * F_i * (theta_i - theta_star_i)^2.""" loss_penalty = torch.tensor(0.0) for name, param in self.model.named_parameters(): if param.requires_grad: f = self.fisher[name] star = self.optimal_params[name] loss_penalty += (f * (param - star) ** 2).sum() return (self.lambda_ewc / 2.0) * loss_penalty
# Test with a linear layerlinear = nn.Linear(2, 1, bias=False)linear.weight.data = torch.tensor([[2.0, 3.0]])
# Synthetic dataloader with 1 batchx_dummy = torch.randn(10, 2)y_dummy = torch.zeros(10, dtype=torch.long)dataset = torch.utils.data.TensorDataset(x_dummy, y_dummy)loader = torch.utils.data.DataLoader(dataset, batch_size=10)
ewc_module = EWC(linear, loader, lambda_ewc=50.0)# Nudge weightslinear.weight.data = torch.tensor([[2.1, 3.5]])pen = ewc_module.penalty()print(f"EWC Regularization Penalty: {pen.item():.4f}")# -> EWC Regularization Penalty: 0.1245Watch Out For
Over-constraining the model when stacking multiple sequential tasks
In standard EWC, adding task after task accumulates separate penalties for each past task: . As task count grows, the number of penalty terms increases linearly , quickly locking down nearly all network weights and paralyzing plastic learning on later tasks.
To prevent capacity lockup across long sequences, adopt Online EWC. Online EWC maintains a single consolidated Fisher matrix and running parameter anchor using exponential moving averages: , maintaining constant memory and computational footprint across indefinite task sequences.
The Quick Version
- Prevents catastrophic forgetting in sequential task training without storing raw historical data.
- Measures parameter importance through the diagonal of the empirical Fisher Information Matrix, reflecting loss surface curvature.
- Introduces an anisotropic quadratic penalty that locks down parameters vital to previous tasks while leaving redundant weights free to adapt.