How the ReProver Retriever Works: Architecture and Training in LeanAgent

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


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

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:

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, 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 links model-side arguments like model_name and max_seq_len with the data module:

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:

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:

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)

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


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


# Inside GeneratorModel.__init__

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

During sequence generation:

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, 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 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 (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), which handles checkpoint restoration and device placement. For automatic integration, provide the checkpoint path when initializing the GeneratorModel in 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.

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 →