Fisher Information Matrix Computation and Elastic Weight Consolidation (EWC) in LeanAgent

LeanAgent implements Elastic Weight Consolidation by estimating the diagonal Fisher Information Matrix through squared gradient accumulation during a dedicated forward-backward pass over training data, then applies these values as a quadratic penalty during fine-tuning to protect previously learned theorem-proving knowledge.

The lean-dojo/leanagent repository provides continual learning capabilities for neural theorem provers through Elastic Weight Consolidation (EWC). By computing the Fisher Information Matrix after initial training sessions, the system quantifies parameter importance and uses these estimates to regularize subsequent fine-tuning, preventing catastrophic forgetting when adapting to new mathematical domains.

Computing the Fisher Information Matrix

The Fisher Information Matrix computation occurs in three distinct phases: gradient accumulation across a static dataset, distributed synchronization across GPUs, and normalization before persistence.

Wrapping the Model with FisherComputationModule

The process begins in compute_fisher.py, where the current PremiseRetriever checkpoint is wrapped in a specialized Lightning module. At lines 50-52, the code instantiates FisherComputationModule(best_model), which intercepts the standard training loop to record gradients rather than update weights. This wrapper manages the lifecycle of Fisher estimation without modifying the underlying model architecture defined in retrieval/model.py.

Accumulating Squared Gradients Across Batches

During the estimation epoch, the module iterates over the retrieval training data using a PyTorch Lightning trainer configured at lines 57-66 of compute_fisher.py. For each batch, fisher_computation_module.py (lines 14-22) computes the contrastive retrieval loss and backpropagates through the network.

The critical accumulation logic resides at lines 62-68 of fisher_computation_module.py. Here, the code captures the square of each parameter's gradient and adds it to a running sum stored in self.fisher_info. This operation approximates the diagonal of the Fisher Information Matrix by computing $\mathbb{E}[(\nabla_\theta \log p(y|\theta))^2]$, which indicates how much each parameter contributes to the model's predictions on the training distribution.

Distributed Synchronization and Normalization

When training across multiple GPUs, individual processes hold partial Fisher estimates. After completing the epoch, fisher_computation_module.py executes dist.all_reduce (lines 84-86) to sum these partial matrices across the distributed world.

Normalization occurs at lines 94-98, where the accumulated values are divided by the total number of samples (dataset_size * world_size). This averaging converts the raw gradient squares into proper Fisher Information estimates, ensuring the magnitude remains stable regardless of dataset size or GPU count.

Persisting the Matrix to Disk

Once computation completes, the global rank-zero process calls save_fisher_info(fisher_file_path) at lines 94-97 of compute_fisher.py. This method serializes the Fisher Information Matrix as a Python pickle file mapping parameter names to their corresponding diagonal Fisher values, enabling retrieval during future training sessions on different datasets.

Applying the Fisher Information Matrix for Elastic Weight Consolidation

During subsequent fine-tuning phases, LeanAgent loads the persisted Fisher matrix and integrates it into the loss function as a regularization term.

Loading and Attaching the FIM

Before initiating training on new data, scripts utilize load_fisher_information defined in leanagent.py (lines 569-575) to deserialize the pickle file produced during the Fisher computation phase. The retrieved dictionary is then attached to the model via set_fisher_info(fisher_info) at lines 83-88 of retrieval/model.py, which stores the matrix in the PremiseRetriever instance for access during the training loop.

Capturing Baseline Parameters

Immediately after loading the Fisher matrix, the training script invokes set_previous_params() at lines 93-95 of retrieval/model.py. This method clones the current model weights (denoted as $\theta^0$) and stores them as reference values. These baseline parameters represent the optimal weights for previously learned tasks, serving as anchors that the EWC penalty will protect during gradient updates.

Integrating the EWC Penalty into Training

During each training step, the training_step method in model.py (lines 32-44) computes the task-specific loss and adds the EWC regularization term:

def training_step(self, batch, batch_idx):
    loss = self.compute_loss(batch)
    if self.ewc_enabled:
        loss += self.ewc_loss()
    return loss

The ewc_loss() method (lines 96-121) implements the core EWC formula, calculating $\frac{\lambda}{2} \sum_i F_i (\theta_i - \theta_i^0)^2$ across all parameters. Here, $F_i$ represents the Fisher Information value for parameter $i$, while $\lambda$ controls regularization strength. The hyperparameter is configurable through set_lambda(lambda_value) (lines 90-92), allowing researchers to balance plasticity and stability based on the divergence between old and new datasets.

