# PremiseRetriever Checkpoint Strategy in LeanAgent: Optimizing for R@10 Metric

> Discover the PremiseRetriever checkpoint strategy in LeanAgent. Learn how R@10 optimization preserves retrieval performance for effective lifelong learning.

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

---

**The PremiseRetriever model saves checkpoints based on validation Recall-at-10 (R@10), keeping the model version that achieves the highest R@10 score to preserve retrieval performance in lifelong learning scenarios.**

The checkpoint strategy in the `lean-dojo/leanagent` repository ensures that the PremiseRetriever—a neural retriever that fetches relevant premises for theorem proving—maintains optimal performance across training epochs. By monitoring the **R@10** metric during validation, the system selects checkpoints that best balance remembering previously learned premises while acquiring new knowledge. This approach directly addresses catastrophic forgetting in progressive training pipelines.

## How R@10 Drives Checkpoint Selection

The PremiseRetriever evaluates performance using **Recall-at-10 (R@10)**, which measures whether the correct premise appears within the top-10 retrieved candidates. After every validation epoch, the trainer logs three key metrics: **R@1**, **R@10**, and **MRR** (Mean Reciprocal Rank).

The checkpoint callback specifically monitors the `R@10` value and preserves the model state whenever validation performance improves. In lifelong learning contexts, maximizing R@10 ensures the retriever maintains access to the most frequently relevant premises—the specific items the theorem prover will actually consider during proof search. A higher R@10 indicates superior preservation of previously learned knowledge alongside new factual acquisition.

## Implementation in Source Code

The checkpoint strategy is implemented across three critical files in the repository, with clear separation between metric computation and training orchestration.

### Metric Computation in retrieval/evaluate.py

The evaluation routine calculates retrieval metrics and logs them for the checkpoint callback to consume. The [`retrieval/evaluate.py`](https://github.com/lean-dojo/leanagent/blob/main/retrieval/evaluate.py) file contains the logic that computes R@10 and reports it via log statements:

```python

# As implemented in retrieval/evaluate.py

logger.info(f"R@1 = {R1} ... R@10 = {R10} ...")

```

These logged values feed directly into PyTorch Lightning’s callback system, enabling automated checkpoint management based on the computed recall metrics.

### ModelCheckpoint Configuration in leanagent.py

The training script ([`leanagent.py`](https://github.com/lean-dojo/leanagent/blob/main/leanagent.py), approximately lines 1150–1165) instantiates a `ModelCheckpoint` callback configured to maximize the R@10 metric:

```python
from pytorch_lightning.callbacks import ModelCheckpoint

checkpoint_callback = ModelCheckpoint(
    dirpath=log_dir / "checkpoints",
    filename="{epoch}-{step}",
    monitor="R@10",          # Metric used for checkpoint ranking

    mode="max",
    save_top_k=-1,           # Retain all checkpoints for analysis

    every_n_epochs=1,
)
trainer = pl.Trainer(
    callbacks=[lr_monitor, checkpoint_callback, early_stop_callback],
    max_epochs=10,
)

```

The `monitor="R@10"` and `mode="max"` parameters ensure the callback identifies the best-performing model according to recall performance. Setting `save_top_k=-1` retains all intermediate checkpoints for later fine-grained analysis, though the "best" model corresponds to the maximal R@10 value.

## Loading Optimized Checkpoints

When deploying the PremiseRetriever, load the checkpoint achieving the highest R@10 to ensure optimal retrieval performance. The `PremiseRetriever` class in [`retrieval/model.py`](https://github.com/lean-dojo/leanagent/blob/main/retrieval/model.py) provides utilities for loading saved states:

```python
from retrieval.model import PremiseRetriever

ckpt_path = find_latest_checkpoint()  # Locates the most recent .ckpt file

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
config = {
    "model_name": "google/byt5-small", 
    "lr": 1e-4, 
    "warmup_steps": 1000, 
    "max_seq_len": 512
}

retriever = PremiseRetriever.load(
    ckpt_path=ckpt_path,
    device=device,
    freeze=False,
    config=config,
)
retriever.eval()

```

This loading pattern preserves the model weights selected by the R@10-optimized checkpoint strategy, ensuring the deployed retriever reflects the best validation performance observed during training.

## Summary

- The PremiseRetriever uses **Recall-at-10 (R@10)** as the primary metric for checkpoint selection, prioritizing models that retrieve correct premises within the top-10 results.
- The `ModelCheckpoint` callback in [`leanagent.py`](https://github.com/lean-dojo/leanagent/blob/main/leanagent.py) monitors `R@10` with `mode="max"` and `save_top_k=-1`, saving all epochs but identifying the best via maximal R@10.
- Metric computation occurs in [`retrieval/evaluate.py`](https://github.com/lean-dojo/leanagent/blob/main/retrieval/evaluate.py), which logs R@1, R@10, and MRR after each validation epoch.
- This strategy mitigates catastrophic forgetting in lifelong learning by ensuring the retriever maintains strong performance on previously seen premises while learning new ones.

## Frequently Asked Questions

### What is R@10 and why does the PremiseRetriever use it instead of accuracy?

**R@10 (Recall-at-10)** measures whether the correct premise appears within the top-10 retrieved candidates for a given query. The PremiseRetriever uses R@10 instead of standard accuracy because theorem proving requires retrieving multiple relevant premises from a large corpus, not just selecting a single correct answer. According to the leanagent source code, maximizing R@10 ensures the system preserves access to the most useful premises—those the prover will actually consider—while mitigating catastrophic forgetting during progressive training.

### Where are the checkpoints saved in the leanagent repository?

Checkpoints are saved to the `checkpoints` subdirectory within the logging directory specified during training. The `ModelCheckpoint` callback in [`leanagent.py`](https://github.com/lean-dojo/leanagent/blob/main/leanagent.py) configures `dirpath=log_dir / "checkpoints"` with filenames formatted as `{epoch}-{step}`. Because `save_top_k=-1`, the system retains checkpoints from every validation epoch, enabling retrospective analysis of model evolution throughout the training process.

### How can I load a specific PremiseRetriever checkpoint?

Load specific checkpoints using the `PremiseRetriever.load()` method defined in [`retrieval/model.py`](https://github.com/lean-dojo/leanagent/blob/main/retrieval/model.py). Pass the checkpoint path, target device, configuration dictionary, and boolean flags for freezing layers. For production deployment, select the checkpoint file corresponding to the highest R@10 value logged during validation, ensuring optimal retrieval performance for downstream theorem proving tasks.

### Does the checkpoint strategy save intermediate models?

**Yes**, the checkpoint strategy saves all intermediate models. The `save_top_k=-1` parameter in the `ModelCheckpoint` configuration ensures every validation epoch produces a persisted checkpoint. While all models are retained for analysis, the training system specifically tracks which checkpoint achieved the maximum R@10 score, designating that version as the "best" model for inference and further fine-tuning operations.