# How the ReProver Retriever Works: Architecture and Training in LeanAgent

> Discover how the ReProver retriever works and its training process. Learn about its T5 encoder, contrastive learning, and EWC for efficient premise selection in LeanAgent.

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

---

**The ReProver retriever uses a T5 encoder with contrastive learning to select relevant Lean 4 premises, training via query-positive-negative triples with optional Elastic Weight Consolidation to prevent catastrophic forgetting.**

The ReProver retriever is a learned premise selection component in the **lean-dojo/leanagent** repository that enables automated theorem proving in Lean 4. It encodes proof states and premises using a T5 encoder, then retrieves the most relevant candidates through contrastive learning to guide the proof search process.

## Core Architecture of the ReProver Retriever

The retriever is implemented in [`retrieval/model.py`](https://github.com/lean-dojo/leanagent/blob/main/retrieval/model.py) as the `PremiseRetriever` class, a PyTorch Lightning module that handles encoding, caching, and retrieval.

### T5 Encoder and Embedding Generation

The retriever uses **T5EncoderModel** from Hugging Face Transformers, configured via `AutoTokenizer` and `AutoModel` (lines 75-77). The `_encode()` method (lines 68-90) processes text through the encoder, averages token embeddings using the attention mask, and applies L2 normalization to produce fixed-size vectors for both proof states (queries) and premises.

```python

# From retrieval/model.py

def _encode(self, input_ids, attention_mask):
    # T5 encoder forward pass

    outputs = self.encoder(input_ids=input_ids, attention_mask=attention_mask)
    # Mean pooling with attention mask

    embeddings = (outputs.last_hidden_state * attention_mask.unsqueeze(-1)).sum(1) / attention_mask.sum(1, keepdim=True)
    # L2 normalization

    return F.normalize(embeddings, p=2, dim=1)

```

### Corpus Caching and Retrieval Mechanism

The retriever maintains a cache of premise embeddings to enable fast nearest-neighbor search. The `corpus_embeddings` tensor stores the matrix of all premise vectors (lines 47-50), while `reindex_corpus()` (lines 63-82) recomputes these embeddings when `embeddings_staled` is set to True, processing the corpus in batches to maintain memory efficiency.

The `retrieve()` method (lines 84-115) tokenizes the query context, obtains its embedding, and fetches the k nearest premises via cosine similarity (computed as the inner product of unit-norm vectors):

```python
def retrieve(self, state, file_name, theorem_full_name, theorem_pos, k):
    # Encode the query context

    query_embedding = self._encode_query(state, file_name, theorem_full_name, theorem_pos)
    # Retrieve nearest premises from cached corpus

    premises, scores = self.corpus.get_nearest_premises(query_embedding, k)
    return premises, scores

```

### Contrastive Loss and EWC Regularization

The `forward()` method (lines 92-117) implements the training logic using contrastive learning. It encodes the context, a positive premise, and a set of negative premises, then computes a similarity matrix and applies MSE loss against binary labels (1 for positive, 0 for negatives).

To prevent catastrophic forgetting during continual learning, the retriever optionally implements **Elastic Weight Consolidation (EWC)**. The `set_previous_params()` method saves a snapshot of model weights and their Fisher information before training, while `ewc_loss()` (lines 93-108) adds a regularization term when `lamda > 0`, penalizing changes to important parameters:

```python
def ewc_loss(self):
    loss = 0
    for name, param in self.named_parameters():
        if name in self.previous_params:
            loss += (self.fisher[name] * (param - self.previous_params[name]).pow(2)).sum()
    return self.lamda * loss

```

## Training Procedure for the ReProver Retriever

Training is orchestrated through the PyTorch Lightning CLI defined in [`retrieval/main.py`](https://github.com/lean-dojo/leanagent/blob/main/retrieval/main.py), which links model and data arguments and manages the optimization loop.

### Data Loading and Contrastive Sampling

The `RetrievalDataModule` reads a JSONL corpus and dynamically creates training triples consisting of a query context, a positive premise (used in the actual proof), and multiple negative premises (irrelevant to the current proof state). These triples feed the contrastive loss function during the forward pass.

### Lightning CLI and Configuration

The CLI class in [`retrieval/main.py`](https://github.com/lean-dojo/leanagent/blob/main/retrieval/main.py) links model-side arguments like `model_name` and `max_seq_len` with the data module:

```python
class CLI(LightningCLI):
    def add_arguments_to_parser(self, parser) -> None:
        parser.link_arguments("model.model_name", "data.model_name")
        parser.link_arguments("data.max_seq_len", "model.max_seq_len")
        parser.add_argument('--data-path', type=str, required=True)

```

Training is launched via:

```bash
python retrieval/main.py fit \
    --config retrieval/confs/cli_lean4_random.yaml \
    --ckpt_path <output_dir>/ckpt \
    --data-path /path/to/lean4_dataset

```

### Validation Metrics and Corpus Re-indexing

During validation, the retriever ensures embeddings are current before computing metrics. The `on_validation_start` hook triggers `reindex_corpus()` to update `corpus_embeddings` with the latest encoder weights.

The `validation_step` (lines 108-151) computes **Recall@K** and **Mean Reciprocal Rank (MRR)** by retrieving premises for validation queries and comparing against ground-truth premises used in actual proofs. Recall@K measures the percentage of queries where the correct premise appears in the top K results, while MRR calculates the average reciprocal rank of the first correct premise.

Optimization uses schedulers and optimizers from `common.get_optimizers`, with checkpointing handled automatically by PyTorch Lightning.

## Practical Implementation Examples

### Loading and Indexing a Pretrained Retriever

To use the ReProver retriever for inference, load a checkpoint and index the premise corpus:

```python
from retrieval.model import PremiseRetriever
import torch

ckpt_path = "path/to/retriever/checkpoint.ckpt"
device = "cuda" if torch.cuda.is_available() else "cpu"

retriever = PremiseRetriever.load(
    ckpt_path=ckpt_path,
    device=device,
    freeze=True,  # keep weights frozen for inference

    config={}
)

# Load and index corpus

retriever.load_corpus(corpus_path)
retriever.reindex_corpus(batch_size=32)

```

### Retrieving Premises During Proof Search

During proof search, the prover queries the retriever with the current proof state:

```python

# state, file_name, theorem_full_name, theorem_pos come from the proof state

k = 10
premises, scores = retriever.retrieve(
    state=state,
    file_name=file_name,
    theorem_full_name=theorem_full_name,
    theorem_pos=theorem_pos,
    k=k
)

# premises is a list of Premise objects; scores are cosine similarities

for p, s in zip(premises, scores):
    print(f"{s:.4f} – {p.serialize()}")

```

### Integration with the Generator

The generator automatically instantiates the retriever when provided with a checkpoint path, as implemented in [`generator/model.py`](https://github.com/lean-dojo/leanagent/blob/main/generator/model.py):

```python

# Inside GeneratorModel.__init__

if ret_ckpt_path is not None:
    self.retriever = PremiseRetriever.load(
        ret_ckpt_path, device, freeze=False, config={}
    )

```

During sequence generation:

```python
if self.retriever is not None:
    retrieved, _ = self.retriever.retrieve(
        state, file, thm_name, thm_pos, k=5
    )
    # augment the generation context with retrieved premises

```

## Summary

- The **ReProver retriever** is implemented as the `PremiseRetriever` class in [`retrieval/model.py`](https://github.com/lean-dojo/leanagent/blob/main/retrieval/model.py), utilizing a **T5 encoder** to embed proof states and premises into dense vectors.
- It employs **contrastive learning** with MSE loss on query-positive-negative triples, with optional **Elastic Weight Consolidation (EWC)** to prevent catastrophic forgetting during continual training.
- The **PyTorch Lightning CLI** in [`retrieval/main.py`](https://github.com/lean-dojo/leanagent/blob/main/retrieval/main.py) orchestrates training, linking model and data configurations while managing optimization and checkpointing.
- Validation computes **Recall@K** and **Mean Reciprocal Rank (MRR)** after re-indexing the corpus to ensure metric accuracy with current encoder weights.
- At inference time, the retriever caches **corpus embeddings** for efficient nearest-neighbor retrieval via cosine similarity, integrating seamlessly with the generator and prover components via the `retrieve()` method.

## Frequently Asked Questions

### What encoder architecture powers the ReProver retriever?

The ReProver retriever uses **T5EncoderModel** from the Hugging Face Transformers library, specifically instantiated through `AutoTokenizer` and `AutoModel` in [`retrieval/model.py`](https://github.com/lean-dojo/leanagent/blob/main/retrieval/model.py) (lines 75-77). The encoder generates contextualized token embeddings that are mean-pooled and L2-normalized by the `_encode()` method (lines 68-90) to produce fixed-size vectors for both proof state queries and candidate premises.

### How does the ReProver retriever prevent catastrophic forgetting during continual learning?

The retriever implements **Elastic Weight Consolidation (EWC)**, a regularization technique that protects important parameters from changing drastically when learning new tasks. Before training begins, `set_previous_params()` saves a snapshot of model weights and computes their Fisher information. During the forward pass, if the `lamda` hyperparameter is greater than zero, the `ewc_loss()` method (lines 93-108) adds a penalty proportional to the squared difference between current and previous parameters, weighted by their importance to previous tasks.

### What validation metrics does the ReProver retriever use to evaluate performance?

During validation, the retriever computes **Recall@K** and **Mean Reciprocal Rank (MRR)**. The `validation_step` (lines 108-151) triggers after `on_validation_start` re-indexes the corpus to ensure embeddings reflect the current encoder weights. Recall@K measures the percentage of proof states where the ground-truth premise appears in the top K retrieved candidates, while MRR calculates the average of reciprocal ranks for the first correct premise across all validation queries.

### How do I integrate the ReProver retriever into the proof generation pipeline?

You can load the retriever using `PremiseRetriever.load()` (lines 124-126 in [`retrieval/model.py`](https://github.com/lean-dojo/leanagent/blob/main/retrieval/model.py)), which handles checkpoint restoration and device placement. For automatic integration, provide the checkpoint path when initializing the `GeneratorModel` in [`generator/model.py`](https://github.com/lean-dojo/leanagent/blob/main/generator/model.py); the generator will instantiate the retriever and call `retrieve()` during sequence generation, passing the current proof state, file name, theorem name, and position to fetch the top-k relevant premises for context augmentation.