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

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 file contains the logic that computes R@10 and reports it via log statements:


# 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, approximately lines 1150–1165) instantiates a ModelCheckpoint callback configured to maximize the R@10 metric:

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 provides utilities for loading saved states:

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 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, 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 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. 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.

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 →