Lightning-AI / Lightning-AI/pytorch-lightning

Move strategy-specific dataloader logic to the stategies

Open
#11,756 1 comment 7 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

data handling design refactor strategy
Dominant language
Python
Stars
31.4k
Forks
3.8k
Avg merge
6d 7h
Merged PRs (30d)
6

Description

## Proposed refactor

### Motivation

Strategies today interact with dataloading, especially in distributed training. It makes sense for the strategy to directly handle this logic.

This would reduce and simplify interactions elsewhere in the trainer, in particular strategy state -> trainer properties -> data connector logic

And would remove the hacky patching that IPUs do right now 😧 https://github.com/PyTorchLightning/pytorch-lightning/blob/cc43d07db1ab77385feff04c01f040c5cad805a9/pytorch_lightning/strategies/ipu.py#L125-L130

### Pitch

The Strategy interface already offers a `process_dataloader` method:
https://github.com/PyTorchLightning/pytorch-lightning/blob/cc43d07db1ab77385feff04c01f040c5cad805a9/pytorch_lightning/strategies/strategy.py#L369-L375

However, there's a ton of strategy-specific logic written in the trainer's data connector:

For example, warnings with DDP spawn:
https://github.com/PyTorchLightning/pytorch-lightning/blob/cc43d07db1ab77385feff04c01f040c5cad805a9/pytorch_lightning/trainer/connectors/data_connector.py#L278-L312

generic multi-processing warning:
https://github.com/PyTorchLightning/pytorch-lightning/blob/cc43d07db1ab77385feff04c01f040c5cad805a9/pytorch_lightning/trainer/connectors/data_connector.py#L314-L322

the is_distributed flag is primarily used here within the trainer: https://github.com/PyTorchLightning/pytorch-lightning/blob/cc43d07db1ab77385feff04c01f040c5cad805a9/pytorch_lightning/trainer/connectors/data_connector.py#L324-L330

Distributed sampler kwargs is strategy specific:
https://github.com/PyTorchLightning/pytorch-lightning/blob/cc43d07db1ab77385feff04c01f040c5cad805a9/pytorch_lightning/trainer/trainer.py#L2126-L2129

IPU check: https://github.com/PyTorchLightning/pytorch-lightning/blob/cc43d07db1ab77385feff04c01f040c5cad805a9/pytorch_lightning/trainer/connectors/data_connector.py#L364

In my opinion we could simplify this by moving relevant logic into `strategy.process_dataloader` instead. Common logic can still be abstracted out into utility functions to share across different strategy classes.

We could either:
- augment `process_dataloader` to acept more metadata, such as the trainer's running stage and if this is a dataloader being used for training/validation/test/predict

```py
def process_dataloader(self, dataloader, use: str) -> Union[dataloader, iterable]:
```

or offer multiple APIs that map to the DataLoader hooks: https://github.com/PyTorchLightning/pytorch-lightning/blob/cc43d07db1ab77385feff04c01f040c5cad805a9/pytorch_lightning/core/hooks.py#L406
```py
def process_train_dataloader(self, dataloader):
def process_val_dataloader(self, dataloader):
def process_test_dataloader(self, dataloader):
def process_predict_dataloader(self, dataloader):
```

Over time, the Trainer flag `replace_sampler_ddp` makes much more sense on the specific distributed strategy constructors instead of on the trainer.

### Additional context

______________________________________________________________________

#### If you enjoy Lightning, check out our other projects! ⚡

- [**Metrics**](https://github.com/PyTorchLightning/metrics): Machine learning metrics for distributed, scalable PyTorch applications.

- [**Lite**](https://pytorch-lightning.readthedocs.io/en/latest/starter/lightning_lite.html): enables pure PyTorch users to scale their existing code on any kind of device while retaining full control over their own loops and optimization logic.

- [**Flash**](https://github.com/PyTorchLightning/lightning-flash): The fastest way to get a Lightning baseline! A collection of tasks for fast prototyping, baselining, fine-tuning, and solving problems with deep learning.

- [**Bolts**](https://github.com/PyTorchLightning/lightning-bolts): Pretrained SOTA Deep Learning models, callbacks, and more for research and production with PyTorch Lightning and PyTorch.

- [**Lightning Transformers**](https://github.com/PyTorchLightning/lightning-transformers): Flexible interface for high-performance research using SOTA Transformers leveraging Pytorch Lightning, Transformers, and Hydra.

cc @justusschock @awaelchli @akihironitta @rohitgr7 @ninginthecloud @tchaton @borda

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 with strategies/strategy.py and trainer/connectors/data_connector.py, then inspect the referenced trainer.py and core/hooks.py locations. Compare the existing process_dataloader interface with the DDP spawn, multiprocessing, distributed sampler, and IPU-specific logic. Done means the strategy owns the relevant dataloader behavior and the trainer/data connector no longer coordinates strategy-specific cases.

Written by the indexing model from the issue text.

Assessment

Tech stack
python, pytorch
Domain
distributed-systems, machine-learning
Issue type
Refactor
Difficulty
5/5
Estimated time
Over a week
Activity status
Stale
Clarity
Mostly clear
Newbie friendliness
35/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.