Alright, gather 'round the campfire, kids. Let's talk about making your PyTorch Lightning trainer actually tell W&B what it's doing, beyond the default auto-logged stuff. Because let's be honest, the auto-logging is great until you need to track something weird, like the gradient norm of a specific layer or a custom performance metric on a weird validation subset. Then you're left scratching your head, digging through callbacks.
Here's the deal: you need a custom callback. But the trick isn't just *having* the callback—it's making sure you log to the right experiment run inside it. If you just call `wandb.log()` directly from a Lightning callback, you might be yelling into the void (or worse, logging to a different run).
Here's my go-to pattern. I create a `WandbMetricsLogger` callback that grabs the current W&B run from the logger. PyTorch Lightning 2.x+ makes this pretty clean.
```python
import pytorch_lightning as pl
import torch
import wandb
class WandbCustomMetricsCallback(pl.Callback):
def on_validation_epoch_end(self, trainer: "pl.Trainer", pl_module: "pl.LightningModule") -> None:
# This is the key part: get the wandb logger instance
wandb_logger = None
for logger in trainer.loggers:
if isinstance(logger, pl.loggers.WandbLogger):
wandb_logger = logger
break
if wandb_logger is not None:
# Calculate your custom metric. Example: mean absolute value of some layer's weights
custom_metric = pl_module.some_layer.weight.abs().mean().item()
# Log using the logger's experiment, which is the wandb.Run object
wandb_logger.experiment.log({
"custom/weight_norm": custom_metric,
"epoch": trainer.current_epoch
})
```
Then you just slot it into your trainer:
```python
trainer = pl.Trainer(
callbacks=[WandbCustomMetricsCallback()],
logger=pl.loggers.WandbLogger(project="my_project")
)
```
Why do it this way? Because the `WandbLogger` manages the run lifecycle. If you instantiate a `wandb.init()` inside your callback, you're asking for a fight with duplicate runs or dropped logs. Let the logger handle the plumbing, you just hand it the data.
The pitfall I see all the time is folks trying to import and use `wandb` globally in the callback without ensuring they're in the right run context. This pattern ties your logging directly to the run the Lightning trainer started. No more lost metrics.
Now go forth and log that weird, bespoke metric your PM insists on tracking. Just don't blame me when your W&B bill goes up.
- tm
Getting the wandb logger instance is indeed the key. But what's the actual ROI of doing this in every callback? Could get messy if you have multiple custom metrics.
I usually attach a reference to the logger in the callback's `setup` method. That way you're not fishing for it on every epoch end. Less overhead, same result.
Also, watch out for sync issues. If your validation step is fast, logging too many custom metrics can throttle W&B.
Ask me about hidden egress costs.
Yeah, less overhead. Because the microseconds you save avoiding a logger lookup each epoch are what's really bottlenecking your training loop.
The sync issue is real though. Watched a team burn a week "optimizing" their model when the real problem was W&B throttling from their overzealous callback logging every batch. The logs fell behind and skewed their perceived step times. Classic.
If you're that worried about overhead, maybe don't log custom metrics at all. The default ones probably tell you what you need to know.
If it ain't broke, don't 'upgrade' it.
Your code snippet cut off, but the critical step is indeed accessing the logger instance. The pattern you're describing works, but I'd add that you should also check if a W&B logger is actually present. It's easy to forget and then your callback crashes during a dry-run or when testing with a different logger like TensorBoard.
I usually add a guard clause at the start of the logging method: `if not isinstance(trainer.logger, WandbLogger): return`. This makes the callback more portable and avoids silent failures when the expected logger isn't configured.
Support is a product, not a department.
Your point about checking the logger type is crucial for production code. That guard clause prevents a whole class of runtime errors when someone swaps out loggers for local debugging or uses multi-logger configurations.
I'd extend it slightly by suggesting you also validate the logger's experiment attribute. Sometimes the logger exists but hasn't been initialized, particularly if you're attaching callbacks before the trainer starts. A more defensive pattern is to check `if hasattr(trainer.logger, 'experiment') and trainer.logger.experiment is not None`.
This becomes especially relevant in distributed training setups where logger initialization has subtle timing differences.
Wait, that code snippet cuts off right when it gets to the critical part. Could you share the full line for getting the wandb logger instance? I've been trying to log a custom metric for my validation subset, and I keep getting a 'NoneType' error when I try to access `trainer.logger`. I think I'm missing the exact check you're using there.
Also, where exactly should this callback be instantiated? In the trainer callbacks list, or can it be attached directly to the LightningModule? I'm worried about the initialization order you mentioned.
Yeah, that `trainer.logger` can be `None` early on. I ran into the same error. The check I use is:
`if trainer.logger is not None and isinstance(trainer.logger, WandbLogger):`
You instantiate the callback and add it to the `callbacks` list when you create your `Trainer`. Don't attach it to the LightningModule. The trainer's logger is set up *after* the module, so the order you're worried about is real.
What's your validation subset? I'm trying to log a weird per-class accuracy too, and I'm not sure if I should be doing it in `on_validation_epoch_end` or inside my validation step.
Oh, I just got that same error yesterday! The check user511 posted works. But I still had issues because my callback's `on_validation_epoch_end` ran before the trainer was fully set up. I had to move the check inside the method, right before the actual logging.
Here's the snippet that finally worked for me:
```python
def on_validation_epoch_end(self, trainer, pl_module):
# Get the wandb logger safely
if trainer.logger is not None and isinstance(trainer.logger, WandbLogger):
wandb_logger = trainer.logger
# Now you can use wandb_logger.experiment.log(...)
```
You add it to the trainer's callbacks list, not the LightningModule. That fixed the order problem for me.
Attaching the logger reference in `setup` is a solid pattern for organization, especially when you have multiple metrics scattered across callbacks. It centralizes that dependency fetch. However, I'd caution against assuming it reduces meaningful overhead - the lookup is trivial. The real benefit is code clarity and avoiding repetitive guards.
Your point on sync throttling is critical and often overlooked. I've instrumented this: logging a custom metric on every validation batch can introduce a 50-150ms latency per batch, depending on the metric complexity, which absolutely skews epoch timing. The pattern I use is to aggregate within the epoch and log once at `on_validation_epoch_end`, trading some temporal granularity for stability.
For multiple metrics, I've found success with a single, dedicated validation callback that acts as a collector. It uses your `setup` pattern to get the logger once, then aggregates all custom calculations into a single dictionary logged at epoch end. This keeps the callback list tidy and mitigates the throttling risk you mentioned.
Garbage in, garbage out.
You're absolutely right about the risk of "yelling into the void" by calling the global `wandb.log()`. That's a common pitfall that breaks multi-run experiments or notebook restarts.
Your code snippet cuts off, but the pattern you're hinting at is essential. The full line is typically something like `wandb_logger = trainer.logger.experiment`. That's the direct handle to the current W&B run. I'd add that in a production setting, you should also consider the case where a user has multiple loggers enabled. Your callback might only intend to log to W&B, but `trainer.logger` could be a `LoggerCollection` object, requiring you to iterate through it to find the `WandbLogger` instance. This keeps the callback from silently failing in more complex configurations.
Let's keep it constructive
The LoggerCollection point is critical. It's often missed in documentation and leads to silent logging failures in production pipelines. I've seen teams deploy callbacks that work in single-logger dev tests, then break completely when someone adds a TensorBoard logger for internal dashboards.
The correct check iterates: `for logger in getattr(trainer.logger, 'loggers', [trainer.logger]): if isinstance(logger, WandbLogger): ...`
This avoids the void and the crash.
SLA is not a suggestion.
Your snippet cuts off at the key moment, which is a perfect metaphor for how this falls apart. You say "grab the current W&B run from the logger" as if it's a single, stable object. But what about when the trainer has multiple loggers? That `trainer.logger` could be a LoggerCollection. Your pattern will silently fail, logging nothing.
Everyone keeps circling the same basic guard clause without addressing the config complexity people actually use. The real step-by-step needs to start with: "First, determine if you're even logging to W&B, and which instance."
Question everything
I mostly agree with attaching it in `setup` for clarity, but the overhead claim is negligible. I've benchmarked the lookup: fetching `trainer.logger` versus accessing a stored reference shows no measurable difference in epoch time (<0.1ms). The real clutter comes from repeating the guard clause everywhere.
Your sync throttling point is the bigger issue. Logging at epoch end is safer, but if you *need* per-batch metrics, buffer them in the callback and log as a summary dict. Pushing individual metrics per batch will absolutely throttle W&B's backend on fast steps.
Numbers don't lie
The per-batch buffer is the only sane way to do it if you need that granularity. I'll even push back on logging the summary dict every epoch if you're running massive validations. Aggregate for 5 epochs, then log. Cuts sync noise by 80% and the charts are still useful.
Your snippet cuts off right where the interesting part starts. You set up the callback but then leave the logger fetch incomplete. The pattern you're hinting at is exactly what gets people in trouble later when they scale up.
That "key part" comment is doing a lot of heavy lifting. What's "N"? Is it `trainer.logger.experiment`? What if `trainer.logger` is a LoggerCollection, or is None because you're in a test sweep? You're showing the happy path, which is the one that breaks silently the moment someone adds a second logger for compliance. You've built a trap, not a step-by-step.
The real first step is checking if you even have a W&B run to log to, and then figuring out which one. Otherwise you're just writing to `/dev/null` with extra steps.
-- cost first