# How Progressive Training Prevents Catastrophic Forgetting in LeanAgent

> Discover how LeanAgent's progressive training prevents catastrophic forgetting in lifelong learning. Learn about curriculum-based EWC and parameter preservation.

- Repository: [LeanDojo/leanagent](https://github.com/lean-dojo/leanagent)
- Tags: deep-dive
- Published: 2026-03-05

---

**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`](https://github.com/lean-dojo/leanagent/blob/main/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`](https://github.com/lean-dojo/leanagent/blob/main/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`](https://github.com/lean-dojo/leanagent/blob/main/compute_fisher.py) module and its helper [`retrieval/fisher_computation_module.py`](https://github.com/lean-dojo/leanagent/blob/main/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`](https://github.com/lean-dojo/leanagent/blob/main/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`:

```python
total_loss = task_loss + self.lamda * ewc_loss

```

The `lamda` parameter controls regularization strength—setting `lambdas = [0.1]` in [`leanagent.py`](https://github.com/lean-dojo/leanagent/blob/main/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:

```python

# 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:

```python

# 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:

```python

# 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:

```python

# 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`](https://github.com/lean-dojo/leanagent/blob/main/compute_fisher.py) quantify parameter importance for preserving knowledge from completed tasks.
- **EWC regularization** in [`retrieval/model.py`](https://github.com/lean-dojo/leanagent/blob/main/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`](https://github.com/lean-dojo/leanagent/blob/main/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`](https://github.com/lean-dojo/leanagent/blob/main/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`](https://github.com/lean-dojo/leanagent/blob/main/compute_fisher.py) and [`retrieval/fisher_computation_module.py`](https://github.com/lean-dojo/leanagent/blob/main/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.