How Progressive Training Prevents Catastrophic Forgetting in LeanAgent

LeanAgent prevents catastrophic forgetting during lifelong learning by combining curriculum-based progressive training with Elastic Weight Consolidation (EWC), which penalizes changes to critically important parameters using Fisher information matrices computed after each repository.

The lean-dojo/leanagent repository implements a lifelong learning system for automated theorem proving that trains sequentially on multiple Lean repositories without erasing previously acquired knowledge. By anchoring learned representations through progressive training with EWC regularization, the system continuously expands its mathematical reasoning capabilities while preserving foundational patterns from earlier tasks.

The Progressive Training Architecture

LeanAgent's lifelong learning workflow centers on curriculum-driven progressive training, where the model encounters repositories in ascending order of difficulty rather than training in isolation on single datasets.

Curriculum-Driven Repository Sequencing

In leanagent.py, the system activates progressive training through the run_progressive_training flag and processes repositories sequentially using sort_repositories_by_difficulty. This ensures easier theorems establish a solid parameter foundation before the model tackles complex mathematical domains. The training loop iterates through repositories one after another without resetting model weights, allowing knowledge to accumulate across the entire curriculum.

Parameter Preservation Across Tasks

Before training on each new repository, the system invokes PremiseRetriever.set_previous_params() in retrieval/model.py to store a snapshot of current weights in previous_params. This creates a baseline reference point that the regularization mechanism uses to detect and penalize deviations from previously learned configurations.

Elastic Weight Consolidation Implementation

The core defense against catastrophic forgetting is Elastic Weight Consolidation (EWC), a regularization technique that treats different model parameters differently based on their importance to prior tasks.

Capturing Parameter Importance with Fisher Information

After completing training on each repository, LeanAgent computes the Fisher information matrix to quantify how much each parameter contributed to the previous task's performance. The compute_fisher.py module and its helper retrieval/fisher_computation_module.py calculate these importance scores, which are then injected into the model via PremiseRetriever.set_fisher_info(). Parameters with high Fisher values are deemed critical for preserving past knowledge.

Computing the EWC Regularization Loss

During subsequent training steps, the total loss function combines standard task loss with the EWC penalty. As implemented in retrieval/model.py (lines 96-122), the PremiseRetriever.ewc_loss() method scales the squared deviation of each parameter from its saved value by the corresponding Fisher entry and the user-controlled coefficient lamda:

total_loss = task_loss + self.lamda * ewc_loss

The lamda parameter controls regularization strength—setting lambdas = [0.1] in leanagent.py activates EWC protection during progressive training. This mathematical formulation anchors important parameters close to their optimal values for previous tasks while allowing less critical weights to adapt to new repositories.

Code Walkthrough: Lifelong Learning in Practice

Enabling progressive training with catastrophic forgetting protection requires configuring the curriculum loop and EWC parameters in the main orchestration file:


# leanagent.py - Enable progressive training curriculum

run_progressive_training = True
if run_progressive_training:
    logger.info("Running progressive training")
    lambdas = [0.1]          # Non-zero lambda activates EWC regularization

# Sort repositories by difficulty for curriculum learning

sorted_repos = sort_repositories_by_difficulty(repositories)

Before each new repository begins, the model saves its current state:


# retrieval/model.py - Snapshot parameters before new task

model.set_previous_params()   # Stores copy of current weights as previous_params

After completing a repository, compute and store Fisher information:


# compute_fisher.py - Calculate parameter importance

fisher_matrix = compute_fisher(model, dataloader)
model.set_fisher_info(fisher_matrix)  # Feed importance scores to retriever

The training step automatically applies EWC regularization without manual intervention:


# retrieval/model.py - Training step with EWC (lines 96-122)

def training_step(self, batch, batch_idx):
    loss = self.task_loss(batch)        # Standard contrastive or classification loss

    loss += self.ewc_loss()              # Add Fisher-weighted parameter deviation penalty

    return loss

Summary

  • Progressive training in LeanAgent processes multiple Lean repositories sequentially using curriculum learning, with run_progressive_training enabling the lifelong learning loop.
  • Parameter snapshots via set_previous_params() establish baseline weights before each new repository training begins.
  • Fisher information matrices computed in compute_fisher.py quantify parameter importance for preserving knowledge from completed tasks.
  • EWC regularization in retrieval/model.py penalizes changes to critical parameters using the formula task_loss + lamda * ewc_loss, preventing catastrophic forgetting.
  • Difficulty-based sorting ensures foundational mathematical concepts are learned before complex domains, stabilizing the knowledge acquisition process.

Frequently Asked Questions

What is catastrophic forgetting in machine learning?

Catastrophic forgetting occurs when a neural network loses previously learned information upon learning new tasks, effectively overwriting old knowledge with new patterns. In LeanAgent's context, this would mean the model forgetting how to prove theorems from earlier mathematical domains when trained on new repositories—a problem mitigated through Elastic Weight Consolidation that protects critical parameters identified by Fisher information.

How does Elastic Weight Consolidation differ from standard regularization?

While standard L2 regularization penalizes all parameter changes uniformly, Elastic Weight Consolidation (EWC) uses the Fisher information matrix to apply selective penalties. Parameters crucial for previous tasks (high Fisher values) receive heavy penalties for deviation, while unimportant parameters can adapt freely to new data, as implemented in retrieval/model.py through the ewc_loss() method.

Where is the progressive training loop implemented in the codebase?

The progressive training orchestration resides in leanagent.py, which controls the outer loop that iterates through repositories when run_progressive_training is enabled. This file handles curriculum ordering via sort_repositories_by_difficulty and coordinates the EWC lifecycle by triggering parameter snapshots and Fisher computation between repository transitions.

What role does the Fisher information matrix play in preventing forgetting?

The Fisher information matrix, computed in compute_fisher.py and retrieval/fisher_computation_module.py, serves as a diagnostic tool that measures how much each parameter contributed to the previous task's loss landscape. By feeding these values to PremiseRetriever.set_fisher_info(), the system identifies which weights must remain stable to preserve past knowledge, allowing the EWC loss term to mathematically anchor these critical parameters during new training.

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 →