Compare commits

..
85 Commits
Author SHA1 Message Date
William Falcon 592fb4e5ba release v0.3.6.2 2019-07-26 23:08:51 -04:00
William Falcon 587c195298 added clean slurm save load test 2019-07-26 23:04:41 -04:00
William Falcon 64586f271d added clean slurm save load test 2019-07-26 23:02:18 -04:00
William Falcon 53b781709e added clean slurm save load test 2019-07-26 22:57:49 -04:00
William Falcon f183ac2a1c added clean slurm save load test 2019-07-26 22:51:33 -04:00
William Falcon 61c82611eb added clean slurm save load test 2019-07-26 22:40:07 -04:00
William Falcon 3224365190 added clean slurm save load test 2019-07-26 22:39:44 -04:00
William Falcon 2a4081e537 added clean slurm save load test 2019-07-26 22:33:31 -04:00
William Falcon 8e3a0443c7 added clean slurm save load test 2019-07-26 22:33:00 -04:00
William Falcon f5a01edfb8 added clean slurm save load test 2019-07-26 22:32:34 -04:00
William Falcon f1de62671d added clean slurm save load test 2019-07-26 22:32:27 -04:00
William Falcon 57edb08bd8 added clean slurm save load test 2019-07-26 22:28:09 -04:00
William Falcon ffa7a0dbab added clean slurm save load test 2019-07-26 22:26:55 -04:00
William Falcon b5419fcd8b added clean slurm save load test 2019-07-26 22:24:01 -04:00
William Falcon c61e13f0ff fixed hpc save, load. cleaned apu 2019-07-26 22:13:41 -04:00
William Falcon a6ae97ac09 fixed hpc save, load. cleaned apu 2019-07-26 22:13:06 -04:00
William Falcon 348223a702 fixed hpc save, load. cleaned apu 2019-07-26 22:09:35 -04:00
William Falcon 64de447545 fixed hpc save, load. cleaned apu 2019-07-26 22:07:02 -04:00
William Falcon 265411572f fixed hpc save, load. cleaned apu 2019-07-26 22:04:27 -04:00
William Falcon 4148c36abd added model save load test 2019-07-26 21:55:01 -04:00
William Falcon 0ee0344820 removed old template 2019-07-26 21:39:53 -04:00
William Falcon a5a80f35ec removed old template 2019-07-26 21:39:28 -04:00
William Falcon 92a1f559b5 remove state_dict 2019-07-26 21:39:01 -04:00
William Falcon aacf1947ea auto state-dict and remove the way the model is loaded during hpc 2019-07-26 21:38:06 -04:00
William Falcon e2c7fa44b7 auto state-dict and remove the way the model is loaded during hpc 2019-07-26 21:37:06 -04:00
William Falcon ff1ed9db7e release v0.3.6.1 2019-07-26 19:11:32 -04:00
William Falcon baf2ccefea Merge pull request #21 from williamFalcon/r
R
2019-07-26 19:10:48 -04:00
William Falcon df37c8418a updated test-tube dep number 2019-07-26 18:57:18 -04:00
William Falcon d6bfb94215 added global rank var name 2019-07-26 18:52:38 -04:00
William Falcon 12f717ad4a added global rank var name 2019-07-26 18:52:02 -04:00
William Falcon 56d41eaa8c Merge pull request #20 from williamFalcon/test2
Test2
2019-07-26 12:47:13 -04:00
William Falcon 84edf35f33 added saving tests to cpu 2019-07-26 12:35:28 -04:00
William Falcon a374a7ea00 added saving tests to cpu 2019-07-26 12:33:35 -04:00
William Falcon fbc1bbd161 added saving tests to cpu 2019-07-26 12:31:26 -04:00
William Falcon 84f03a1335 added saving tests to cpu 2019-07-26 12:29:19 -04:00
William Falcon 1a835969a6 added saving tests to cpu 2019-07-26 12:14:58 -04:00
William Falcon 2ee8f157ce added checkpoint test on cpu 2019-07-26 11:51:25 -04:00
William Falcon 51a5cc36e3 added checkpoint test on cpu 2019-07-26 11:50:02 -04:00
William Falcon c4b37d1efe updated readme 2019-07-25 20:13:22 -04:00
William Falcon 0489ed1e89 updated readme 2019-07-25 19:55:22 -04:00
William Falcon 677edc46d8 removed exception crashing from val 2019-07-25 19:49:45 -04:00
William Falcon 7e52f6ea97 cleaned up some if statements 2019-07-25 17:14:33 -04:00
William Falcon 7e728d97e7 removed save model logging 2019-07-25 14:36:22 -04:00
William Falcon 08bf9e16ae updated docs 2019-07-25 12:46:11 -04:00
William Falcon 7166b1acbc updated docs 2019-07-25 12:44:48 -04:00
William Falcon d18f38c0d7 updated docs 2019-07-25 12:40:09 -04:00
William Falcon f844f110af updated docs 2019-07-25 12:37:59 -04:00
William Falcon 1cbe54f8ba updated docs 2019-07-25 12:35:28 -04:00
William Falcon 79a79fb27d updated examples 2019-07-25 12:33:53 -04:00
William Falcon b1cd5d9d31 updated examples 2019-07-25 12:30:59 -04:00
William Falcon 6bd58de40e updated examples 2019-07-25 12:30:18 -04:00
William Falcon a1dd4d3e2c release v0.3.6 2019-07-25 12:22:50 -04:00
William Falcon b914866131 updated docs 2019-07-25 12:12:45 -04:00
William Falcon e182559c83 updated docs 2019-07-25 12:11:49 -04:00
William Falcon 9b99a02061 removed hparams req 2019-07-25 12:09:09 -04:00
William Falcon 20227b1382 removed hparams req 2019-07-25 12:08:00 -04:00
William Falcon d0d5653b06 removed hparams req 2019-07-25 12:04:20 -04:00
William Falcon b0d38d532d updated docs 2019-07-25 12:01:52 -04:00
William Falcon 4562580461 updated docs 2019-07-25 11:58:06 -04:00
William Falcon d272f29c88 updated docs 2019-07-25 11:52:54 -04:00
William Falcon 600c755460 updated docs 2019-07-25 11:44:25 -04:00
William Falcon d09a9e2c96 release v0.3.51 2019-07-25 11:38:57 -04:00
William Falcon 0f79e9d74e updated docs 2019-07-25 11:35:11 -04:00
William Falcon 9fa8120805 updated docs 2019-07-25 11:30:17 -04:00
William Falcon 715bf23105 updated docs 2019-07-25 11:28:34 -04:00
William Falcon 88ac4a0849 testing multiple calles 2019-07-25 11:19:58 -04:00
William Falcon 383746b87a testing multiple calles 2019-07-25 11:19:20 -04:00
William Falcon fffc09830f switched cpu amp order 2019-07-25 11:11:14 -04:00
William Falcon aadf8e16aa switched cpu amp order 2019-07-25 11:10:21 -04:00
William Falcon 4b04dc06d4 switched cpu amp order 2019-07-25 11:08:31 -04:00
William Falcon 0e42d28415 fixed root node addr 2019-07-25 11:05:15 -04:00
William Falcon 09dba13cde updated test models with lazy decorators 2019-07-25 11:01:08 -04:00
William Falcon 42a45bb273 updated test models with lazy decorators 2019-07-25 11:00:35 -04:00
William Falcon 5604e955eb updated test models with lazy decorators 2019-07-25 10:59:10 -04:00
William Falcon 6d34224e68 updated test models with lazy decorators 2019-07-25 10:56:42 -04:00
William Falcon 24a3246bc1 updated test models with lazy decorators 2019-07-25 10:56:03 -04:00
William Falcon 39b15855ed added lazy decorator 2019-07-25 10:39:48 -04:00
William Falcon c6da6eb46c updated readme 2019-07-25 10:33:35 -04:00
William Falcon d23d25646a cleaned readme 2019-07-25 10:26:47 -04:00
William Falcon bd6521a584 cleaned readme 2019-07-25 10:25:41 -04:00
William Falcon 2ce3e3e108 cleaned readme 2019-07-25 10:25:12 -04:00
William Falcon deeb82d28f cleaned readme 2019-07-25 10:23:51 -04:00
William Falcon 74817c2fb1 cleaned readme 2019-07-25 10:11:51 -04:00
William Falcon b989358c9b added downloads badge 2019-07-25 09:55:50 -04:00
William Falcon 0d47561a31 added downloads badge 2019-07-25 09:55:30 -04:00
15 changed files with 565 additions and 476 deletions
+92 -19
View File
@@ -9,8 +9,10 @@
<p align="center"> <p align="center">
The Keras for ML researchers using PyTorch. More control. Less boilerplate. The Keras for ML researchers using PyTorch. More control. Less boilerplate.
</p> </p>
<p align="center"> <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://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>
<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://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://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> <a href="https://williamfalcon.github.io/pytorch-lightning/"><img src="https://readthedocs.org/projects/pytorch-lightning/badge/?version=latest"></a>
@@ -25,17 +27,85 @@ pip install pytorch-lightning
**[View the docs here](https://williamfalcon.github.io/pytorch-lightning/)** **[View the docs here](https://williamfalcon.github.io/pytorch-lightning/)**
## What is it? ## 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? ## Why do I want to use lightning?
Because you don't want to define a training loop, validation loop, gradient clipping, checkpointing, loading, 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.
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. 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!
## How do I do use it?
To use lightning do 2 things: 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). 1. [Define a LightningModel](https://williamfalcon.github.io/pytorch-lightning/LightningModule/RequiredTrainerInterface/)
2. [Define a LightningModel](https://github.com/williamFalcon/pytorch-lightning/blob/master/examples/new_project_templates/lightning_module_template.py). ```python
import pytorch_lightning as ptl
import torch
from torch.nn import functional as F
from torch.utils.data import DataLoader
from torchvision.datasets import MNIST
class CoolModel(ptl.LightningModule):
def __init(self):
super(CoolModel, self).__init__()
# not the best model...
self.l1 = torch.nn.Linear(28 * 28, 10)
def forward(self, x):
return torch.relu(self.l1(x))
def my_loss(self, y_hat, y):
return F.cross_entropy(y_hat, y)
def training_step(self, batch, batch_nb):
x, y = batch
y_hat = self.forward(x)
return {'tng_loss': self.my_loss(y_hat, y)}
def validation_step(self, batch, batch_nb):
x, y = batch
y_hat = self.forward(x)
return {'val_loss': self.my_loss(y_hat, y)}
def validation_end(self, outputs):
avg_loss = torch.stack([x for x in outputs['val_loss']]).mean()
return avg_loss
def configure_optimizers(self):
return [torch.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? ## What does lightning control for me?
Everything! Everything!
@@ -150,7 +220,7 @@ And run tensorboard from that dir
tensorboard --logdir /some/path tensorboard --logdir /some/path
``` ```
## Lightning automatically automates all of the following ([each is also configurable](https://williamfalcon.github.io/pytorch-lightning/Trainer/)): ## Lightning automates all of the following ([each is also configurable](https://williamfalcon.github.io/pytorch-lightning/Trainer/)):
###### Checkpointing ###### Checkpointing
@@ -217,19 +287,22 @@ pip install pytorch-lightning
# clone lightning for the demo # clone lightning for the demo
git clone https://github.com/williamFalcon/pytorch-lightning.git git clone https://github.com/williamFalcon/pytorch-lightning.git
cd examples/new_project_templates/ cd pytorch_lightning/examples/new_project_templates/
# run demo (on cpu) # all of the following demos use the SAME model to show no modification needs to be made to your code
python trainer_gpu_cluster_template.py
# train on cpu
python single_cpu_template.py
# train on multiple-gpus
python single_gpu_node_template.py --gpus "0,1"
# train on 32 gpus on a cluster (run on a SLURM managed cluster)
python multi_node_cluster_template.py --nb_gpu_nodes 4 --gpus '0,1,2,3,4,5,6,7'
``` ```
Without changing the model AT ALL, you can run the model on a single gpu, over multiple gpus, or over multiple nodes. ## Bleeding edge
If you can't wait for the next release, install the most up to date code with:
```bash ```bash
# run a grid search on two gpus pip install git+https://github.com/williamFalcon/pytorch-lightning.git@master --upgrade
python fully_featured_trainer.py --gpus "0;1" ```
# run single model on multiple gpus
python fully_featured_trainer.py --gpus "0;1" --interactive
```
@@ -26,6 +26,58 @@ Otherwise, to Define a Lightning Module, implement the following methods:
- [update_tng_log_metrics](RequiredTrainerInterface.md#update_tng_log_metrics) - [update_tng_log_metrics](RequiredTrainerInterface.md#update_tng_log_metrics)
- [add_model_specific_args](RequiredTrainerInterface.md#add_model_specific_args) - [add_model_specific_args](RequiredTrainerInterface.md#add_model_specific_args)
---
**Minimal example**
```python
import pytorch_lightning as ptl
import torch
from torch.nn import functional as F
from torch.utils.data import DataLoader
from torchvision.datasets import MNIST
class CoolModel(ptl.LightningModule):
def __init(self):
super(CoolModel, self).__init__()
# not the best model...
self.l1 = torch.nn.Linear(28 * 28, 10)
def forward(self, x):
return torch.relu(self.l1(x))
def my_loss(self, y_hat, y):
return F.cross_entropy(y_hat, y)
def training_step(self, batch, batch_nb):
x, y = batch
y_hat = self.forward(x)
return {'tng_loss': self.my_loss(y_hat, y)}
def validation_step(self, batch, batch_nb):
x, y = batch
y_hat = self.forward(x)
return {'val_loss': self.my_loss(y_hat, y)}
def validation_end(self, outputs):
avg_loss = torch.stack([x for x in outputs['val_loss']]).mean()
return avg_loss
def configure_optimizers(self):
return [torch.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)
```
--- ---
### training_step ### training_step
@@ -237,10 +289,10 @@ def load_model_specific(self, checkpoint):
### tng_dataloader ### tng_dataloader
``` {.python} ``` {.python}
@property @ptl.data_loader
def tng_dataloader(self) 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 ##### Return
Pytorch DataLoader Pytorch DataLoader
@@ -248,32 +300,26 @@ Pytorch DataLoader
**Example** **Example**
``` {.python} ``` {.python}
@property @ptl.data_loader
def tng_dataloader(self): def tng_dataloader(self):
if self._tng_dataloader is None: transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,), (1.0,))])
try: dataset = MNIST(root='/path/to/mnist/', train=True, transform=transform, download=True)
transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,), (1.0,))]) loader = torch.utils.data.DataLoader(
dataset = MNIST(root='/path/to/mnist/', train=True, transform=transform, download=True) dataset=dataset,
loader = torch.utils.data.DataLoader( batch_size=self.hparams.batch_size,
dataset=dataset, shuffle=True
batch_size=self.hparams.batch_size, )
shuffle=True return loader
)
self._tng_dataloader = loader
except Exception as e:
raise e
return self._tng_dataloader
``` ```
--- ---
### val_dataloader ### val_dataloader
``` {.python} ``` {.python}
@property @ptl.data_loader
def tng_dataloader(self) 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 ##### Return
Pytorch DataLoader Pytorch DataLoader
@@ -281,32 +327,27 @@ Pytorch DataLoader
**Example** **Example**
``` {.python} ``` {.python}
@property @ptl.data_loader
def val_dataloader(self): def val_dataloader(self):
if self._val_dataloader is None: transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,), (1.0,))])
try: dataset = MNIST(root='/path/to/mnist/', train=False, transform=transform, download=True)
transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,), (1.0,))]) loader = torch.utils.data.DataLoader(
dataset = MNIST(root='/path/to/mnist/', train=False, transform=transform, download=True) dataset=dataset,
loader = torch.utils.data.DataLoader( batch_size=self.hparams.batch_size,
dataset=dataset, shuffle=True
batch_size=self.hparams.batch_size, )
shuffle=True
) return loader
self._val_dataloader = loader
except Exception as e:
raise e
return self._val_dataloader
``` ```
--- ---
### test_dataloader ### test_dataloader
``` {.python} ``` {.python}
@property @ptl.data_loader
def test_dataloader(self) 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 ##### Return
Pytorch DataLoader Pytorch DataLoader
@@ -314,22 +355,17 @@ Pytorch DataLoader
**Example** **Example**
``` {.python} ``` {.python}
@property @ptl.data_loader
def test_dataloader(self): def test_dataloader(self):
if self._test_dataloader is None: transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,), (1.0,))])
try: dataset = MNIST(root='/path/to/mnist/', train=False, transform=transform, download=True)
transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,), (1.0,))]) loader = torch.utils.data.DataLoader(
dataset = MNIST(root='/path/to/mnist/', train=False, transform=transform, download=True) dataset=dataset,
loader = torch.utils.data.DataLoader( batch_size=self.hparams.batch_size,
dataset=dataset, shuffle=True
batch_size=self.hparams.batch_size, )
shuffle=True
) return loader
self._test_dataloader = loader
except Exception as e:
raise e
return self._test_dataloader
``` ```
--- ---
+5 -4
View File
@@ -1,10 +1,11 @@
###### New project Quick Start ###### New project Quick Start
To start a new project define these two files. To start a new project define these two files.
1. [Define a LightningModule](/LightningModule/RequiredTrainerInterface/#template-model-definition) 1. [Define a LightningModule](/pytorch-lightning/LightningModule/RequiredTrainerInterface/)
2. Pick a trainer 2. [Define a trainer](https://williamfalcon.github.io/pytorch-lightning/Trainer/)
- [Basic CPU Trainer](https://github.com/williamFalcon/pytorch-lightning/blob/master/examples/new_project_templates/trainer_cpu_template.py) - [Basic CPU Trainer Template](https://github.com/williamFalcon/pytorch-lightning/blob/master/pytorch_lightning/examples/new_project_templates/single_cpu_template.py)
- [GPU cluster Trainer](https://github.com/williamFalcon/pytorch-lightning/blob/master/examples/new_project_templates/trainer_gpu_cluster_template.py) - [Multi-GPU Trainer Template](https://github.com/williamFalcon/pytorch-lightning/blob/master/pytorch_lightning/examples/new_project_templates/single_gpu_node_template.py)
- [GPU cluster Trainer Template](https://github.com/williamFalcon/pytorch-lightning/blob/master/pytorch_lightning/examples/new_project_templates/multi_node_cluster_template.py)
###### Docs shortcuts ###### Docs shortcuts
- [LightningModule](LightningModule/RequiredTrainerInterface/) - [LightningModule](LightningModule/RequiredTrainerInterface/)
+2 -1
View File
@@ -1,2 +1,3 @@
from .models import Trainer 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 import DataLoader
from torch.utils.data.distributed import DistributedSampler from torch.utils.data.distributed import DistributedSampler
import pytorch_lightning as ptl
from pytorch_lightning.root_module.root_module import LightningModule from pytorch_lightning.root_module.root_module import LightningModule
@@ -24,7 +25,8 @@ class LightningTemplateModel(LightningModule):
:param hparams: :param hparams:
""" """
# init superclass # init superclass
super(LightningTemplateModel, self).__init__(hparams) super(LightningTemplateModel, self).__init__()
self.hparams = hparams
self.batch_size = hparams.batch_size self.batch_size = hparams.batch_size
@@ -152,17 +154,6 @@ class LightningTemplateModel(LightningModule):
tqdm_dic = {'val_loss': val_loss_mean.item(), 'val_acc': val_acc_mean.item()} tqdm_dic = {'val_loss': val_loss_mean.item(), 'val_acc': val_acc_mean.item()}
return tqdm_dic return tqdm_dic
# ---------------------
# MODEL SAVING
# ---------------------
def get_save_dict(self):
checkpoint = {'state_dict': self.state_dict()}
return checkpoint
def load_model_specific(self, checkpoint):
self.load_state_dict(checkpoint['state_dict'])
pass
# --------------------- # ---------------------
# TRAINING SETUP # TRAINING SETUP
# --------------------- # ---------------------
@@ -200,35 +191,20 @@ class LightningTemplateModel(LightningModule):
return loader return loader
@property @ptl.data_loader
def tng_dataloader(self): def tng_dataloader(self):
if self._tng_dataloader is None: print('tng data loader called')
try: return self.__dataloader(train=True)
self._tng_dataloader = self.__dataloader(train=True)
except Exception as e:
print(e)
raise e
return self._tng_dataloader
@property @ptl.data_loader
def val_dataloader(self): def val_dataloader(self):
if self._val_dataloader is None: print('val data loader called')
try: return self.__dataloader(train=False)
self._val_dataloader = self.__dataloader(train=False)
except Exception as e:
print(e)
raise e
return self._val_dataloader
@property @ptl.data_loader
def test_dataloader(self): def test_dataloader(self):
if self._test_dataloader is None: print('test data loader called')
try: return self.__dataloader(train=False)
self._test_dataloader = self.__dataloader(train=False)
except Exception as e:
print(e)
raise e
return self._test_dataloader
@staticmethod @staticmethod
def add_model_specific_args(parent_parser, root_dir): # pragma: no cover def add_model_specific_args(parent_parser, root_dir): # pragma: no cover
@@ -1,203 +0,0 @@
import torch.nn as nn
import numpy as np
from pytorch_lightning.root_module.root_module import LightningModule
from test_tube import HyperOptArgumentParser
from torchvision.datasets import MNIST
import torchvision.transforms as transforms
import torch
import torch.nn.functional as F
class ExampleModel1(LightningModule):
"""
Sample model to show how to define a template
"""
def __init__(self, hparams):
# init superclass
super(ExampleModel1, self).__init__(hparams)
self.batch_size = hparams.batch_size
# build model
self.__build_model()
# ---------------------
# MODEL SETUP
# ---------------------
def __build_model(self):
"""
Layout model
:return:
"""
self.c_d1 = nn.Linear(in_features=self.hparams.in_features, out_features=self.hparams.hidden_dim)
self.c_d1_bn = nn.BatchNorm1d(self.hparams.hidden_dim)
self.c_d1_drop = nn.Dropout(self.hparams.drop_prob)
self.c_d2 = nn.Linear(in_features=self.hparams.hidden_dim, out_features=self.hparams.out_features)
# ---------------------
# TRAINING
# ---------------------
def forward(self, x):
x = self.c_d1(x)
x = F.tanh(x)
x = self.c_d1_bn(x)
x = self.c_d1_drop(x)
x = self.c_d2(x)
logits = F.log_softmax(x, dim=1)
return logits
def loss(self, labels, logits):
nll = F.nll_loss(logits, labels)
return nll
def training_step(self, data_batch):
"""
Called inside the training loop
:param data_batch:
:return:
"""
# forward pass
x, y = data_batch
x = x.view(x.size(0), -1)
y_hat = self.forward(x)
# calculate loss
loss_val = self.loss(y, y_hat)
tqdm_dic = {'jefe': 1}
return loss_val, tqdm_dic
def validation_step(self, data_batch):
"""
Called inside the validation loop
:param data_batch:
:return:
"""
x, y = data_batch
x = x.view(x.size(0), -1)
y_hat = self.forward(x)
loss_val = self.loss(y, y_hat)
# acc
labels_hat = torch.argmax(y_hat, dim=1)
val_acc = torch.sum(y == labels_hat).item() / (len(y) * 1.0)
output = {'y_hat': y_hat, 'val_loss': loss_val.item(), 'val_acc': val_acc}
return output
def validation_end(self, outputs):
"""
Called at the end of validation to aggregate outputs
:param outputs: list of individual outputs of each validation step
:return:
"""
val_loss_mean = 0
accs = []
for output in outputs:
val_loss_mean += output['val_loss']
accs.append(output['val_acc'])
val_loss_mean /= len(outputs)
tqdm_dic = {'val_loss': val_loss_mean, 'val_acc': np.mean(accs)}
return tqdm_dic
def update_tng_log_metrics(self, logs):
return logs
# ---------------------
# MODEL SAVING
# ---------------------
def get_save_dict(self):
checkpoint = {
'state_dict': self.state_dict(),
}
return checkpoint
def load_model_specific(self, checkpoint):
self.load_state_dict(checkpoint['state_dict'])
pass
# ---------------------
# TRAINING SETUP
# ---------------------
def configure_optimizers(self):
"""
return whatever optimizers we want here
:return: list of optimizers
"""
optimizer = self.choose_optimizer(self.hparams.optimizer_name, self.parameters(), {'lr': self.hparams.learning_rate}, 'optimizer')
self.optimizers = [optimizer]
return self.optimizers
def __dataloader(self, train):
# init data generators
transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,), (1.0,))])
dataset = MNIST(root=self.hparams.data_root, train=train, transform=transform, download=True)
loader = torch.utils.data.DataLoader(
dataset=dataset,
batch_size=self.hparams.batch_size,
shuffle=True
)
return loader
@property
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
@property
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
@property
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
@staticmethod
def add_model_specific_args(parent_parser):
parser = HyperOptArgumentParser(strategy=parent_parser.strategy, parents=[parent_parser])
# param overwrites
# parser.set_defaults(gradient_clip=5.0)
# network params
parser.opt_list('--drop_prob', default=0.2, options=[0.2, 0.5], type=float, tunable=False)
parser.add_argument('--in_features', default=28*28)
parser.add_argument('--hidden_dim', default=500)
parser.add_argument('--out_features', default=10)
# data
parser.add_argument('--data_root', default='/Users/williamfalcon/Developer/personal/research_lib/research_proj/datasets/mnist', type=str)
# training params (opt)
parser.opt_list('--learning_rate', default=0.001, type=float, options=[0.0001, 0.0005, 0.001, 0.005],
tunable=False)
parser.opt_list('--batch_size', default=256, type=int, options=[32, 64, 128, 256], tunable=False)
parser.opt_list('--optimizer_name', default='adam', type=str, options=['adam'], tunable=False)
return parser
+42 -41
View File
@@ -161,7 +161,6 @@ class Trainer(TrainerIO):
self.nb_tng_batches = None self.nb_tng_batches = None
self.nb_test_batches = None self.nb_test_batches = None
# gpus come in as a string. # gpus come in as a string.
# if gpus = -1 then use all available devices # if gpus = -1 then use all available devices
# otherwise, split the string using commas # otherwise, split the string using commas
@@ -438,14 +437,14 @@ class Trainer(TrainerIO):
# ON CPU # ON CPU
else: 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 # run through amp wrapper
if self.use_amp: if self.use_amp:
raise MisconfigurationException('amp + cpu is not supported. Please use a GPU option') 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) self.__run_pretrain_routine(model)
# return 1 when finished # return 1 when finished
@@ -500,6 +499,9 @@ class Trainer(TrainerIO):
self.proc_rank = self.node_rank * len(self.data_parallel_device_ids) + gpu_nb self.proc_rank = self.node_rank * len(self.data_parallel_device_ids) + gpu_nb
self.world_size = self.nb_gpu_nodes * len(self.data_parallel_device_ids) self.world_size = self.nb_gpu_nodes * len(self.data_parallel_device_ids)
# let the exp know the rank to avoid overwriting logs
self.experiment.rank = self.proc_rank
# set up server using proc 0's ip address # set up server using proc 0's ip address
# try to init for 20 times at max in case ports are taken # try to init for 20 times at max in case ports are taken
# where to store ip_table # where to store ip_table
@@ -544,25 +546,25 @@ class Trainer(TrainerIO):
os.environ['MASTER_PORT'] = f'{port}' os.environ['MASTER_PORT'] = f'{port}'
# figure out the root node addr # 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) root_node = self.resolve_root_node_address(root_node)
os.environ['MASTER_ADDR'] = root_node os.environ['MASTER_ADDR'] = root_node
dist.init_process_group("nccl", rank=self.proc_rank, world_size=self.world_size) dist.init_process_group("nccl", rank=self.proc_rank, world_size=self.world_size)
def resolve_root_node_address(self, root_node): def resolve_root_node_address(self, root_node):
try: if '[' in root_node:
if '[' in root_node: name = root_node.split('[')[0]
name = root_node.split('[')[0] number = root_node.split(',')[0]
number = root_node.split(',')[0] if '-' in number:
if '-' in number: number = number.split('-')[0]
number = number.split('-')[0]
number = re.sub('[^0-9]', '', number) number = re.sub('[^0-9]', '', number)
root_node = name + number root_node = name + number
except Exception as e:
root_node = '127.0.0.2'
return root_node return root_node
@@ -608,14 +610,18 @@ class Trainer(TrainerIO):
if self.proc_rank == 0: if self.proc_rank == 0:
self.experiment.save() self.experiment.save()
# track model now.
# if cluster resets state, the model will update with the saved weights
self.model = model
# enable cluster checkpointing # enable cluster checkpointing
# also restores training state
if self.cluster is not None: # pragma: no cover if self.cluster is not None: # pragma: no cover
self.enable_auto_hpc_walltime_manager() self.enable_auto_hpc_walltime_manager()
# --------------------------- # ---------------------------
# CORE TRAINING LOOP # CORE TRAINING LOOP
# --------------------------- # ---------------------------
self.model = model
self.__train() self.__train()
def __train(self): def __train(self):
@@ -854,30 +860,25 @@ class Trainer(TrainerIO):
elif not can_check_epoch: elif not can_check_epoch:
return return
try: # hook
# hook if self.__is_function_implemented('on_pre_performance_check'):
if self.__is_function_implemented('on_pre_performance_check'): model = self.__get_model()
model = self.__get_model() model.on_pre_performance_check()
model.on_pre_performance_check()
# use full val set on end of epoch # use full val set on end of epoch
# use a small portion otherwise # use a small portion otherwise
max_batches = None if not self.fast_dev_run else 1 max_batches = None if not self.fast_dev_run else 1
model_specific_tqdm_metrics_dic = self.validate( model_specific_tqdm_metrics_dic = self.validate(
self.model, self.model,
self.val_dataloader, self.val_dataloader,
max_batches max_batches
) )
self.__add_tqdm_metrics(model_specific_tqdm_metrics_dic) self.__add_tqdm_metrics(model_specific_tqdm_metrics_dic)
# hook # hook
if self.__is_function_implemented('on_post_performance_check'): if self.__is_function_implemented('on_post_performance_check'):
model = self.__get_model() model = self.__get_model()
model.on_post_performance_check() model.on_post_performance_check()
except Exception as e:
print(e)
print(traceback.print_exc())
if self.progress_bar: if self.progress_bar:
# add model specific metrics # add model specific metrics
@@ -885,6 +886,6 @@ class Trainer(TrainerIO):
self.prog_bar.set_postfix(**tqdm_metrics) self.prog_bar.set_postfix(**tqdm_metrics)
# model checkpointing # model checkpointing
if self.proc_rank == 0 and self.checkpoint_callback: if self.proc_rank == 0 and self.checkpoint_callback is not None:
print('save callback...') print('save callback...')
self.checkpoint_callback.on_epoch_end(epoch=self.current_epoch, logs=self.__tng_tqdm_dic) self.checkpoint_callback.on_epoch_end(epoch=self.current_epoch, logs=self.__tng_tqdm_dic)
@@ -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
+26 -21
View File
@@ -4,34 +4,36 @@ import re
import pdb import pdb
from pytorch_lightning.pt_overrides.override_data_parallel import LightningDistributedDataParallel, LightningDataParallel from pytorch_lightning.pt_overrides.override_data_parallel import LightningDistributedDataParallel, LightningDataParallel
class ModelIO(object): class ModelIO(object):
def load_model_specific(self, checkpoint): def on_load_checkpoint(self, checkpoint):
""" """
Do something with the checkpoint Do something with the checkpoint
Gives model a chance to load something before state_dict is restored
:param checkpoint: :param checkpoint:
:return: :return:
""" """
raise NotImplementedError pass
def get_save_dict(self): def on_save_checkpoint(self, checkpoint):
""" """
Return specific things for the model Give the model a chance to add something to the checkpoint.
:return: state_dict is already there
""" """
raise NotImplementedError pass
# ------------------------- # -------------------------
# OPTIONAL HOOKS # OPTIONAL HOOKS
# ------------------------- # -------------------------
def on_hpc_save(self): def on_hpc_save(self, checkpoint):
""" """
Hook to do whatever you need right before Slurm manager saves the model Hook to do whatever you need right before Slurm manager saves the model
:return: :return:
""" """
pass pass
def on_hpc_load(self): def on_hpc_load(self, checkpoint):
""" """
Hook to do whatever you need right before Slurm manager loads the model Hook to do whatever you need right before Slurm manager loads the model
:return: :return:
@@ -42,7 +44,6 @@ class ModelIO(object):
class TrainerIO(object): class TrainerIO(object):
def __get_model(self): def __get_model(self):
print(type(self.model))
is_dp_module = type(self.model) is LightningDistributedDataParallel or type(self.model) is LightningDataParallel is_dp_module = type(self.model) is LightningDistributedDataParallel or type(self.model) is LightningDataParallel
model = self.model.module if is_dp_module else self.model model = self.model.module if is_dp_module else self.model
return model return model
@@ -76,18 +77,19 @@ class TrainerIO(object):
checkpoint['optimizer_states'] = optimizer_states checkpoint['optimizer_states'] = optimizer_states
# request what to save from the model # add the state_dict from the model
model = self.__get_model() model = self.__get_model()
checkpoint_dict = model.get_save_dict() checkpoint['state_dict'] = model.state_dict()
# give the model a chance to add a few things
model.on_save_checkpoint(checkpoint)
# merge trainer and model saving items
checkpoint.update(checkpoint_dict)
return checkpoint return checkpoint
# -------------------- # --------------------
# HPC IO # HPC IO
# -------------------- # --------------------
def enable_auto_hpc_walltime_manager(self): # pragma: no cover def enable_auto_hpc_walltime_manager(self):
if self.cluster is None: if self.cluster is None:
return return
@@ -150,13 +152,14 @@ class TrainerIO(object):
# give model a chance to do something on hpc_save # give model a chance to do something on hpc_save
model = self.__get_model() model = self.__get_model()
model.on_hpc_save() checkpoint = self.dump_checkpoint()
# request what to save from the model model.on_hpc_save(checkpoint)
checkpoint_dict = self.dump_checkpoint()
# do the actual save # do the actual save
torch.save(checkpoint_dict, filepath) torch.save(checkpoint, filepath)
return filepath
def hpc_load(self, folderpath, on_gpu): def hpc_load(self, folderpath, on_gpu):
filepath = '{}/hpc_ckpt_{}.ckpt'.format(folderpath, self.max_ckpt_in_folder(folderpath)) filepath = '{}/hpc_ckpt_{}.ckpt'.format(folderpath, self.max_ckpt_in_folder(folderpath))
@@ -166,15 +169,17 @@ class TrainerIO(object):
else: else:
checkpoint = torch.load(filepath, map_location=lambda storage, loc: storage) checkpoint = torch.load(filepath, map_location=lambda storage, loc: storage)
# load training state # load training state (affects trainer only)
self.restore_training_state(checkpoint) self.restore_training_state(checkpoint)
# load model state # load model state
model = self.__get_model() model = self.__get_model()
model.load_model_specific(checkpoint)
# load the state_dict on the model automatically
model.load_state_dict(checkpoint['state_dict'])
# call model hook # call model hook
model.on_hpc_load() model.on_hpc_load(checkpoint)
def max_ckpt_in_folder(self, path): def max_ckpt_in_folder(self, path):
files = os.listdir(path) files = os.listdir(path)
+26 -37
View File
@@ -1,18 +1,15 @@
import os
import torch import torch
import math
from pytorch_lightning.root_module.memory import ModelSummary from pytorch_lightning.root_module.memory import ModelSummary
from pytorch_lightning.root_module.grads import GradInformation 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.model_saving import ModelIO, load_hparams_from_tags_csv
from pytorch_lightning.root_module.hooks import ModelHooks from pytorch_lightning.root_module.hooks import ModelHooks
from pytorch_lightning.root_module.decorators import data_loader
class LightningModule(GradInformation, ModelIO, ModelHooks): class LightningModule(GradInformation, ModelIO, ModelHooks):
def __init__(self, hparams): def __init__(self, *args, **kwargs):
super(LightningModule, self).__init__() super(LightningModule, self).__init__(*args, **kwargs)
self.hparams = hparams
self.dtype = torch.FloatTensor self.dtype = torch.FloatTensor
self.exp_save_path = None self.exp_save_path = None
@@ -26,11 +23,6 @@ class LightningModule(GradInformation, ModelIO, ModelHooks):
# track if gpu was requested for checkpointing # track if gpu was requested for checkpointing
self.on_gpu = False 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): def forward(self, *args, **kwargs):
""" """
Expand model in into whatever you need. Expand model in into whatever you need.
@@ -71,27 +63,7 @@ class LightningModule(GradInformation, ModelIO, ModelHooks):
""" """
raise NotImplementedError raise NotImplementedError
def loss(self, *args, **kwargs): @data_loader
"""
Expand model_out into your components
:param model_out:
:return:
"""
raise NotImplementedError
def summarize(self):
model_summary = ModelSummary(self)
print(model_summary)
def freeze(self):
for param in self.parameters():
param.requires_grad = False
def unfreeze(self):
for param in self.parameters():
param.requires_grad = True
@property
def tng_dataloader(self): def tng_dataloader(self):
""" """
Implement a function to load an h5py of this data Implement a function to load an h5py of this data
@@ -99,7 +71,7 @@ class LightningModule(GradInformation, ModelIO, ModelHooks):
""" """
raise NotImplementedError raise NotImplementedError
@property @data_loader
def test_dataloader(self): def test_dataloader(self):
""" """
Implement a function to load an h5py of this data Implement a function to load an h5py of this data
@@ -107,7 +79,7 @@ class LightningModule(GradInformation, ModelIO, ModelHooks):
""" """
raise NotImplementedError raise NotImplementedError
@property @data_loader
def val_dataloader(self): def val_dataloader(self):
""" """
Implement a function to load an h5py of this data Implement a function to load an h5py of this data
@@ -136,9 +108,26 @@ class LightningModule(GradInformation, ModelIO, ModelHooks):
else: else:
checkpoint = torch.load(weights_path, map_location=lambda storage, loc: storage) checkpoint = torch.load(weights_path, map_location=lambda storage, loc: storage)
# load the state_dict on the model automatically
model = cls(hparams) model = cls(hparams)
model.load_state_dict(checkpoint['state_dict'])
# give model a chance to load something
model.on_load_checkpoint(checkpoint)
# allow model to load
model.load_model_specific(checkpoint)
model.load_state_dict(checkpoint['state_dict'], strict=False)
return model return model
def summarize(self):
model_summary = ModelSummary(self)
print(model_summary)
def freeze(self):
for param in self.parameters():
param.requires_grad = False
def unfreeze(self):
for param in self.parameters():
param.requires_grad = True
@@ -11,6 +11,7 @@ from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler from torch.utils.data.distributed import DistributedSampler
from pytorch_lightning.root_module.root_module import LightningModule from pytorch_lightning.root_module.root_module import LightningModule
import pytorch_lightning as ptl
class LightningTestModel(LightningModule): class LightningTestModel(LightningModule):
@@ -24,7 +25,8 @@ class LightningTestModel(LightningModule):
:param hparams: :param hparams:
""" """
# init superclass # init superclass
super(LightningTestModel, self).__init__(hparams) super(LightningTestModel, self).__init__()
self.hparams = hparams
self.batch_size = hparams.batch_size self.batch_size = hparams.batch_size
@@ -169,17 +171,6 @@ class LightningTestModel(LightningModule):
def on_tng_metrics(self, logs): def on_tng_metrics(self, logs):
logs['some_tensor_to_test'] = torch.rand(1) logs['some_tensor_to_test'] = torch.rand(1)
# ---------------------
# MODEL SAVING
# ---------------------
def get_save_dict(self):
checkpoint = {'state_dict': self.state_dict()}
return checkpoint
def load_model_specific(self, checkpoint):
self.load_state_dict(checkpoint['state_dict'])
pass
# --------------------- # ---------------------
# TRAINING SETUP # TRAINING SETUP
# --------------------- # ---------------------
@@ -217,35 +208,17 @@ class LightningTestModel(LightningModule):
return loader return loader
@property @ptl.data_loader
def tng_dataloader(self): def tng_dataloader(self):
if self._tng_dataloader is None: return self.__dataloader(train=True)
try:
self._tng_dataloader = self.__dataloader(train=True)
except Exception as e:
print(e)
raise e
return self._tng_dataloader
@property @ptl.data_loader
def val_dataloader(self): def val_dataloader(self):
if self._val_dataloader is None: return self.__dataloader(train=False)
try:
self._val_dataloader = self.__dataloader(train=False)
except Exception as e:
print(e)
raise e
return self._val_dataloader
@property @ptl.data_loader
def test_dataloader(self): def test_dataloader(self):
if self._test_dataloader is None: return self.__dataloader(train=False)
try:
self._test_dataloader = self.__dataloader(train=False)
except Exception as e:
print(e)
raise e
return self._test_dataloader
@staticmethod @staticmethod
def add_model_specific_args(parent_parser, root_dir): def add_model_specific_args(parent_parser, root_dir):
+2 -2
View File
@@ -7,7 +7,7 @@ from setuptools import setup, find_packages
# http://blog.ionelmc.ro/2014/05/25/python-packaging/ # http://blog.ionelmc.ro/2014/05/25/python-packaging/
setup( setup(
name="pytorch-lightning", name="pytorch-lightning",
version='0.3.5', version='0.3.6.2',
description="The Keras for ML researchers using PyTorch", description="The Keras for ML researchers using PyTorch",
author="William Falcon", author="William Falcon",
author_email="waf2107@columbia.edu", author_email="waf2107@columbia.edu",
@@ -19,7 +19,7 @@ setup(
install_requires=[ install_requires=[
"torch>=1.1.0", "torch>=1.1.0",
"tqdm", "tqdm",
"test-tube>=0.6.7.1", "test-tube>=0.6.7.4",
], ],
packages=find_packages(), packages=find_packages(),
long_description=open("README.md", encoding="utf-8").read(), long_description=open("README.md", encoding="utf-8").read(),
+53 -15
View File
@@ -11,6 +11,55 @@ import os
import shutil import shutil
import pdb import pdb
import pytorch_lightning as ptl
import torch
from torch.nn import functional as F
from torch.utils.data import DataLoader
from torchvision.datasets import MNIST
class CoolModel(ptl.LightningModule):
def __init(self):
super(CoolModel, self).__init__()
# not the best model...
self.l1 = torch.nn.Linear(28 * 28, 10)
def forward(self, x):
return torch.relu(self.l1(x))
def my_loss(self, y_hat, y):
return F.cross_entropy(y_hat, y)
def training_step(self, batch, batch_nb):
x, y = batch
y_hat = self.forward(x)
return {'tng_loss': self.my_loss(y_hat, y)}
def validation_step(self, batch, batch_nb):
x, y = batch
y_hat = self.forward(x)
return {'val_loss': self.my_loss(y_hat, y)}
def validation_end(self, outputs):
avg_loss = torch.stack([x for x in outputs['val_loss']]).mean()
return avg_loss
def configure_optimizers(self):
return [torch.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)
def get_model(): def get_model():
# set up model with these hyperparams # set up model with these hyperparams
@@ -91,14 +140,12 @@ 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})' 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() save_dir = init_save_dir()
model, hparams = get_model()
# exp file to get meta # exp file to get meta
exp = get_exp(False) exp = get_exp(False)
exp.argparse(hparams)
exp.save() exp.save()
# exp file to get weights # exp file to get weights
@@ -111,9 +158,10 @@ def mainasdf():
max_nb_epochs=1, max_nb_epochs=1,
gpus=[0, 1], gpus=[0, 1],
distributed_backend='dp', distributed_backend='dp',
use_amp=True
) )
model = CoolModel()
result = trainer.fit(model) result = trainer.fit(model)
# correct result and ok accuracy # correct result and ok accuracy
@@ -128,15 +176,5 @@ def mainasdf():
clear_save_dir() clear_save_dir()
if __name__ == '__main__': if __name__ == '__main__':
import subprocess main()
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))
+192 -10
View File
@@ -3,16 +3,18 @@ from pytorch_lightning import Trainer
from pytorch_lightning.examples.new_project_templates.lightning_module_template import LightningTemplateModel from pytorch_lightning.examples.new_project_templates.lightning_module_template import LightningTemplateModel
from pytorch_lightning.testing_models.lm_test_module import LightningTestModel from pytorch_lightning.testing_models.lm_test_module import LightningTestModel
from argparse import Namespace from argparse import Namespace
from test_tube import Experiment from test_tube import Experiment, SlurmCluster
from pytorch_lightning.callbacks import ModelCheckpoint, EarlyStopping from pytorch_lightning.callbacks import ModelCheckpoint, EarlyStopping
from pytorch_lightning.utils.debugging import MisconfigurationException from pytorch_lightning.utils.debugging import MisconfigurationException
from pytorch_lightning.root_module import memory from pytorch_lightning.root_module import memory
from pytorch_lightning.models.trainer import reduce_distributed_output from pytorch_lightning.models.trainer import reduce_distributed_output
from pytorch_lightning.root_module import model_saving
import numpy as np import numpy as np
import warnings import warnings
import torch import torch
import os import os
import shutil import shutil
import pdb
SEED = 2334 SEED = 2334
torch.manual_seed(SEED) torch.manual_seed(SEED)
@@ -22,6 +24,108 @@ np.random.seed(SEED)
# ------------------------------------------------------------------------ # ------------------------------------------------------------------------
# TESTS # TESTS
# ------------------------------------------------------------------------ # ------------------------------------------------------------------------
def test_cpu_slurm_save_load():
"""
Verify model save/load/checkpoint on CPU
:return:
"""
hparams = get_hparams()
model = LightningTestModel(hparams)
save_dir = init_save_dir()
# exp file to get meta
exp = get_exp(False)
exp.argparse(hparams)
exp.save()
cluster_a = SlurmCluster()
trainer_options = dict(
max_nb_epochs=1,
cluster=cluster_a,
experiment=exp,
checkpoint_callback=ModelCheckpoint(save_dir)
)
# fit model
trainer = Trainer(**trainer_options)
result = trainer.fit(model)
real_global_step = trainer.global_step
# traning complete
assert result == 1, 'amp + ddp model failed to complete'
# predict with trained model before saving
# make a prediction
for batch in model.test_dataloader:
break
x, y = batch
x = x.view(x.size(0), -1)
model.eval()
pred_before_saving = model(x)
# test registering a save function
trainer.enable_auto_hpc_walltime_manager()
# test HPC saving
# simulate snapshot on slurm
saved_filepath = trainer.hpc_save(save_dir, exp)
assert os.path.exists(saved_filepath)
# wipe-out trainer and model
# retrain with not much data... this simulates picking training back up after slurm
# we want to see if the weights come back correctly
continue_tng_hparams = get_hparams(continue_training=True, hpc_exp_number=cluster_a.hpc_exp_number)
trainer_options = dict(
max_nb_epochs=1,
cluster=SlurmCluster(continue_tng_hparams),
experiment=exp,
checkpoint_callback=ModelCheckpoint(save_dir),
)
trainer = Trainer(**trainer_options)
model = LightningTestModel(hparams)
# set the epoch start hook so we can predict before the model does the full training
def assert_pred_same():
assert trainer.global_step == real_global_step and trainer.global_step > 0
# predict with loaded model to make sure answers are the same
trainer.model.eval()
new_pred = trainer.model(x)
assert torch.all(torch.eq(pred_before_saving, new_pred)).item() == 1
model.on_epoch_start = assert_pred_same
# by calling fit again, we trigger training, loading weights from the cluster
# and our hook to predict using current model before any more weight updates
trainer.fit(model)
clear_save_dir()
def test_loading_meta_tags():
hparams = get_hparams()
save_dir = init_save_dir()
# save tags
exp = get_exp(False)
exp.tag({'some_str':'a_str', 'an_int': 1, 'a_float': 2.0})
exp.argparse(hparams)
exp.save()
# load tags
tags_path = exp.get_data_path(exp.name, exp.version) + '/meta_tags.csv'
tags = model_saving.load_hparams_from_tags_csv(tags_path)
assert tags.batch_size == 32 and tags.hidden_dim == 1000
clear_save_dir()
def test_dp_output_reduce(): def test_dp_output_reduce():
# test identity when we have a single gpu # test identity when we have a single gpu
@@ -43,6 +147,74 @@ def test_dp_output_reduce():
assert reduced['b']['c'] == out['b']['c'] assert reduced['b']['c'] == out['b']['c']
def test_model_saving_loading():
"""
Tests use case where trainer saves the model, and user loads it from tags independently
:return:
"""
hparams = get_hparams()
model = LightningTestModel(hparams)
save_dir = init_save_dir()
# exp file to get meta
exp = get_exp(False)
exp.argparse(hparams)
exp.save()
trainer_options = dict(
max_nb_epochs=1,
cluster=SlurmCluster(),
experiment=exp,
checkpoint_callback=ModelCheckpoint(save_dir)
)
# fit model
trainer = Trainer(**trainer_options)
result = trainer.fit(model)
# traning complete
assert result == 1, 'amp + ddp model failed to complete'
# make a prediction
for batch in model.test_dataloader:
break
x, y = batch
x = x.view(x.size(0), -1)
# generate preds before saving model
model.eval()
pred_before_saving = model(x)
# save model
new_weights_path = os.path.join(save_dir, 'save_test.ckpt')
trainer.save_checkpoint(new_weights_path)
# load new model
tags_path = exp.get_data_path(exp.name, exp.version)
tags_path = os.path.join(tags_path, 'meta_tags.csv')
model_2 = LightningTestModel.load_from_metrics(weights_path=new_weights_path, tags_csv=tags_path, on_gpu=False)
model_2.eval()
# make prediction
# assert that both predictions are the same
new_pred = model_2(x)
assert torch.all(torch.eq(pred_before_saving, new_pred)).item() == 1
clear_save_dir()
def test_model_freeze_unfreeze():
hparams = get_hparams()
model = LightningTestModel(hparams)
model.freeze()
model.unfreeze()
def test_amp_gpu_ddp_slurm_managed(): def test_amp_gpu_ddp_slurm_managed():
""" """
Make sure DDP + AMP work Make sure DDP + AMP work
@@ -55,7 +227,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') warnings.warn('test_amp_gpu_ddp cannot run. Rerun on a node with 2+ GPUs to run this test')
return return
# simulate setting slurm flags
os.environ['MASTER_PORT'] = str(np.random.randint(12000, 19000, 1)[0]) os.environ['MASTER_PORT'] = str(np.random.randint(12000, 19000, 1)[0])
os.environ['SLURM_LOCALID'] = str(0)
hparams = get_hparams() hparams = get_hparams()
model = LightningTestModel(hparams) model = LightningTestModel(hparams)
@@ -409,16 +583,24 @@ def run_gpu_model_test(trainer_options, model, hparams, on_gpu=True):
clear_save_dir() clear_save_dir()
def get_hparams(): def get_hparams(continue_training=False, hpc_exp_number=0):
root_dir = os.path.dirname(os.path.realpath(__file__)) root_dir = os.path.dirname(os.path.realpath(__file__))
hparams = Namespace(**{'drop_prob': 0.2,
'batch_size': 32, args = {
'in_features': 28*28, 'drop_prob': 0.2,
'learning_rate': 0.001*8, 'batch_size': 32,
'optimizer_name': 'adam', 'in_features': 28*28,
'data_root': os.path.join(root_dir, 'mnist'), 'learning_rate': 0.001*8,
'out_features': 10, 'optimizer_name': 'adam',
'hidden_dim': 1000}) 'data_root': os.path.join(root_dir, 'mnist'),
'out_features': 10,
'hidden_dim': 1000}
if continue_training:
args['test_tube_do_checkpoint_load'] = True
args['hpc_exp_number'] = hpc_exp_number
hparams = Namespace(**args)
return hparams return hparams