Lightning-AI / Lightning-AI/pytorch-lightning

Add `WandbLogger` callback for customizing checkpoint artifact logging

Open
#17,913 1 comment 4 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

callback logger: wandb refactor
Dominant language
Python
Stars
31.4k
Forks
3.8k
Avg merge
6d 7h
Merged PRs (30d)
6

Description

### Outline & Motivation

It could be useful to add a callback to `WandbLogger` to allow custom handling of checkpoint artifacts. Examples of use cases:

1. I'm already writing checkpoints to persistent storage (e.g. using a `ModelCheckpoint` writing to S3), so I just want `WandbLogger` to log [reference artifacts](https://docs.wandb.ai/guides/artifacts/track-external-files) to them.
2. I want to add additional files or metadata to my WandB checkpoint artifacts.

We could refactor `WandbLogger` slightly:

```python
class WandbLogger:
def on_log_checkpoint_artifact(self, artifact, checkpoint_timestamp, path, score, tag):
artifact.add_file(path, name="model.ckpt")
return artifact

def _scan_and_log_checkpoints(self, checkpoint_callback: ModelCheckpoint) -> None:
# get checkpoints to be saved with associated score
checkpoints = _scan_checkpoints(checkpoint_callback, self._logged_model_time)

# log iteratively all new checkpoints
for t, p, s, tag in checkpoints:
metadata = (
{
"score": s.item() if isinstance(s, Tensor) else s,
"original_filename": Path(p).name,
checkpoint_callback.__class__.__name__: {
k: getattr(checkpoint_callback, k)
for k in [
"monitor",
"mode",
"save_last",
"save_top_k",
"save_weights_only",
"_every_n_train_steps",
]
# ensure it does not break if `ModelCheckpoint` args change
if hasattr(checkpoint_callback, k)
},
}
if _WANDB_GREATER_EQUAL_0_10_22
else None
)
if not self._checkpoint_name:
self._checkpoint_name = f"model-{self.experiment.id}"
artifact = wandb.Artifact(name=self._checkpoint_name, type="model", metadata=metadata)

# Handle artifact logic here
artifact = self.on_log_checkpoint_artifact(artifact, t, p, s, tag)

aliases = ["latest", "best"] if p == checkpoint_callback.best_model_path else ["latest"]
self.experiment.log_artifact(artifact, aliases=aliases)
# remember logged models - timestamp needed in case filename didn't change (lastkckpt or custom name)
self._logged_model_time[p] = t
```

Then, if users want custom artifact logging, they can subclass `WandbLogger` and override `on_log_checkpoint_artifact`:

```python
class ReferenceArtifactLogger(WandbLogger):
def on_log_checkpoint_artifact(self, artifact, checkpoint_timestamp, path, score, tag):
artifact.add_reference(path)
return artifact

### Pitch

_No response_

### Additional context

_No response_

cc @lantiga @justusschock @morganmcg1 @borisdayma @scottire @parambharat @awaelchli

Contributor guide

Open the contributing guide

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

Research direction

Start by locating WandbLogger and its _scan_and_log_checkpoints entry point, then review how ModelCheckpoint values and checkpoint artifacts are currently handled. Done means custom WandbLogger subclasses can alter artifact handling, including adding references or extra metadata, without breaking the existing logging flow.

Written by the indexing model from the issue text.

Assessment

Tech stack
python
Domain
machine-learning
Issue type
Feature
Difficulty
3/5
Estimated time
1-2 days
Activity status
Stale
Clarity
Mostly clear
Newbie friendliness
35/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.