mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-12 12:40:20 +08:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d09a9e2c96 | ||
|
|
0f79e9d74e | ||
|
|
9fa8120805 | ||
|
|
715bf23105 | ||
|
|
88ac4a0849 | ||
|
|
383746b87a | ||
|
|
fffc09830f | ||
|
|
aadf8e16aa | ||
|
|
4b04dc06d4 | ||
|
|
0e42d28415 | ||
|
|
09dba13cde | ||
|
|
42a45bb273 | ||
|
|
5604e955eb | ||
|
|
6d34224e68 | ||
|
|
24a3246bc1 | ||
|
|
39b15855ed | ||
|
|
c6da6eb46c | ||
|
|
d23d25646a | ||
|
|
bd6521a584 | ||
|
|
2ce3e3e108 | ||
|
|
deeb82d28f | ||
|
|
74817c2fb1 | ||
|
|
b989358c9b | ||
|
|
0d47561a31 |
@@ -11,6 +11,10 @@
|
||||
</p>
|
||||
<p align="center">
|
||||
<a href="https://badge.fury.io/py/pytorch-lightning"><img src="https://badge.fury.io/py/pytorch-lightning.svg" alt="PyPI version" height="18"></a>
|
||||
<a href="https://pepy.tech/project/pytorch-lightning"><img src="https://pepy.tech/badge/pytorch-lightning" alt="PyPI version" height="18"></a>
|
||||
</p>
|
||||
|
||||
<p align="center">
|
||||
<a href="https://github.com/williamFalcon/pytorch-lightning/tree/master/tests"><img src="https://github.com/williamFalcon/pytorch-lightning/blob/master/coverage.svg"></a>
|
||||
<a href="https://travis-ci.org/williamFalcon/pytorch-lightning"><img src="https://travis-ci.org/williamFalcon/pytorch-lightning.svg?branch=master"></a>
|
||||
<a href="https://williamfalcon.github.io/pytorch-lightning/"><img src="https://readthedocs.org/projects/pytorch-lightning/badge/?version=latest"></a>
|
||||
@@ -25,17 +29,72 @@ pip install pytorch-lightning
|
||||
**[View the docs here](https://williamfalcon.github.io/pytorch-lightning/)**
|
||||
|
||||
## What is it?
|
||||
Keras and fast.ai are too abstract for researchers. Lightning abstracts the full training loop but gives you control in the critical points.
|
||||
Lightning defers training and validation loop logic to you. It guarantees correct, modern best practices for the core training logic.
|
||||
|
||||
|
||||
## Why do I want to use lightning?
|
||||
Because you don't want to define a training loop, validation loop, gradient clipping, checkpointing, loading,
|
||||
gpu training, etc... every time you start a project. Let lightning handle all of that for you! Just define your
|
||||
data and what happens in the training, testing and validation loop and lightning will do the rest.
|
||||
When starting a new project the last thing you want to do is recode a training loop, model loading/saving, distributed training, when to validate, etc... You're likely to spend a long time ironing out all the bugs without even getting to the core of your research.
|
||||
|
||||
With lightning, you guarantee those parts of your code work so you can focus on what the meat of the research: Data and training, validation loop logic. Don't worry about multiple gpus or speeding up your code, lightning will do that for you!
|
||||
|
||||
To use lightning do 2 things:
|
||||
1. [Define a Trainer](https://github.com/williamFalcon/pytorch-lightning/blob/master/examples/new_project_templates/trainer_cpu_template.py).
|
||||
2. [Define a LightningModel](https://github.com/williamFalcon/pytorch-lightning/blob/master/examples/new_project_templates/lightning_module_template.py).
|
||||
1. [Define a LightningModel](https://williamfalcon.github.io/pytorch-lightning/LightningModule/RequiredTrainerInterface/)
|
||||
```python
|
||||
import pytorch_lightning as ptl
|
||||
import torch
|
||||
|
||||
class CoolModel(ptl.LightningModule):
|
||||
|
||||
def __init(self):
|
||||
self.l1 = torch.nn.Linear(28*28, 10)
|
||||
|
||||
def forward(self, x):
|
||||
return self.l1(x)
|
||||
|
||||
def training_step(self, batch, batch_nb):
|
||||
x, y = batch
|
||||
y_hat = self.forward(x)
|
||||
return {'tng_loss': some_loss(y_hat, y)}
|
||||
|
||||
def validation_step(self, batch, batch_nb):
|
||||
x, y = batch
|
||||
y_hat = self.forward(x)
|
||||
return {'val_loss': some_loss(y_hat, y)}
|
||||
|
||||
def configure_optimizers(self):
|
||||
return [optim.Adam(self.parameters(), lr=0.02)]
|
||||
|
||||
@ptl.data_loader
|
||||
def tng_dataloader(self):
|
||||
return DataLoader(MNIST('path/to/save', train=True), batch_size=32)
|
||||
|
||||
@ptl.data_loader
|
||||
def val_dataloader(self):
|
||||
return DataLoader(MNIST('path/to/save', train=False), batch_size=32)
|
||||
|
||||
@ptl.data_loader
|
||||
def test_dataloader(self):
|
||||
return DataLoader(MNIST('path/to/save', train=False), batch_size=32)
|
||||
|
||||
```
|
||||
|
||||
2. Fit with a [trainer](https://williamfalcon.github.io/pytorch-lightning/Trainer/)
|
||||
```python
|
||||
from pytorch_lightning import Trainer
|
||||
from test_tube import Experiment
|
||||
|
||||
model = CoolModel()
|
||||
|
||||
# fit on 32 gpus across 4 nodes
|
||||
exp = Experiment(save_dir='some/dir')
|
||||
trainer = Trainer(experiment=exp, nb_gpu_nodes=4, gpus=[0,1,2,3,4,5,6,7])
|
||||
|
||||
trainer.fit(model)
|
||||
|
||||
# see all experiment metrics here
|
||||
# tensorboard --log_dir some/dir
|
||||
```
|
||||
|
||||
|
||||
## What does lightning control for me?
|
||||
Everything!
|
||||
|
||||
@@ -237,10 +237,10 @@ def load_model_specific(self, checkpoint):
|
||||
### tng_dataloader
|
||||
|
||||
``` {.python}
|
||||
@property
|
||||
@ptl.data_loader
|
||||
def tng_dataloader(self)
|
||||
```
|
||||
Called by lightning during training loop. Define it as a property.
|
||||
Called by lightning during training loop. Make sure to use the @ptl.data_loader decorator, this ensures not calling this function until the data are needed.
|
||||
|
||||
##### Return
|
||||
Pytorch DataLoader
|
||||
@@ -248,32 +248,26 @@ Pytorch DataLoader
|
||||
**Example**
|
||||
|
||||
``` {.python}
|
||||
@property
|
||||
@ptl.data_loader
|
||||
def tng_dataloader(self):
|
||||
if self._tng_dataloader is None:
|
||||
try:
|
||||
transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,), (1.0,))])
|
||||
dataset = MNIST(root='/path/to/mnist/', train=True, transform=transform, download=True)
|
||||
loader = torch.utils.data.DataLoader(
|
||||
dataset=dataset,
|
||||
batch_size=self.hparams.batch_size,
|
||||
shuffle=True
|
||||
)
|
||||
self._tng_dataloader = loader
|
||||
except Exception as e:
|
||||
raise e
|
||||
|
||||
return self._tng_dataloader
|
||||
transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,), (1.0,))])
|
||||
dataset = MNIST(root='/path/to/mnist/', train=True, transform=transform, download=True)
|
||||
loader = torch.utils.data.DataLoader(
|
||||
dataset=dataset,
|
||||
batch_size=self.hparams.batch_size,
|
||||
shuffle=True
|
||||
)
|
||||
return loader
|
||||
```
|
||||
|
||||
---
|
||||
### val_dataloader
|
||||
|
||||
``` {.python}
|
||||
@property
|
||||
@ptl.data_loader
|
||||
def tng_dataloader(self)
|
||||
```
|
||||
Called by lightning during validation loop. Define it as a property.
|
||||
Called by lightning during validation loop. Make sure to use the @ptl.data_loader decorator, this ensures not calling this function until the data are needed.
|
||||
|
||||
##### Return
|
||||
Pytorch DataLoader
|
||||
@@ -281,32 +275,27 @@ Pytorch DataLoader
|
||||
**Example**
|
||||
|
||||
``` {.python}
|
||||
@property
|
||||
@ptl.data_loader
|
||||
def val_dataloader(self):
|
||||
if self._val_dataloader is None:
|
||||
try:
|
||||
transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,), (1.0,))])
|
||||
dataset = MNIST(root='/path/to/mnist/', train=False, transform=transform, download=True)
|
||||
loader = torch.utils.data.DataLoader(
|
||||
dataset=dataset,
|
||||
batch_size=self.hparams.batch_size,
|
||||
shuffle=True
|
||||
)
|
||||
self._val_dataloader = loader
|
||||
except Exception as e:
|
||||
raise e
|
||||
|
||||
return self._val_dataloader
|
||||
transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,), (1.0,))])
|
||||
dataset = MNIST(root='/path/to/mnist/', train=False, transform=transform, download=True)
|
||||
loader = torch.utils.data.DataLoader(
|
||||
dataset=dataset,
|
||||
batch_size=self.hparams.batch_size,
|
||||
shuffle=True
|
||||
)
|
||||
|
||||
return loader
|
||||
```
|
||||
|
||||
---
|
||||
### test_dataloader
|
||||
|
||||
``` {.python}
|
||||
@property
|
||||
@ptl.data_loader
|
||||
def test_dataloader(self)
|
||||
```
|
||||
Called by lightning during test loop. Define it as a property.
|
||||
Called by lightning during test loop. Make sure to use the @ptl.data_loader decorator, this ensures not calling this function until the data are needed.
|
||||
|
||||
##### Return
|
||||
Pytorch DataLoader
|
||||
@@ -314,22 +303,17 @@ Pytorch DataLoader
|
||||
**Example**
|
||||
|
||||
``` {.python}
|
||||
@property
|
||||
@ptl.data_loader
|
||||
def test_dataloader(self):
|
||||
if self._test_dataloader is None:
|
||||
try:
|
||||
transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,), (1.0,))])
|
||||
dataset = MNIST(root='/path/to/mnist/', train=False, transform=transform, download=True)
|
||||
loader = torch.utils.data.DataLoader(
|
||||
dataset=dataset,
|
||||
batch_size=self.hparams.batch_size,
|
||||
shuffle=True
|
||||
)
|
||||
self._test_dataloader = loader
|
||||
except Exception as e:
|
||||
raise e
|
||||
|
||||
return self._test_dataloader
|
||||
transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,), (1.0,))])
|
||||
dataset = MNIST(root='/path/to/mnist/', train=False, transform=transform, download=True)
|
||||
loader = torch.utils.data.DataLoader(
|
||||
dataset=dataset,
|
||||
batch_size=self.hparams.batch_size,
|
||||
shuffle=True
|
||||
)
|
||||
|
||||
return loader
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
@@ -1,2 +1,3 @@
|
||||
from .models import Trainer
|
||||
from .root_module.root_module import LightningModule
|
||||
from .root_module.root_module import LightningModule
|
||||
from .root_module.decorators import data_loader
|
||||
@@ -10,6 +10,7 @@ from torch import optim
|
||||
from torch.utils.data import DataLoader
|
||||
from torch.utils.data.distributed import DistributedSampler
|
||||
|
||||
import pytorch_lightning as ptl
|
||||
from pytorch_lightning.root_module.root_module import LightningModule
|
||||
|
||||
|
||||
@@ -200,35 +201,20 @@ class LightningTemplateModel(LightningModule):
|
||||
|
||||
return loader
|
||||
|
||||
@property
|
||||
@ptl.data_loader
|
||||
def tng_dataloader(self):
|
||||
if self._tng_dataloader is None:
|
||||
try:
|
||||
self._tng_dataloader = self.__dataloader(train=True)
|
||||
except Exception as e:
|
||||
print(e)
|
||||
raise e
|
||||
return self._tng_dataloader
|
||||
print('tng data loader called')
|
||||
return self.__dataloader(train=True)
|
||||
|
||||
@property
|
||||
@ptl.data_loader
|
||||
def val_dataloader(self):
|
||||
if self._val_dataloader is None:
|
||||
try:
|
||||
self._val_dataloader = self.__dataloader(train=False)
|
||||
except Exception as e:
|
||||
print(e)
|
||||
raise e
|
||||
return self._val_dataloader
|
||||
print('val data loader called')
|
||||
return self.__dataloader(train=False)
|
||||
|
||||
@property
|
||||
@ptl.data_loader
|
||||
def test_dataloader(self):
|
||||
if self._test_dataloader is None:
|
||||
try:
|
||||
self._test_dataloader = self.__dataloader(train=False)
|
||||
except Exception as e:
|
||||
print(e)
|
||||
raise e
|
||||
return self._test_dataloader
|
||||
print('test data loader called')
|
||||
return self.__dataloader(train=False)
|
||||
|
||||
@staticmethod
|
||||
def add_model_specific_args(parent_parser, root_dir): # pragma: no cover
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import torch.nn as nn
|
||||
import numpy as np
|
||||
from pytorch_lightning.root_module.root_module import LightningModule
|
||||
from pytorch_lightning import LightningModule
|
||||
from test_tube import HyperOptArgumentParser
|
||||
from torchvision.datasets import MNIST
|
||||
import torchvision.transforms as transforms
|
||||
@@ -149,7 +149,7 @@ class ExampleModel1(LightningModule):
|
||||
|
||||
return loader
|
||||
|
||||
@property
|
||||
@data_loader
|
||||
def tng_dataloader(self):
|
||||
if self._tng_dataloader is None:
|
||||
try:
|
||||
|
||||
@@ -438,14 +438,14 @@ class Trainer(TrainerIO):
|
||||
|
||||
# ON CPU
|
||||
else:
|
||||
# CHOOSE OPTIMIZER
|
||||
# filter out the weights that were done on gpu so we can load on good old cpus
|
||||
self.optimizers = model.configure_optimizers()
|
||||
|
||||
# run through amp wrapper
|
||||
if self.use_amp:
|
||||
raise MisconfigurationException('amp + cpu is not supported. Please use a GPU option')
|
||||
|
||||
# CHOOSE OPTIMIZER
|
||||
# filter out the weights that were done on gpu so we can load on good old cpus
|
||||
self.optimizers = model.configure_optimizers()
|
||||
|
||||
self.__run_pretrain_routine(model)
|
||||
|
||||
# return 1 when finished
|
||||
@@ -544,25 +544,25 @@ class Trainer(TrainerIO):
|
||||
os.environ['MASTER_PORT'] = f'{port}'
|
||||
|
||||
# figure out the root node addr
|
||||
root_node = os.environ['SLURM_NODELIST'].split(' ')[0]
|
||||
try:
|
||||
root_node = os.environ['SLURM_NODELIST'].split(' ')[0]
|
||||
except Exception as e:
|
||||
root_node = '127.0.0.2'
|
||||
|
||||
root_node = self.resolve_root_node_address(root_node)
|
||||
os.environ['MASTER_ADDR'] = root_node
|
||||
|
||||
dist.init_process_group("nccl", rank=self.proc_rank, world_size=self.world_size)
|
||||
|
||||
def resolve_root_node_address(self, root_node):
|
||||
try:
|
||||
if '[' in root_node:
|
||||
name = root_node.split('[')[0]
|
||||
number = root_node.split(',')[0]
|
||||
if '-' in number:
|
||||
number = number.split('-')[0]
|
||||
if '[' in root_node:
|
||||
name = root_node.split('[')[0]
|
||||
number = root_node.split(',')[0]
|
||||
if '-' in number:
|
||||
number = number.split('-')[0]
|
||||
|
||||
number = re.sub('[^0-9]', '', number)
|
||||
root_node = name + number
|
||||
|
||||
except Exception as e:
|
||||
root_node = '127.0.0.2'
|
||||
number = re.sub('[^0-9]', '', number)
|
||||
root_node = name + number
|
||||
|
||||
return root_node
|
||||
|
||||
|
||||
@@ -0,0 +1,17 @@
|
||||
|
||||
def data_loader(fn):
|
||||
"""
|
||||
Decorator to make any fx with this use the lazy property
|
||||
:param fn:
|
||||
:return:
|
||||
"""
|
||||
|
||||
attr_name = '_lazy_' + fn.__name__
|
||||
|
||||
@property
|
||||
def _data_loader(self):
|
||||
if not hasattr(self, attr_name):
|
||||
setattr(self, attr_name, fn(self))
|
||||
return getattr(self, attr_name)
|
||||
|
||||
return _data_loader
|
||||
@@ -1,11 +1,9 @@
|
||||
import os
|
||||
import torch
|
||||
import math
|
||||
|
||||
from pytorch_lightning.root_module.memory import ModelSummary
|
||||
from pytorch_lightning.root_module.grads import GradInformation
|
||||
from pytorch_lightning.root_module.model_saving import ModelIO, load_hparams_from_tags_csv
|
||||
from pytorch_lightning.root_module.hooks import ModelHooks
|
||||
from pytorch_lightning.root_module.decorators import data_loader
|
||||
|
||||
|
||||
class LightningModule(GradInformation, ModelIO, ModelHooks):
|
||||
@@ -26,11 +24,6 @@ class LightningModule(GradInformation, ModelIO, ModelHooks):
|
||||
# track if gpu was requested for checkpointing
|
||||
self.on_gpu = False
|
||||
|
||||
# computed vars for the dataloaders
|
||||
self._tng_dataloader = None
|
||||
self._val_dataloader = None
|
||||
self._test_dataloader = None
|
||||
|
||||
def forward(self, *args, **kwargs):
|
||||
"""
|
||||
Expand model in into whatever you need.
|
||||
@@ -91,7 +84,7 @@ class LightningModule(GradInformation, ModelIO, ModelHooks):
|
||||
for param in self.parameters():
|
||||
param.requires_grad = True
|
||||
|
||||
@property
|
||||
@data_loader
|
||||
def tng_dataloader(self):
|
||||
"""
|
||||
Implement a function to load an h5py of this data
|
||||
@@ -99,7 +92,7 @@ class LightningModule(GradInformation, ModelIO, ModelHooks):
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@property
|
||||
@data_loader
|
||||
def test_dataloader(self):
|
||||
"""
|
||||
Implement a function to load an h5py of this data
|
||||
@@ -107,7 +100,7 @@ class LightningModule(GradInformation, ModelIO, ModelHooks):
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@property
|
||||
@data_loader
|
||||
def val_dataloader(self):
|
||||
"""
|
||||
Implement a function to load an h5py of this data
|
||||
@@ -142,3 +135,6 @@ class LightningModule(GradInformation, ModelIO, ModelHooks):
|
||||
model.load_model_specific(checkpoint)
|
||||
model.load_state_dict(checkpoint['state_dict'], strict=False)
|
||||
return model
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -11,6 +11,7 @@ from torch.utils.data import DataLoader
|
||||
from torch.utils.data.distributed import DistributedSampler
|
||||
|
||||
from pytorch_lightning.root_module.root_module import LightningModule
|
||||
import pytorch_lightning as ptl
|
||||
|
||||
|
||||
class LightningTestModel(LightningModule):
|
||||
@@ -217,35 +218,17 @@ class LightningTestModel(LightningModule):
|
||||
|
||||
return loader
|
||||
|
||||
@property
|
||||
@ptl.data_loader
|
||||
def tng_dataloader(self):
|
||||
if self._tng_dataloader is None:
|
||||
try:
|
||||
self._tng_dataloader = self.__dataloader(train=True)
|
||||
except Exception as e:
|
||||
print(e)
|
||||
raise e
|
||||
return self._tng_dataloader
|
||||
return self.__dataloader(train=True)
|
||||
|
||||
@property
|
||||
@ptl.data_loader
|
||||
def val_dataloader(self):
|
||||
if self._val_dataloader is None:
|
||||
try:
|
||||
self._val_dataloader = self.__dataloader(train=False)
|
||||
except Exception as e:
|
||||
print(e)
|
||||
raise e
|
||||
return self._val_dataloader
|
||||
return self.__dataloader(train=False)
|
||||
|
||||
@property
|
||||
@ptl.data_loader
|
||||
def test_dataloader(self):
|
||||
if self._test_dataloader is None:
|
||||
try:
|
||||
self._test_dataloader = self.__dataloader(train=False)
|
||||
except Exception as e:
|
||||
print(e)
|
||||
raise e
|
||||
return self._test_dataloader
|
||||
return self.__dataloader(train=False)
|
||||
|
||||
@staticmethod
|
||||
def add_model_specific_args(parent_parser, root_dir):
|
||||
|
||||
@@ -7,7 +7,7 @@ from setuptools import setup, find_packages
|
||||
# http://blog.ionelmc.ro/2014/05/25/python-packaging/
|
||||
setup(
|
||||
name="pytorch-lightning",
|
||||
version='0.3.5',
|
||||
version='0.3.51',
|
||||
description="The Keras for ML researchers using PyTorch",
|
||||
author="William Falcon",
|
||||
author_email="waf2107@columbia.edu",
|
||||
|
||||
+2
-13
@@ -91,7 +91,7 @@ def run_prediction(dataloader, trained_model):
|
||||
assert val_acc > 0.70, f'this model is expected to get > 0.7 in test set (it got {val_acc})'
|
||||
|
||||
|
||||
def mainasdf():
|
||||
def main():
|
||||
|
||||
save_dir = init_save_dir()
|
||||
model, hparams = get_model()
|
||||
@@ -111,7 +111,6 @@ def mainasdf():
|
||||
max_nb_epochs=1,
|
||||
gpus=[0, 1],
|
||||
distributed_backend='dp',
|
||||
use_amp=True
|
||||
)
|
||||
|
||||
result = trainer.fit(model)
|
||||
@@ -128,15 +127,5 @@ def mainasdf():
|
||||
clear_save_dir()
|
||||
|
||||
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
import subprocess
|
||||
import re
|
||||
|
||||
print('getting pid')
|
||||
command = "lsof -i :%s | awk '{print $2}'" % 12910
|
||||
pids = subprocess.check_output(command, shell=True)
|
||||
pids = pids.strip()
|
||||
|
||||
print(len(pids))
|
||||
main()
|
||||
|
||||
@@ -55,7 +55,9 @@ def test_amp_gpu_ddp_slurm_managed():
|
||||
warnings.warn('test_amp_gpu_ddp cannot run. Rerun on a node with 2+ GPUs to run this test')
|
||||
return
|
||||
|
||||
# simulate setting slurm flags
|
||||
os.environ['MASTER_PORT'] = str(np.random.randint(12000, 19000, 1)[0])
|
||||
os.environ['SLURM_LOCALID'] = str(0)
|
||||
|
||||
hparams = get_hparams()
|
||||
model = LightningTestModel(hparams)
|
||||
|
||||
Reference in New Issue
Block a user