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

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.

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 (line 45), the repository imports the standard callback suite:

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:

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:

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:

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 (lines 1570-1575):

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

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, allowing LeanAgent to scale its monitoring capabilities without modifying core training logic.

Summary

  • LeanAgent leverages PyTorch Lightning's callback architecture in 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 (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 (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 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.

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 →