Practical Implementation Examples

Computing and Saving the Fisher Information Matrix

The following pattern from compute_fisher.py demonstrates how to generate a Fisher matrix after completing initial training:

from retrieval.fisher_computation_module import FisherComputationModule
import pytorch_lightning as pl
from datetime import timedelta

# Load the trained checkpoint

model = PremiseRetriever.load(
    checkpoint_path, 
    device="cuda", 
    freeze=False, 
    config=cfg
)

# Wrap in Fisher computation module

fisher_module = FisherComputationModule(model)

# Configure distributed trainer

trainer = pl.Trainer(
    accelerator="gpu",
    devices=4,
    strategy=pl.strategies.DDPStrategy(timeout=timedelta(seconds=7*24*60*60)),
    max_epochs=1,
    precision="bf16-mixed",
)

# Accumulate squared gradients over training data

trainer.fit(fisher_module, datamodule=data_module)

# Persist on master process only

if trainer.is_global_zero:
    fisher_module.save_fisher_info("checkpoints/fisher_info.pkl")

Loading Fisher Information and Enabling EWC for Fine-Tuning

When adapting the model to a new mathematical corpus, load the previously computed Fisher matrix to activate EWC regularization:

from leanagent import load_fisher_information, find_latest_fisher
from retrieval.model import PremiseRetriever

# Initialize model for new task

model = PremiseRetriever.load(new_checkpoint_path)

# Load and attach Fisher Information Matrix

fisher_path = find_latest_fisher()
fisher_info = load_fisher_information(fisher_path)
model.set_fisher_info(fisher_info)

# Establish baseline parameters from current weights

model.set_previous_params()

# Set regularization strength (lambda)

model.set_lambda(0.4)

# Train with EWC penalty automatically applied

trainer.fit(model, datamodule=new_task_datamodule)

Summary

  • FisherComputationModule in retrieval/fisher_computation_module.py handles gradient accumulation and distributed synchronization to estimate the diagonal Fisher Information Matrix.
  • The matrix is persisted via save_fisher_info() in compute_fisher.py and later retrieved using load_fisher_information() from leanagent.py.
  • During fine-tuning, set_fisher_info() and set_previous_params() prepare the model, while ewc_loss() computes the quadratic penalty $\frac{\lambda}{2} \sum F (\theta - \theta^0)^2$.
  • The regularization strength is controlled through set_lambda(), allowing dynamic adjustment of the stability-plasticity trade-off.

Frequently Asked Questions

What does the Fisher Information Matrix represent in LeanAgent's EWC implementation?

The Fisher Information Matrix represents the importance of each neural network parameter for the tasks the model has already learned. In LeanAgent, only the diagonal elements are computed, measuring how sensitive the model's predictions are to changes in each individual weight. Parameters with high Fisher values are deemed critical for previous theorem-proving capabilities and are therefore penalized more heavily if they drift during fine-tuning on new data.

Why does the implementation use squared gradients rather than computing the full Fisher matrix?

Computing the full Fisher Information Matrix would require $O(n^2)$ memory for $n$ parameters, which is infeasible for large transformer models used in LeanAgent. The repository approximates the diagonal using squared gradients ($\mathbb{E}[g^2]$), which requires only $O(n)$ additional memory. This diagonal approximation, implemented in fisher_computation_module.py lines 62-68, captures sufficient information about parameter importance while remaining computationally tractable across multiple GPUs.

How should the lambda hyperparameter be tuned for optimal EWC performance?

The lambda parameter controls the strength of the EWC regularization relative to the new task's loss. According to the set_lambda() method in retrieval/model.py (lines 90-92), typical values range between 0.1 and 10.0 depending on the similarity between old and new datasets. Higher values (e.g., 0.4 or above) enforce stronger parameter protection suitable for highly divergent mathematical domains, while lower values permit more flexible adaptation when tasks are closely related.

Can Fisher Information computation be distributed across multiple GPUs?

Yes, LeanAgent explicitly supports distributed Fisher computation through PyTorch's distributed backend. The all_reduce operation at lines 84-86 of fisher_computation_module.py synchronizes squared gradients across all processes, while the normalization step at lines 94-98 divides by the global dataset size to produce consistent estimates regardless of the number of GPUs used during computation.

Have a question about this repo?

These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →