# How to Use Custom PyTorch Lightning Callbacks for Monitoring Training Progress in LeanAgent

> Learn how to leverage custom PyTorch Lightning callbacks in LeanAgent to monitor training progress effectively. Automatically save checkpoints, track metrics, and log learning rates with ease.

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

---

**LeanAgent uses PyTorch Lightning's callback system to automatically save checkpoints, stop training when metrics plateau, and log learning rates, all configured through the `Trainer` in [`leanagent.py`](https://github.com/lean-dojo/leanagent/blob/main/leanagent.py).**

The `lean-dojo/leanagent` repository implements a premise-retrieval model for Lean theorem proving. To manage distributed training across multiple GPUs and ensure reproducibility, the codebase leverages **custom PyTorch Lightning callbacks** for comprehensive monitoring. These hooks integrate directly with the `pl.Trainer` to automate checkpointing, early stopping, and metric logging without manual intervention.

## Core Callbacks Defined in LeanAgent

In [`leanagent.py`](https://github.com/lean-dojo/leanagent/blob/main/leanagent.py) (line 45), the repository imports the standard callback suite:

```python
from pytorch_lightning.callbacks import ModelCheckpoint, EarlyStopping, LearningRateMonitor, Callback

```

These four classes form the monitoring backbone. While `ModelCheckpoint`, `EarlyStopping`, and `LearningRateMonitor` are instantiated with project-specific configurations, the base `Callback` class remains available for future custom extensions.

### ModelCheckpoint for Metric-Based Versioning

The `ModelCheckpoint` callback (configured around line 1554) saves model states after every epoch using filenames that encode performance metrics:

```python
checkpoint_callback = ModelCheckpoint(
    dirpath=RAID_DIR + "/" + CHECKPOINT_DIR,
    filename=dir_name + f"_lambda_{lambda_value}" + "_{epoch}-{Recall@10_val:.2f}",
    save_top_k=-1,            # Retain every checkpoint, not just the best

    every_n_epochs=1,         # Trigger after each training epoch

    monitor="Recall@10_val",  # Metric embedded in filename

    mode="max",
    verbose=True,
)

```

This configuration ensures that every checkpoint file includes the validation **Recall@10** score, creating a complete history of model versions for downstream analysis and reproducibility.

### EarlyStopping to Prevent Overfitting

To avoid wasteful computation when validation metrics stagnate, LeanAgent configures `EarlyStopping` (around line 1569) to monitor the same `Recall@10_val` metric:

```python
early_stop_callback = EarlyStopping(
    monitor="Recall@10_val",
    patience=5,      # Halt training after 5 epochs without improvement

    mode="max",
    verbose=True,
)

```

When the validation recall fails to increase for five consecutive epochs, the callback triggers a graceful training termination, returning the best model state automatically.

### LearningRateMonitor for Schedule Tracking

The `LearningRateMonitor` (instantiated near line 1554) logs the optimizer's learning rate at every step:

```python
lr_monitor = LearningRateMonitor(logging_interval='step')

```

By setting `logging_interval='step'`, LeanAgent captures fine-grained learning rate changes—critical when using schedulers like cosine annealing or warm restarts. These values stream automatically to TensorBoard or CSV logs via Lightning's logger integration.

## Registering Callbacks with the Trainer

All configured callbacks are registered with the `pl.Trainer` in [`leanagent.py`](https://github.com/lean-dojo/leanagent/blob/main/leanagent.py) (lines 1570-1575):

```python
trainer = pl.Trainer(
    accelerator="gpu",
    precision="bf16-mixed",
    strategy=ddp_strategy,
    devices=4,
    callbacks=[lr_monitor, checkpoint_callback, early_stop_callback],
    max_epochs=current_epoch + epochs_per_repo,
    log_every_n_steps=1,
    default_root_dir=custom_log_dir,
)

```

The `callbacks` list accepts any `Callback` subclass. LeanAgent passes the three instantiated monitors, ensuring they execute automatically at the appropriate training lifecycle hooks (e.g., `on_train_epoch_end`, `on_validation_end`).

## Extending with Custom Callback Subclasses

While the current implementation uses standard Lightning callbacks, [`leanagent.py`](https://github.com/lean-dojo/leanagent/blob/main/leanagent.py) imports the base `Callback` class (line 45) to facilitate future extensions. Developers can subclass `Callback` to implement project-specific logic, such as custom metric aggregation or distributed logging.

For example, to log the average parameter norm after each epoch:

```python
from pytorch_lightning.callbacks import Callback
import torch

class ParameterNormLogger(Callback):
    def on_train_epoch_end(self, trainer, pl_module):
        avg_norm = torch.mean(
            torch.stack([p.norm() for p in pl_module.parameters()])
        )
        trainer.logger.experiment.add_scalar(
            "avg_param_norm", avg_norm, trainer.current_epoch
        )

# Usage: add to the callbacks list in leanagent.py

# callbacks=[lr_monitor, checkpoint_callback, early_stop_callback, ParameterNormLogger()]

```

This pattern integrates seamlessly with the existing `Trainer` configuration in [`leanagent.py`](https://github.com/lean-dojo/leanagent/blob/main/leanagent.py), allowing LeanAgent to scale its monitoring capabilities without modifying core training logic.

## Summary

- **LeanAgent** leverages PyTorch Lightning's callback architecture in [`leanagent.py`](https://github.com/lean-dojo/leanagent/blob/main/leanagent.py) to automate training monitoring.
- **ModelCheckpoint** saves every epoch's state with `Recall@10_val` embedded in filenames, ensuring full reproducibility.
- **EarlyStopping** halts training after 5 epochs of stagnant validation recall, conserving compute resources.
- **LearningRateMonitor** logs learning rates at every step, capturing scheduler dynamics for analysis.
- The base **Callback** class is imported for future extensions, allowing custom monitoring logic without framework modifications.

## Frequently Asked Questions

### How does LeanAgent determine when to save model checkpoints?

LeanAgent configures `ModelCheckpoint` in [`leanagent.py`](https://github.com/lean-dojo/leanagent/blob/main/leanagent.py) (around line 1554) to trigger `every_n_epochs=1`, saving a checkpoint after each training epoch. The filename pattern includes the validation `Recall@10` metric (`{epoch}-{Recall@10_val:.2f}`), creating a versioned history tied to model performance.

### What metric does LeanAgent monitor for early stopping?

The `EarlyStopping` callback monitors `Recall@10_val` (validation recall at rank 10). As implemented in [`leanagent.py`](https://github.com/lean-dojo/leanagent/blob/main/leanagent.py) (around line 1569), it uses `patience=5` and `mode="max"`, meaning training stops automatically if the recall score fails to improve for five consecutive epochs.

### Can I add custom logging to LeanAgent without modifying the core training loop?

Yes. Because [`leanagent.py`](https://github.com/lean-dojo/leanagent/blob/main/leanagent.py) imports the base `Callback` class from PyTorch Lightning (line 45), you can subclass `Callback` to implement `on_train_epoch_end`, `on_validation_batch_end`, or other hooks. Pass your custom callback instance to the `callbacks` list in the `pl.Trainer` initialization (around line 1570) to inject logic without altering the main training script.