Lightning-AI / Lightning-AI/pytorch-lightning

Edge case causes incorrect filesystem to be selected for finding cloud checkpoints

Open
#17,912 6 comments 1 reaction 0 assignees View on GitHub

Nobody has claimed this yet.

3rd party bug checkpointing help wanted ver: 2.0.x
Dominant language
Python
Stars
31.4k
Forks
3.8k
Avg merge
6d 7h
Merged PRs (30d)
6

Description

### Bug description

When both of the following happen together:
1. a logger is used with a cloud (e.g `s3://` or `gcs://` protocol) save dir
2. a `ModelCheckpoint` is used without passing a `dirpath`

The desired behaviors are:
1. the checkpoint directory is resolved (via `ModelCheckpoint.__resolve_ckpt_dir`) to `$logger.save_dir/$logger.name/$logger.version/checkpoints` and the `ModelCheckpoint` callback saves them there.
2. `ModelCheckpoint._find_last_checkpoints` will find `$logger.save_dir/$logger.name/$logger.version/checkpoints/last.ckpt`. If will first check if that path exists on the filesystem instantiated in `ModelCheckpoint.__init_ckpt_dir`.

Desired behavior 1 works, 2 does not. There are two bugs:
1. `ModelCheckpoint.__init_ckpt_dir` will [select the wrong filesystem](https://github.com/Lightning-AI/lightning/blob/6eae2310d6dae086596e5bdddd08e8cd3884336e/src/lightning/pytorch/callbacks/model_checkpoint.py#L443) when `dirpath` is `None`, causing `ModelCheckpoint._find_last_checkpoints` to [not find](https://github.com/Lightning-AI/lightning/blob/6eae2310d6dae086596e5bdddd08e8cd3884336e/src/lightning/pytorch/callbacks/model_checkpoint.py#L612) the cloud filepaths.
2. Even if the correct `ModelCheckpoint._fs` were used, `_find_last_checkpoints` returns a set of paths with their protocols stripped (due to the [call to _fs.ls](https://github.com/Lightning-AI/lightning/blob/6eae2310d6dae086596e5bdddd08e8cd3884336e/src/lightning/pytorch/callbacks/model_checkpoint.py#L615)). This causes `_CheckpointConnector_parse_ckpt_path` to then also [select the wrong filesystem](https://github.com/Lightning-AI/lightning/blob/master/src/lightning/pytorch/trainer/connectors/checkpoint_connector.py#L185-L186), resulting in no checkpoints found.
### What version are you seeing the problem on?

v2.0; but likely also present on others

### How to reproduce the bug

1. Use a logger with a cloud save dir
2. Create some cloud checkpoint, e.g. `s3://.../logger_name/logger_version/checkpoints/last.ckpt`
6. From a new job, try to resume training using `ckpt_path="last"`
7. A warning will be emitted about how lightning couldn't find the checkpoint

### Error messages and logs

```
UserWarning: .fit(ckpt_path="last") is set, but there is no last checkpoint available. No checkpoint will be loaded.
```

### Environment

Current environment

```
#- Lightning Component: ModelCheckpoint
#- PyTorch Lightning Version: 2.0.3
#- Lightning App Version: N/A
#- PyTorch Version: 2.0.1
#- Python version: 3.10.11
#- OS: Linux
#- CUDA/cuDNN version: 11.7
#- GPU models and configuration: 1x T4
#- How you installed Lightning: `pip`
#- Running environment of LightningApp (e.g. local, cloud): AWS Sagemaker
```

### More info

Here is my current workaround for S3 checkpoints:

```python
from s3fs import S3FileSystem

class S3ModelCheckpoint(ModelCheckpoint):
def __init__(self, *args: str | None, **kwargs: str | None) -> None:
super().__init__(*args, **kwargs)
self._fs = S3FileSystem()

def _find_last_checkpoints(self, trainer: "L.Trainer") -> set[str]:
ckpts = {"s3://" + ckpt for ckpt in super()._find_last_checkpoints(trainer)}
return ckpts
```

The [Universal Pathlib](https://github.com/fsspec/universal_pathlib) project fixes the behavior of cloud paths so that the procotols aren't stripped off. Could be worth looking into, to prevent these sorts of edge cases from occurring.

cc @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 in src/lightning/pytorch/callbacks/model_checkpoint.py, reading __init_ckpt_dir and _find_last_checkpoints, then inspect src/lightning/pytorch/trainer/connectors/checkpoint_connector.py around checkpoint path parsing. Reproduce with a cloud logger save directory, ModelCheckpoint without dirpath, and ckpt_path="last". Done means the cloud filesystem is selected and the checkpoint retains its protocol so resuming finds last.ckpt.

Written by the indexing model from the issue text.

Assessment

Tech stack
aws, python, pytorch
Domain
cloud, machine-learning
Issue type
Bug
Difficulty
4/5
Estimated time
3-5 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.