Compare commits

..
157 Commits
Author SHA1 Message Date
William Falcon 9757841e67 release v0.2.5.2 2019-07-18 17:59:39 -04:00
William Falcon 0ac7a8590b added slurm managed flag catch for non-slurm peeps 2019-07-18 17:59:16 -04:00
William Falcon 6e12431e6b added slurm managed flag catch for non-slurm peeps 2019-07-18 17:58:38 -04:00
William Falcon c2e2298586 release v0.2.5.1 2019-07-18 17:14:34 -04:00
William Falcon 319feb7da5 removed printing. added auto process gen if slurm tasks do not match 2019-07-18 17:13:57 -04:00
William Falcon 5195124d4e added slurm no process warning 2019-07-18 17:06:56 -04:00
William Falcon 4e67983f23 added slurm no process warning 2019-07-18 17:05:09 -04:00
William Falcon c02b6c4c88 added slurm no process warning 2019-07-18 17:03:27 -04:00
William Falcon ad44d9168b added slurm no process warning 2019-07-18 16:47:46 -04:00
William Falcon 53a1a6d462 removed print lines 2019-07-18 16:37:48 -04:00
William Falcon 59d60eaf18 testing single process ddp 2019-07-18 15:06:20 -04:00
William Falcon 112be99b19 testing single process ddp 2019-07-18 14:57:56 -04:00
William Falcon 0e67773d2e testing single process ddp 2019-07-18 14:53:01 -04:00
William Falcon 394cdeeb8b added epoch flag back 2019-07-18 13:32:36 -04:00
William Falcon d0a8292e02 release v0.2.5 2019-07-18 12:13:00 -04:00
William Falcon d7409afed9 added arg docs 2019-07-18 12:11:59 -04:00
William Falcon f01cb63234 added arg docs 2019-07-18 12:10:07 -04:00
William Falcon 8be7480f31 added arg docs 2019-07-18 12:09:25 -04:00
William Falcon 751bc7c695 added arg docs 2019-07-18 12:08:47 -04:00
William Falcon 3be26dbb95 added arg docs 2019-07-18 12:08:17 -04:00
William Falcon 2ca0864ce8 added arg docs 2019-07-18 12:07:11 -04:00
William Falcon b1041220ac added arg docs 2019-07-18 12:05:52 -04:00
William Falcon da842c0cd6 added arg docs 2019-07-18 12:04:45 -04:00
William Falcon c4971e8432 added arg docs 2019-07-18 12:04:19 -04:00
William Falcon e81dbce38c set dp as default backend 2019-07-18 11:59:14 -04:00
William Falcon 0d992689d5 set dp as default backend 2019-07-18 11:58:27 -04:00
William Falcon 4085b3fa69 set dp as default backend 2019-07-18 11:57:39 -04:00
William Falcon b684bb55c5 set dp as default backend 2019-07-18 11:56:48 -04:00
William Falcon f98f88ff08 set dp as default backend 2019-07-18 11:51:43 -04:00
William Falcon f0955df4f0 set dp as default backend 2019-07-18 11:50:23 -04:00
William Falcon 7744c7117d set dp as default backend 2019-07-18 11:49:42 -04:00
William Falcon 22f4d6e26e set dp as default backend 2019-07-18 11:49:28 -04:00
William Falcon d49a83dec0 set dp as default backend 2019-07-18 11:48:16 -04:00
William Falcon 6d1d5ef68e set dp as default backend 2019-07-18 11:45:55 -04:00
William Falcon 3a1525222d set dp as default backend 2019-07-18 11:42:47 -04:00
William Falcon f650253cae set dp as default backend 2019-07-18 11:40:10 -04:00
William Falcon c67c84b443 set dp as default backend 2019-07-18 11:40:00 -04:00
William Falcon 4db32984c6 set dp as default backend 2019-07-18 11:39:13 -04:00
William Falcon 81d39786d9 set dp as default backend 2019-07-18 11:39:06 -04:00
William Falcon 63de076765 set dp as default backend 2019-07-18 11:36:48 -04:00
William Falcon 256ca62a3c set dp as default backend 2019-07-18 11:36:31 -04:00
William Falcon 39d04eb795 set dp as default backend 2019-07-18 11:35:59 -04:00
William Falcon e02857fcce set dp as default backend 2019-07-18 11:33:51 -04:00
William Falcon c163caf8cb set dp as default backend 2019-07-18 11:31:45 -04:00
William Falcon 2096a0aa84 set dp as default backend 2019-07-18 11:29:38 -04:00
William Falcon c253f96c53 set dp as default backend 2019-07-18 11:29:21 -04:00
William Falcon 551daca047 set dp as default backend 2019-07-18 11:25:02 -04:00
William Falcon ded0abead7 set dp as default backend 2019-07-18 11:21:35 -04:00
William Falcon e86b191691 set dp as default backend 2019-07-18 11:20:11 -04:00
William Falcon 3321e8c541 set dp as default backend 2019-07-18 11:18:19 -04:00
William Falcon bc3a805202 set dp as default backend 2019-07-18 11:16:16 -04:00
William Falcon 162b9f4f27 set dp as default backend 2019-07-18 11:15:21 -04:00
William Falcon e5bc3ea5b4 added training router 2019-07-18 11:09:37 -04:00
William Falcon baa139f97a added training router 2019-07-18 11:09:00 -04:00
William Falcon 470f3e6d29 added training router 2019-07-18 11:08:48 -04:00
William Falcon c12a0b57da added dp and ddp flag 2019-07-18 11:03:16 -04:00
William Falcon e7ecfa15f8 added option and flag 2019-07-18 10:56:45 -04:00
William Falcon 9051eb0039 updated docs 2019-07-17 15:56:55 -04:00
William Falcon bb8dbfca09 release v0.2.4.1 2019-07-17 10:04:14 -04:00
William Falcon 0240c70780 updated required deps 2019-07-17 10:03:58 -04:00
William Falcon a41abad5b2 Update trainer.py 2019-07-16 17:02:21 -04:00
William Falcon a83588b14e Update trainer.py 2019-07-16 13:12:56 -04:00
William Falcon 80192752b7 Merge pull request #13 from cinjon/on_tng_metrics
add a hook for on_tng_metrics so that users get access to the grad_no…
2019-07-16 12:59:16 -04:00
Cinjon Resnick fbd3873a0f add a hook for on_tng_metrics so that users get access to the grad_norm and mem_map dicts. 2019-07-16 12:51:48 -04:00
William Falcon 28cfddbe65 accept dist sampler classes 2019-07-16 12:44:58 -04:00
William Falcon b4bdb283ce release v0.2.4 2019-07-16 10:05:14 -04:00
William Falcon 967e57f071 early stop starts counting once min epochs met 2019-07-16 10:00:03 -04:00
William Falcon d12f6b7dd8 added summary flag 2019-07-15 21:11:29 -04:00
William Falcon 182c025c88 removed validation call 2019-07-15 20:48:46 -04:00
William Falcon 58e6199ce8 removed validation call 2019-07-15 14:56:56 -04:00
William Falcon 6a33f0d483 made early stop checkpoint optional 2019-07-15 14:54:38 -04:00
William Falcon dd230a93e8 made early stop checkpoint optional 2019-07-15 14:53:37 -04:00
William Falcon 3aa9cfc18e made checkpoint callback optional 2019-07-15 13:18:56 -04:00
William Falcon e57f461323 made checkpoint callback optional 2019-07-15 13:17:38 -04:00
William Falcon ab00514ef6 fixed metrics request not forced anymore 2019-07-15 13:03:08 -04:00
William Falcon 1dd58b4687 Merge branch 'master' of https://github.com/williamFalcon/pytorch-lightning 2019-07-15 13:01:17 -04:00
William Falcon b4b8a3dfde fixed none bug 2019-07-15 13:01:08 -04:00
William Falcon d8782c7b90 Update README.md 2019-07-15 09:21:30 -04:00
William Falcon d5878e9a72 release v0.2.3 2019-07-14 18:15:15 -04:00
William Falcon ad24bef1c9 removed print statements 2019-07-14 18:12:41 -04:00
William Falcon 50246a5066 working on single gpu init speed 2019-07-14 17:33:48 -04:00
William Falcon 21914cb1c1 working on single gpu init speed 2019-07-14 17:15:20 -04:00
William Falcon 904935cf98 working on single gpu init speed 2019-07-14 17:11:52 -04:00
William Falcon 468e75c180 working on single gpu init speed 2019-07-14 17:10:13 -04:00
William Falcon 849f52b7a6 modified single gpu init 2019-07-14 17:01:18 -04:00
William Falcon e520297781 modified single gpu init 2019-07-14 16:57:15 -04:00
William Falcon cefc27112d ddp flag change 2019-07-13 22:28:08 -04:00
William Falcon 6876f60098 merge 2019-07-13 22:21:17 -04:00
William Falcon fc1653e337 Merge branch 'nccl' of https://github.com/williamFalcon/pytorch-lightning into nccl 2019-07-13 22:19:41 -04:00
William Falcon e9f5913dac enabling gpu size = 1 to run without data parallel 2019-07-13 22:16:10 -04:00
William Falcon 7da82c2560 added fallback local init 2019-07-13 22:16:10 -04:00
William Falcon a2639c6894 added fallback local init 2019-07-13 22:16:10 -04:00
William Falcon eb05fa316f added fallback local init 2019-07-13 22:16:10 -04:00
William Falcon 6d55adb0d8 fixed nccl init 2019-07-13 22:16:10 -04:00
William Falcon cff0500a63 fixed nccl init 2019-07-13 22:16:10 -04:00
William Falcon f3ca184fb6 fixed nccl init 2019-07-13 22:16:10 -04:00
William Falcon 3239c9fdf8 fixed nccl init 2019-07-13 22:16:10 -04:00
William Falcon 7e37f68a5b fixed nccl init 2019-07-13 22:16:10 -04:00
William Falcon 960937ebe9 fixed nccl init 2019-07-13 22:16:10 -04:00
William Falcon a87784b4c5 fixed nccl init 2019-07-13 22:16:10 -04:00
William Falcon 5812efcf24 fixed nccl init 2019-07-13 22:16:10 -04:00
William Falcon e82014ec6c fixed nccl init 2019-07-13 22:16:10 -04:00
William Falcon dc87a4fc91 fixed nccl init 2019-07-13 22:16:10 -04:00
William Falcon 4f5eef2e78 fixed nccl init 2019-07-13 22:16:10 -04:00
William Falcon 6c02afefca fixed nccl init 2019-07-13 22:16:10 -04:00
William Falcon 4696e12641 fixed nccl init 2019-07-13 22:16:10 -04:00
William Falcon 7c688fbf2e enabling gpu size = 1 to run without data parallel 2019-07-13 22:09:17 -04:00
William Falcon 9ccfc7bd33 added fallback local init 2019-07-13 22:03:36 -04:00
William Falcon 52a98d76d8 added fallback local init 2019-07-13 10:16:50 -04:00
William Falcon 8b0cda84e7 added fallback local init 2019-07-13 10:13:52 -04:00
William Falcon 9f41a9e8b7 fixed nccl init 2019-07-12 16:35:20 -04:00
William Falcon b7baa96186 fixed nccl init 2019-07-12 16:29:44 -04:00
William Falcon faa2d4fa8b fixed nccl init 2019-07-12 16:23:20 -04:00
William Falcon 4f5da45fae fixed nccl init 2019-07-12 16:17:50 -04:00
William Falcon 7e54ad3f7c fixed nccl init 2019-07-12 16:16:46 -04:00
William Falcon 3bf366bcd8 fixed nccl init 2019-07-12 16:08:23 -04:00
William Falcon 6219f24a03 fixed nccl init 2019-07-12 16:07:57 -04:00
William Falcon 0bd81db538 fixed nccl init 2019-07-12 16:05:46 -04:00
William Falcon c84700814d fixed nccl init 2019-07-12 16:03:17 -04:00
William Falcon c244599ae8 fixed nccl init 2019-07-12 15:59:33 -04:00
William Falcon d99b121379 fixed nccl init 2019-07-12 15:59:12 -04:00
William Falcon 91b869d043 fixed nccl init 2019-07-12 15:55:28 -04:00
William Falcon 08e1ab64b5 fixed nccl init 2019-07-12 15:53:45 -04:00
William Falcon c1b21fb1e4 Merge pull request #11 from cinjon/modulefix
trainer: module fix.
2019-07-12 15:28:48 -04:00
William Falcon 8451bb7745 fixed nccl init 2019-07-12 15:25:34 -04:00
William Falcon 1a1771cfd8 fixed nccl init 2019-07-12 15:24:42 -04:00
William Falcon 1952e9be49 fixed nccl init 2019-07-12 15:11:32 -04:00
William Falcon 19391b1df1 fixed nccl init 2019-07-12 15:04:20 -04:00
William Falcon 369174c4d3 fixed nccl init 2019-07-12 14:36:00 -04:00
William Falcon 5ba0a2ed48 fixed nccl init 2019-07-12 14:28:49 -04:00
William Falcon 88061b2284 Merge branch 'master' of https://github.com/williamFalcon/pytorch-lightning 2019-07-12 13:42:53 -04:00
William Falcon ba38037917 fixed nccl init 2019-07-12 13:39:58 -04:00
William Falcon a7bb731a1d testing env init 2019-07-12 13:19:10 -04:00
William Falcon 58531888e0 testing env init 2019-07-12 13:17:33 -04:00
William Falcon 56ac885f03 Merge pull request #12 from cinjon/commafix
root_module: fix comma splits.
2019-07-12 13:13:57 -04:00
William Falcon 5e033fd97a testing env init 2019-07-12 13:11:08 -04:00
William Falcon 5d14b97aa6 testing file init 2019-07-12 12:57:54 -04:00
William Falcon 0b0addbcbe testing file init 2019-07-12 12:56:44 -04:00
Cinjon Resnick 098d518398 trainer: module fix. 2019-07-12 12:54:35 -04:00
William Falcon ba111e681e testing file init 2019-07-12 12:41:54 -04:00
Cinjon Resnick 3de053c903 root_module: fix comma splits. 2019-07-12 12:38:39 -04:00
William Falcon ac1bd57b8b testing file init 2019-07-12 12:33:54 -04:00
William Falcon 3f0fab9160 reset master 2019-07-12 12:32:36 -04:00
William Falcon 24c13aadc0 testing file init 2019-07-12 12:06:19 -04:00
William Falcon 885bad3555 testing master_Addr flag 2019-07-12 11:55:14 -04:00
William Falcon 6dde1d7ae3 testing master_Addr flag 2019-07-12 11:43:05 -04:00
William Falcon c223960edb testing master_Addr flag 2019-07-12 11:30:57 -04:00
William Falcon 32646cf2ee release v0.2.2 2019-07-11 16:19:11 -04:00
William Falcon 415ee4903b simplify trainer output 2019-07-11 15:23:33 -04:00
William Falcon a21dc5a187 simplify trainer output 2019-07-11 15:15:22 -04:00
William Falcon 0929908229 simplify trainer output 2019-07-11 15:08:45 -04:00
William Falcon cc12a1c8fa added clarifying comments 2019-07-11 14:58:47 -04:00
William Falcon 91b3a0aac6 added clarifying comments 2019-07-11 14:57:26 -04:00
William Falcon ed35f4e076 updated amp use 2019-07-11 14:35:41 -04:00
William Falcon c4781cb415 Merge branch 'master' of https://github.com/williamFalcon/pytorch-lightning 2019-07-11 14:18:07 -04:00
William Falcon 730a06640b updated amp use 2019-07-11 14:17:43 -04:00
William Falcon 6eb25edb31 release v0.21 2019-07-09 19:56:11 -04:00
12 changed files with 505 additions and 102 deletions
+2 -2
View File
@@ -116,8 +116,8 @@ def validation_end(self, outputs):
return tqdm_dic return tqdm_dic
``` ```
## TensorboardX ## Tensorboard
Lightning is fully integrated with tensorboardX. Lightning is fully integrated with tensorboard.
<p align="center"> <p align="center">
<a href="https://williamfalcon.github.io/pytorch-lightning/"> <a href="https://williamfalcon.github.io/pytorch-lightning/">
+1 -1
View File
@@ -19,7 +19,7 @@ Cut the learning rate by 10 at every epoch listed in this list.
trainer = Trainer(lr_scheduler_milestones=None) trainer = Trainer(lr_scheduler_milestones=None)
# cut LR by 10 at 100, 200, and 300 epochs # cut LR by 10 at 100, 200, and 300 epochs
trainer = Trainer(lr_scheduler_milestones=[100, 200, 300]) trainer = Trainer(lr_scheduler_milestones='100, 200, 300')
``` ```
--- ---
@@ -84,9 +84,10 @@ class LightningTemplateModel(LightningModule):
loss_val = self.loss(y, y_hat) loss_val = self.loss(y, y_hat)
output = OrderedDict({ output = OrderedDict({
'loss': loss_val, 'loss': loss_val
'tqdm_metrics': {}
}) })
# can also return just a scalar instead of a dict (return loss_val)
return output return output
def validation_step(self, data_batch, batch_i): def validation_step(self, data_batch, batch_i):
@@ -107,8 +108,10 @@ class LightningTemplateModel(LightningModule):
output = OrderedDict({ output = OrderedDict({
'val_loss': loss_val, 'val_loss': loss_val,
'val_acc': torch.tensor(val_acc), 'val_acc': torch.tensor(val_acc).cuda(loss_val.device.index),
}) })
# can also return just a scalar instead of a dict (return loss_val)
return output return output
def validation_end(self, outputs): def validation_end(self, outputs):
@@ -117,6 +120,10 @@ class LightningTemplateModel(LightningModule):
:param outputs: list of individual outputs of each validation step :param outputs: list of individual outputs of each validation step
:return: :return:
""" """
# if returned a scalar from validation_step, outputs is a list of tensor scalars
# we return just the average in this case (if we want)
# return torch.stack(outputs).mean()
val_loss_mean = 0 val_loss_mean = 0
val_acc_mean = 0 val_acc_mean = 0
for output in outputs: for output in outputs:
@@ -0,0 +1,112 @@
"""
Runs a model on a single node across N-gpus.
"""
import os
import sys
import numpy as np
from time import sleep
import torch
from test_tube import HyperOptArgumentParser, Experiment, SlurmCluster
from pytorch_lightning.models.trainer import Trainer
from pytorch_lightning.utils.arg_parse import add_default_args
from pytorch_lightning.callbacks import EarlyStopping, ModelCheckpoint
SEED = 2334
torch.manual_seed(SEED)
np.random.seed(SEED)
from lightning_module_template import LightningTemplateModel
def main(hparams):
"""
Main training routine specific for this project
:param hparams:
:return:
"""
# ------------------------
# 1 INIT LIGHTNING MODEL
# ------------------------
print('loading model...')
model = LightningTemplateModel(hparams)
print('model built')
# ------------------------
# 2 INIT TEST TUBE EXP
# ------------------------
# init experiment
exp = Experiment(
name=hyperparams.experiment_name,
save_dir=hyperparams.test_tube_save_path,
autosave=False,
description='test demo'
)
exp.argparse(hparams)
exp.save()
# ------------------------
# 3 DEFINE CALLBACKS
# ------------------------
model_save_path = '{}/{}/{}'.format(hparams.model_save_path, exp.name, exp.version)
early_stop = EarlyStopping(
monitor='val_acc',
patience=3,
verbose=True,
mode='max'
)
checkpoint = ModelCheckpoint(
filepath=model_save_path,
save_best_only=True,
verbose=True,
monitor='val_loss',
mode='min'
)
# ------------------------
# 4 INIT TRAINER
# ------------------------
trainer = Trainer(
experiment=exp,
checkpoint_callback=checkpoint,
early_stop_callback=early_stop,
gpus=hparams.gpus,
)
# ------------------------
# 5 START TRAINING
# ------------------------
trainer.fit(model)
if __name__ == '__main__':
# dirs
root_dir = os.path.dirname(os.path.realpath(__file__))
demo_log_dir = os.path.join(root_dir, 'pt_lightning_demo_logs')
checkpoint_dir = os.path.join(demo_log_dir, 'model_weights')
test_tube_dir = os.path.join(demo_log_dir, 'test_tube_data')
# although we user hyperOptParser, we are using it only as argparse right now
parent_parser = HyperOptArgumentParser(strategy='grid_search', add_help=False)
# gpu args
parent_parser.add_argument('--gpus', type=str, default='-1', help='how many gpus to use in the node. -1 uses all the gpus on the node')
parent_parser.add_argument('--test_tube_save_path', type=str, default=test_tube_dir, help='where to save logs')
parent_parser.add_argument('--model_save_path', type=str, default=checkpoint_dir, help='where to save model')
parent_parser.add_argument('--experiment_name', type=str, default='pt_lightning_exp_a', help='test tube exp name')
# allow model to overwrite or extend args
parser = LightningTemplateModel.add_model_specific_args(parent_parser, root_dir)
hyperparams = parser.parse_args()
# ---------------------
# RUN TRAINING
# ---------------------
# run on HPC cluster
print(f'RUNNING INTERACTIVE MODE ON GPUS. gpu ids: {hyperparams.gpus}')
main(hyperparams)
@@ -0,0 +1,112 @@
"""
Runs a model on a single node across N-gpus.
"""
import os
import sys
import numpy as np
from time import sleep
import torch
from test_tube import HyperOptArgumentParser, Experiment, SlurmCluster
from pytorch_lightning.models.trainer import Trainer
from pytorch_lightning.utils.arg_parse import add_default_args
from pytorch_lightning.callbacks import EarlyStopping, ModelCheckpoint
SEED = 2334
torch.manual_seed(SEED)
np.random.seed(SEED)
from lightning_module_template import LightningTemplateModel
def main(hparams):
"""
Main training routine specific for this project
:param hparams:
:return:
"""
# ------------------------
# 1 INIT LIGHTNING MODEL
# ------------------------
print('loading model...')
model = LightningTemplateModel(hparams)
print('model built')
# ------------------------
# 2 INIT TEST TUBE EXP
# ------------------------
# init experiment
exp = Experiment(
name=hyperparams.experiment_name,
save_dir=hyperparams.test_tube_save_path,
autosave=False,
description='test demo'
)
exp.argparse(hparams)
exp.save()
# ------------------------
# 3 DEFINE CALLBACKS
# ------------------------
model_save_path = '{}/{}/{}'.format(hparams.model_save_path, exp.name, exp.version)
early_stop = EarlyStopping(
monitor='val_acc',
patience=3,
verbose=True,
mode='max'
)
checkpoint = ModelCheckpoint(
filepath=model_save_path,
save_best_only=True,
verbose=True,
monitor='val_loss',
mode='min'
)
# ------------------------
# 4 INIT TRAINER
# ------------------------
trainer = Trainer(
experiment=exp,
checkpoint_callback=checkpoint,
early_stop_callback=early_stop,
gpus=hparams.gpus,
)
# ------------------------
# 5 START TRAINING
# ------------------------
trainer.fit(model)
if __name__ == '__main__':
# dirs
root_dir = os.path.dirname(os.path.realpath(__file__))
demo_log_dir = os.path.join(root_dir, 'pt_lightning_demo_logs')
checkpoint_dir = os.path.join(demo_log_dir, 'model_weights')
test_tube_dir = os.path.join(demo_log_dir, 'test_tube_data')
# although we user hyperOptParser, we are using it only as argparse right now
parent_parser = HyperOptArgumentParser(strategy='grid_search', add_help=False)
# gpu args
parent_parser.add_argument('--gpus', type=str, default='0', help='how many gpus to use in the node. -1 uses all the gpus on the node')
parent_parser.add_argument('--test_tube_save_path', type=str, default=test_tube_dir, help='where to save logs')
parent_parser.add_argument('--model_save_path', type=str, default=checkpoint_dir, help='where to save model')
parent_parser.add_argument('--experiment_name', type=str, default='pt_lightning_exp_a', help='test tube exp name')
# allow model to overwrite or extend args
parser = LightningTemplateModel.add_model_specific_args(parent_parser, root_dir)
hyperparams = parser.parse_args()
# ---------------------
# RUN TRAINING
# ---------------------
# run on HPC cluster
print(f'RUNNING INTERACTIVE MODE ON GPUS. gpu ids: {hyperparams.gpus}')
main(hyperparams)
+233 -88
View File
@@ -6,6 +6,7 @@ import subprocess
import traceback import traceback
import warnings import warnings
import os import os
import pdb
import torch import torch
from torch.utils.data.distributed import DistributedSampler from torch.utils.data.distributed import DistributedSampler
@@ -17,7 +18,7 @@ import tqdm
from pytorch_lightning.root_module.memory import get_gpu_memory_map from pytorch_lightning.root_module.memory import get_gpu_memory_map
from pytorch_lightning.root_module.model_saving import TrainerIO from pytorch_lightning.root_module.model_saving import TrainerIO
from pytorch_lightning.pt_overrides.override_data_parallel import LightningDistributedDataParallel from pytorch_lightning.pt_overrides.override_data_parallel import LightningDistributedDataParallel, LightningDataParallel
try: try:
@@ -27,11 +28,33 @@ except ModuleNotFoundError:
APEX_AVAILABLE = False APEX_AVAILABLE = False
def reduce_distributed_output(output, nb_gpus):
if nb_gpus <= 1:
return output
# when using DP, we get one output per gpu
# average outputs and return
if type(output) is torch.Tensor:
return output.mean()
for k, v in output.items():
# recurse on nested dics
if isinstance(output[k], dict):
output[k] = reduce_distributed_output(output[k], nb_gpus)
# reduce only metrics that have the same nb of gpus
elif output[k].size(0) == nb_gpus:
reduced = torch.mean(output[k])
output[k] = reduced
return output
class Trainer(TrainerIO): class Trainer(TrainerIO):
def __init__(self, def __init__(self,
experiment, experiment,
checkpoint_callback, early_stop_callback, early_stop_callback=None,
checkpoint_callback=None,
gradient_clip=0, gradient_clip=0,
cluster=None, cluster=None,
process_position=0, process_position=0,
@@ -44,20 +67,57 @@ class Trainer(TrainerIO):
check_val_every_n_epoch=1, check_val_every_n_epoch=1,
fast_dev_run=False, fast_dev_run=False,
accumulate_grad_batches=1, accumulate_grad_batches=1,
enable_early_stop=True, max_nb_epochs=1000, min_nb_epochs=1, max_nb_epochs=1000, min_nb_epochs=1,
train_percent_check=1.0, val_percent_check=1.0, test_percent_check=1.0, val_check_interval=0.95, train_percent_check=1.0, val_percent_check=1.0, test_percent_check=1.0,
val_check_interval=0.95,
log_save_interval=100, add_log_row_interval=10, log_save_interval=100, add_log_row_interval=10,
lr_scheduler_milestones=None, lr_scheduler_milestones=None,
distributed_backend='dp',
use_amp=False, use_amp=False,
print_nan_grads=False, print_nan_grads=False,
print_weights_summary=True,
amp_level='O2', amp_level='O2',
nb_sanity_val_steps=5): nb_sanity_val_steps=5):
"""
:param experiment: Test-tube experiment
:param early_stop_callback: from pytorch_lightning import EarlyStopping
:param checkpoint_callback: from pytorch_lightning import Checkpoint
:param gradient_clip:
:param cluster:
:param process_position:
:param current_gpu_name:
:param nb_gpu_nodes:
:param gpus:
:param progress_bar:
:param overfit_pct:
:param track_grad_norm:
:param check_val_every_n_epoch:
:param fast_dev_run:
:param accumulate_grad_batches:
:param max_nb_epochs:
:param min_nb_epochs:
:param train_percent_check:
:param val_percent_check:
:param test_percent_check:
:param val_check_interval:
:param log_save_interval:
:param add_log_row_interval:
:param lr_scheduler_milestones:
:param distributed_backend: 'np' to use DistributedParallel, 'ddp' to use DistributedDataParallel
:param use_amp:
:param print_nan_grads:
:param print_weights_summary:
:param amp_level:
:param nb_sanity_val_steps:
"""
# Transfer params # Transfer params
self.nb_gpu_nodes = nb_gpu_nodes self.nb_gpu_nodes = nb_gpu_nodes
self.gradient_clip = gradient_clip self.gradient_clip = gradient_clip
self.check_val_every_n_epoch = check_val_every_n_epoch self.check_val_every_n_epoch = check_val_every_n_epoch
self.enable_early_stop = enable_early_stop self.enable_early_stop = early_stop_callback is not None
self.track_grad_norm = track_grad_norm self.track_grad_norm = track_grad_norm
self.fast_dev_run = fast_dev_run self.fast_dev_run = fast_dev_run
self.on_gpu = gpus is not None and torch.cuda.is_available() self.on_gpu = gpus is not None and torch.cuda.is_available()
@@ -67,8 +127,12 @@ class Trainer(TrainerIO):
self.cluster = cluster self.cluster = cluster
self.process_position = process_position self.process_position = process_position
self.current_gpu_name = current_gpu_name self.current_gpu_name = current_gpu_name
self.print_weights_summary = print_weights_summary
self.checkpoint_callback = checkpoint_callback self.checkpoint_callback = checkpoint_callback
self.checkpoint_callback.save_function = self.save_checkpoint
if self.checkpoint_callback is not None:
self.checkpoint_callback.save_function = self.save_checkpoint
self.early_stop = early_stop_callback self.early_stop = early_stop_callback
self.model = None self.model = None
self.max_nb_epochs = max_nb_epochs self.max_nb_epochs = max_nb_epochs
@@ -82,6 +146,8 @@ class Trainer(TrainerIO):
self.print_nan_grads = print_nan_grads self.print_nan_grads = print_nan_grads
self.data_parallel_device_ids = None self.data_parallel_device_ids = None
self.world_size = 1 self.world_size = 1
self.use_ddp = False
self.use_dp = False
# 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
@@ -95,8 +161,14 @@ class Trainer(TrainerIO):
# set the correct cuda visible devices (using pci order) # set the correct cuda visible devices (using pci order)
os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID" os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID"
os.environ["CUDA_VISIBLE_DEVICES"] = ','.join([str(x) for x in self.data_parallel_device_ids]) os.environ["CUDA_VISIBLE_DEVICES"] = ','.join([str(x) for x in self.data_parallel_device_ids])
print(f'VISIBLE GPUS: {os.environ["CUDA_VISIBLE_DEVICES"]}')
self.data_parallel = self.data_parallel_device_ids is not None and len(self.data_parallel_device_ids) > 0 # make DP and DDP mutually exclusive
# single GPU will also use DP with devices=[0]
have_gpus = self.data_parallel_device_ids is not None and len(self.data_parallel_device_ids) > 0
if have_gpus:
self.use_dp = distributed_backend == 'dp'
self.use_ddp = distributed_backend == 'ddp'
# process info # process info
self.proc_rank = 0 self.proc_rank = 0
@@ -137,6 +209,10 @@ class Trainer(TrainerIO):
''' '''
warnings.warn(msg) warnings.warn(msg)
@property
def data_parallel(self):
return self.use_dp or self.use_ddp
def __determine_data_use_amount(self, train_percent_check, val_percent_check, test_percent_check, overfit_pct): def __determine_data_use_amount(self, train_percent_check, val_percent_check, test_percent_check, overfit_pct):
""" """
Use less data for debugging purposes Use less data for debugging purposes
@@ -149,8 +225,12 @@ class Trainer(TrainerIO):
self.val_percent_check = overfit_pct self.val_percent_check = overfit_pct
self.test_percent_check = overfit_pct self.test_percent_check = overfit_pct
def __get_model(self):
return self.model.module if self.data_parallel else self.model
def __is_function_implemented(self, f_name): def __is_function_implemented(self, f_name):
f_op = getattr(self.model, f_name, None) model = self.__get_model()
f_op = getattr(model, f_name, None)
return callable(f_op) return callable(f_op)
@property @property
@@ -208,9 +288,6 @@ class Trainer(TrainerIO):
:param max_batches: Scalar :param max_batches: Scalar
:return: :return:
""" """
if self.proc_rank == 0:
print('validating...')
# enable eval mode # enable eval mode
model.zero_grad() model.zero_grad()
model.eval() model.eval()
@@ -234,8 +311,12 @@ class Trainer(TrainerIO):
# ----------------- # -----------------
# RUN VALIDATION STEP # RUN VALIDATION STEP
# ----------------- # -----------------
if self.data_parallel: if self.use_ddp:
output = model(data_batch, batch_i) output = model(data_batch, batch_i)
elif self.use_dp:
output = model(data_batch, batch_i)
output = reduce_distributed_output(output, len(self.data_parallel_device_ids))
else: else:
output = model.validation_step(data_batch, batch_i) output = model.validation_step(data_batch, batch_i)
@@ -269,7 +350,7 @@ class Trainer(TrainerIO):
self.test_dataloader = model.test_dataloader self.test_dataloader = model.test_dataloader
self.val_dataloader = model.val_dataloader self.val_dataloader = model.val_dataloader
if self.on_gpu and type(self.tng_dataloader.sampler) is not DistributedSampler: if self.use_ddp and not isinstance(self.tng_dataloader.sampler, DistributedSampler):
msg = ''' msg = '''
when using multiple gpus and multiple nodes you must pass a DistributedSampler to DataLoader(sampler). when using multiple gpus and multiple nodes you must pass a DistributedSampler to DataLoader(sampler).
@@ -288,10 +369,65 @@ class Trainer(TrainerIO):
# MODEL TRAINING # MODEL TRAINING
# ----------------------------- # -----------------------------
def fit(self, model): def fit(self, model):
# when using multi-node or DDP within a node start each module in a separate process
if self.use_ddp:
# must copy only the meta of the exp so it survives pickle/unpickle when going to new process
self.experiment = self.experiment.get_meta_copy()
# whenever we have the correct number of tasks, we let slurm manage processes
# otherwise we launch the required number of processes
try:
nb_slurm_tasks = int(os.environ['SLURM_NTASKS'])
nb_requested_gpus = len(self.data_parallel_device_ids)
is_slurm_managing_tasks = nb_slurm_tasks == nb_requested_gpus
except Exception as e:
# likely not on slurm, so set the slurm managed flag to false
is_slurm_managing_tasks = False
if is_slurm_managing_tasks:
task = int(os.environ['SLURM_LOCALID'])
self.ddp_train(task, model)
else:
msg = f"""
You requested {nb_requested_gpus} GPUs but launched {nb_slurm_tasks} slurm tasks.
We will launch {nb_requested_gpus} processes for you.
We recommend you let slurm manage the processes by setting: --ntasks-per-node={nb_requested_gpus}
If you're not using SLURM, ignore this message!
"""
warnings.warn(msg)
mp.spawn(self.ddp_train, nprocs=len(self.data_parallel_device_ids), args=(model, ))
# 1 gpu or dp option triggers training using DP module
# easier to avoid NCCL issues
elif self.use_dp:
self.dp_train(model)
# 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:
# An example
model, optimizers = amp.initialize(
model, self.optimizers, opt_level=self.amp_level,
)
self.optimizers = optimizers
self.__run_pretrain_routine(model)
def dp_train(self, model):
# CHOOSE OPTIMIZER # CHOOSE OPTIMIZER
# filter out the weights that were done on gpu so we can load on good old cpus # filter out the weights that were done on gpu so we can load on good old cpus
self.optimizers = model.configure_optimizers() self.optimizers = model.configure_optimizers()
model.cuda(self.data_parallel_device_ids[0])
model = LightningDataParallel(model, device_ids=self.data_parallel_device_ids)
# run through amp wrapper # run through amp wrapper
if self.use_amp: if self.use_amp:
# An example # An example
@@ -300,15 +436,9 @@ class Trainer(TrainerIO):
) )
self.optimizers = optimizers self.optimizers = optimizers
# when using gpus, first thing we do is spawn a new process between each worker self.__run_pretrain_routine(model)
# applies to single gpu, multi-gpu and multi-nodes
if self.on_gpu:
self.experiment = self.experiment.get_meta_copy()
mp.spawn(self.dp_train, nprocs=len(self.data_parallel_device_ids), args=(model, ))
else:
self.__run_pretrain_routine(model)
def dp_train(self, gpu_nb, model): def ddp_train(self, gpu_nb, model):
""" """
Entry point into a DP thread Entry point into a DP thread
:param gpu_nb: :param gpu_nb:
@@ -324,6 +454,8 @@ class Trainer(TrainerIO):
node_rank = 0 node_rank = 0
# recover original exp before went into process # recover original exp before went into process
# init in write mode only on proc 0
self.experiment.debug = self.proc_rank > 0
self.experiment = self.experiment.get_non_ddp_exp() self.experiment = self.experiment.get_non_ddp_exp()
# show progbar only on prog_rank 0 # show progbar only on prog_rank 0
@@ -334,60 +466,56 @@ class Trainer(TrainerIO):
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)
# set up server using proc 0's ip address # set up server using proc 0's ip address
ip = self.__get_root_node_ip(self.proc_rank, self.nb_gpu_nodes) # try to init for 20 times at max in case ports are taken
dist.init_process_group("nccl", init_method=f'tcp://{ip}:12001', rank=self.proc_rank, world_size=self.world_size) # where to store ip_table
self.__init_tcp_connection()
# CHOOSE OPTIMIZER
# filter out the weights that were done on gpu so we can load on good old cpus
self.optimizers = model.configure_optimizers()
# MODEL
# copy model to each gpu # copy model to each gpu
torch.cuda.set_device(gpu_nb) torch.cuda.set_device(gpu_nb)
model.cuda(gpu_nb) model.cuda(gpu_nb)
# AMP
# run through amp wrapper before going to distributed DP
if self.use_amp:
# An example
model, optimizers = amp.initialize(
model, self.optimizers, opt_level=self.amp_level,
)
self.optimizers = optimizers
model = LightningDistributedDataParallel(model, device_ids=[gpu_nb]) model = LightningDistributedDataParallel(model, device_ids=[gpu_nb])
# continue training routine # continue training routine
self.__run_pretrain_routine(model) self.__run_pretrain_routine(model)
def __get_root_node_ip(self, world_gpu_nb, nb_gpu_nodes): def __init_tcp_connection(self):
""" """
Resolves the ip address of proc 0. Connect all procs in the world using the env:// init
Proc 0 writes address to a file. Every other process waits until the ip is available before it starts Use the first node as the root address
:param port:
:param world_gpu_nb: gpu number amongst all the world gpus :param tries:
:param nb_gpu_nodes:
:param ip_file_dir:
:return: :return:
""" """
# on one node we use localhost try:
if nb_gpu_nodes == 1: port = os.environ['MASTER_PORT']
return '127.0.0.1' except Exception as e:
port = 12910
os.environ['MASTER_PORT'] = f'{port}'
# where to store ip_table try:
ip_file_dir = os.path.join(self.cluster.log_path, 'ip_tables') root_node = os.environ['SLURM_NODELIST'].split(' ')[0]
except Exception as e:
root_node = '127.0.0.2'
# the first gpu in the world becomes the host os.environ['MASTER_ADDR'] = root_node
# this is based on its global rank
# it communicates its ip by saving an ip_table to the slurm cluster logging dir
# every other process waits for this ip to appear before continuing
ip_table_name = f'.ip_meta_' + os.environ['SLURM_JOB_ID']
ip_file = os.path.join(ip_file_dir, ip_table_name)
os.makedirs(ip_file_dir, exist_ok=True)
if world_gpu_nb == 0: sleep(self.proc_rank*0.5)
# get the proc 0 IP dist.init_process_group("nccl", rank=self.proc_rank, world_size=self.world_size)
root_ip = subprocess.run(['hostname', '-I'], stdout=subprocess.PIPE).stdout.decode('utf-8')
root_ip = root_ip.split(' ')[0]
# save the ip to the file
with open(file=ip_file, mode='w') as f:
f.write(root_ip)
return root_ip
else:
# wait up to 120 seconds until proc 0 writes
# once written, read proc 0's address and use it to configure server
for i in range(0, 120):
sleep(1.0)
if os.path.exists(ip_file):
ip = list(open(file=ip_file, mode='r'))[0]
return ip
def __run_pretrain_routine(self, model): def __run_pretrain_routine(self, model):
""" """
@@ -396,7 +524,7 @@ class Trainer(TrainerIO):
:return: :return:
""" """
ref_model = model ref_model = model
if self.on_gpu: if self.data_parallel:
ref_model = model.module ref_model = model.module
ref_model.trainer = self ref_model.trainer = self
@@ -417,7 +545,7 @@ class Trainer(TrainerIO):
self.lr_schedulers.append(scheduler) self.lr_schedulers.append(scheduler)
# print model summary # print model summary
if self.proc_rank == 0: if self.proc_rank == 0 and self.print_weights_summary:
ref_model.summarize() ref_model.summarize()
# give model convenience properties # give model convenience properties
@@ -448,12 +576,12 @@ class Trainer(TrainerIO):
for lr_scheduler in self.lr_schedulers: for lr_scheduler in self.lr_schedulers:
lr_scheduler.step() lr_scheduler.step()
model = self.model.module if self.data_parallel else self.model model = self.__get_model()
model.current_epoch = epoch_nb model.current_epoch = epoch_nb
# hook # hook
if self.__is_function_implemented('on_epoch_start'): if self.__is_function_implemented('on_epoch_start'):
model = self.model.module if self.data_parallel else self.model model = self.__get_model()
model.on_epoch_start() model.on_epoch_start()
self.current_epoch = epoch_nb self.current_epoch = epoch_nb
@@ -468,7 +596,7 @@ class Trainer(TrainerIO):
self.batch_nb = batch_nb self.batch_nb = batch_nb
self.global_step += 1 self.global_step += 1
model = self.model.module if self.data_parallel else self.model model = self.__get_model()
model.global_step = self.global_step model.global_step = self.global_step
# stop when the flag is changed or we've gone past the amount requested in the batches # stop when the flag is changed or we've gone past the amount requested in the batches
@@ -500,10 +628,8 @@ class Trainer(TrainerIO):
# count items in memory # count items in memory
# nb_params, nb_tensors = count_mem_items() # nb_params, nb_tensors = count_mem_items()
if self.data_parallel: model = self.__get_model()
metrics = self.model.module.update_tng_log_metrics(self.__tng_tqdm_dic) metrics = model.update_tng_log_metrics(self.__tng_tqdm_dic)
else:
metrics = self.model.update_tng_log_metrics(self.__tng_tqdm_dic)
# add gpu memory # add gpu memory
if self.on_gpu: if self.on_gpu:
@@ -512,11 +638,13 @@ class Trainer(TrainerIO):
# add norms # add norms
if self.track_grad_norm > 0: if self.track_grad_norm > 0:
model = self.model.module if self.data_parallel else self.model model = self.__get_model()
grad_norm_dic = model.grad_norm(self.track_grad_norm) grad_norm_dic = model.grad_norm(self.track_grad_norm)
metrics.update(grad_norm_dic) metrics.update(grad_norm_dic)
if self.__is_function_implemented('on_tng_metrics'):
model.on_tng_metrics(metrics)
# log metrics # log metrics
scalar_metrics = self.__metrics_to_scalars(metrics, blacklist=self.__log_vals_blacklist()) scalar_metrics = self.__metrics_to_scalars(metrics, blacklist=self.__log_vals_blacklist())
if self.proc_rank == 0: if self.proc_rank == 0:
@@ -525,7 +653,7 @@ class Trainer(TrainerIO):
# hook # hook
if self.__is_function_implemented('on_batch_end'): if self.__is_function_implemented('on_batch_end'):
model = self.model.module if self.data_parallel else self.model model = self.__get_model()
model.on_batch_end() model.on_batch_end()
# end epoch early # end epoch early
@@ -534,13 +662,13 @@ class Trainer(TrainerIO):
# hook # hook
if self.__is_function_implemented('on_epoch_end'): if self.__is_function_implemented('on_epoch_end'):
model = self.model.module if self.data_parallel else self.model model = self.__get_model()
model.on_epoch_end() model.on_epoch_end()
# early stopping # early stopping
if self.enable_early_stop: met_min_epochs = epoch_nb > self.min_nb_epochs
if self.enable_early_stop and met_min_epochs:
should_stop = self.early_stop_callback.on_epoch_end(epoch=epoch_nb, logs=self.__tng_tqdm_dic) should_stop = self.early_stop_callback.on_epoch_end(epoch=epoch_nb, logs=self.__tng_tqdm_dic)
met_min_epochs = epoch_nb > self.min_nb_epochs
# stop training # stop training
stop = should_stop and met_min_epochs stop = should_stop and met_min_epochs
@@ -563,7 +691,7 @@ class Trainer(TrainerIO):
def __log_vals_blacklist(self): def __log_vals_blacklist(self):
"""avoid logging some vals lightning uses to maintain state""" """avoid logging some vals lightning uses to maintain state"""
blacklist = {'batch_nb', 'v_nb', 'epoch', 'gpu'} blacklist = {'batch_nb', 'v_nb', 'gpu'}
return blacklist return blacklist
def __run_tng_batch(self, data_batch, batch_nb): def __run_tng_batch(self, data_batch, batch_nb):
@@ -572,8 +700,8 @@ class Trainer(TrainerIO):
# hook # hook
if self.__is_function_implemented('on_batch_start'): if self.__is_function_implemented('on_batch_start'):
model = self.model.module if self.data_parallel else self.model model_ref = self.__get_model()
response = model.on_batch_start(data_batch) response = model_ref.on_batch_start(data_batch)
if response == -1: if response == -1:
return -1 return -1
@@ -583,13 +711,26 @@ class Trainer(TrainerIO):
# forward pass # forward pass
# return a scalar value and a dic with tqdm metrics # return a scalar value and a dic with tqdm metrics
if self.data_parallel: if self.use_ddp:
output = self.model(data_batch, batch_nb) output = self.model(data_batch, batch_nb)
elif self.use_dp:
output = self.model(data_batch, batch_nb)
output = reduce_distributed_output(output, len(self.data_parallel_device_ids))
else: else:
output = self.model.training_step(data_batch, batch_nb) output = self.model.training_step(data_batch, batch_nb)
model_specific_tqdm_metrics_dic = output['tqdm_metrics'] try:
loss = output['loss'] model_specific_tqdm_metrics_dic = output['tqdm_metrics']
except Exception as e:
model_specific_tqdm_metrics_dic = {}
# if output dict doesn't have the keyword loss
# then assume the output=loss if scalar
try:
loss = output['loss']
except Exception as e:
if type(output) is torch.Tensor:
loss = output
self.__add_tqdm_metrics(model_specific_tqdm_metrics_dic) self.__add_tqdm_metrics(model_specific_tqdm_metrics_dic)
@@ -603,10 +744,11 @@ class Trainer(TrainerIO):
loss.backward() loss.backward()
if self.print_nan_grads: if self.print_nan_grads:
model = self.model.module if self.data_parallel else self.model model = self.__get_model()
for param in model.parameters(): for param in model.parameters():
print(param.grad.float().sum()) print(param.grad.float().sum())
# avoid memory leaks
self.batch_loss_value += loss.item() self.batch_loss_value += loss.item()
# gradient update with accumulated gradients # gradient update with accumulated gradients
@@ -614,7 +756,7 @@ class Trainer(TrainerIO):
# clip gradients # clip gradients
if self.gradient_clip > 0: if self.gradient_clip > 0:
model = self.model.module if self.data_parallel else self.model model = self.__get_model()
torch.nn.utils.clip_grad_norm(model.parameters(), self.gradient_clip) torch.nn.utils.clip_grad_norm(model.parameters(), self.gradient_clip)
# update gradients across all optimizers # update gradients across all optimizers
@@ -640,7 +782,8 @@ class Trainer(TrainerIO):
# activate batch end hook # activate batch end hook
if self.__is_function_implemented('on_batch_end'): if self.__is_function_implemented('on_batch_end'):
self.model.on_batch_end() model = self.__get_model()
model.on_batch_end()
return 0 return 0
@@ -655,7 +798,8 @@ class Trainer(TrainerIO):
try: try:
# hook # hook
if self.__is_function_implemented('on_pre_performance_check'): if self.__is_function_implemented('on_pre_performance_check'):
self.model.on_pre_performance_check() model = self.__get_model()
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
@@ -669,7 +813,8 @@ class Trainer(TrainerIO):
# hook # hook
if self.__is_function_implemented('on_post_performance_check'): if self.__is_function_implemented('on_post_performance_check'):
self.model.on_post_performance_check() model = self.__get_model()
model.on_post_performance_check()
except Exception as e: except Exception as e:
print(e) print(e)
@@ -681,6 +826,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: if self.proc_rank == 0 and self.checkpoint_callback:
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)
@@ -1,6 +1,7 @@
from torch.nn import DataParallel from torch.nn import DataParallel
from torch.nn.parallel import DistributedDataParallel from torch.nn.parallel import DistributedDataParallel
import itertools import itertools
from itertools import chain
import threading import threading
import torch import torch
@@ -42,6 +43,29 @@ class LightningDataParallel(DataParallel):
Override the forward call in lightning so it goes to training and validation step respectively Override the forward call in lightning so it goes to training and validation step respectively
""" """
def forward(self, *inputs, **kwargs):
if not self.device_ids:
return self.module(*inputs, **kwargs)
for t in chain(self.module.parameters(), self.module.buffers()):
if t.device != self.src_device_obj:
raise RuntimeError("module must have its parameters and buffers "
"on device {} (device_ids[0]) but found one of "
"them on device: {}".format(self.src_device_obj, t.device))
inputs, kwargs = self.scatter(inputs, kwargs, self.device_ids)
if len(self.device_ids) == 1:
# lightning
if self.module.training:
return self.module.training_step(*inputs[0], **kwargs[0])
else:
return self.module.validation_step(*inputs[0], **kwargs[0])
replicas = self.replicate(self.module, self.device_ids[:len(inputs)])
outputs = self.parallel_apply(replicas, inputs, kwargs)
return self.gather(outputs, self.output_device)
def parallel_apply(self, replicas, inputs, kwargs): def parallel_apply(self, replicas, inputs, kwargs):
return parallel_apply(replicas, inputs, kwargs, self.device_ids[:len(replicas)]) return parallel_apply(replicas, inputs, kwargs, self.device_ids[:len(replicas)])
+3
View File
@@ -19,3 +19,6 @@ class ModelHooks(torch.nn.Module):
def on_post_performance_check(self): def on_post_performance_check(self):
pass pass
def on_tng_metrics(self, metrics):
pass
@@ -2,7 +2,7 @@ import torch
import os import os
import re import re
import pdb import pdb
from pytorch_lightning.pt_overrides.override_data_parallel import LightningDistributedDataParallel from pytorch_lightning.pt_overrides.override_data_parallel import LightningDistributedDataParallel, LightningDataParallel
class ModelIO(object): class ModelIO(object):
@@ -66,7 +66,8 @@ class TrainerIO(object):
checkpoint['optimizer_states'] = optimizer_states checkpoint['optimizer_states'] = optimizer_states
# request what to save from the model # request what to save from the model
model = self.model.module if type(self.model) is LightningDistributedDataParallel else self.model 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
checkpoint_dict = model.get_save_dict() checkpoint_dict = model.get_save_dict()
# merge trainer and model saving items # merge trainer and model saving items
+2 -2
View File
@@ -78,7 +78,7 @@ class LightningModule(GradInformation, ModelIO, OptimizerConfig, ModelHooks):
:param logs: :param logs:
:return: :return:
""" """
raise NotImplementedError return logs
def loss(self, *args, **kwargs): def loss(self, *args, **kwargs):
""" """
@@ -129,7 +129,7 @@ class LightningModule(GradInformation, ModelIO, OptimizerConfig, ModelHooks):
def get_process_position(gpus): def get_process_position(gpus):
try: try:
current_gpu = os.environ["CUDA_VISIBLE_DEVICES"] current_gpu = os.environ["CUDA_VISIBLE_DEVICES"]
gpu_ids = gpus.split(';') gpu_ids = gpus.split(',')
process_position = gpu_ids.index(current_gpu) process_position = gpu_ids.index(current_gpu)
return process_position, current_gpu return process_position, current_gpu
except Exception as e: except Exception as e:
+1 -1
View File
@@ -1,4 +1,3 @@
from matplotlib import pyplot as plt
import numpy as np import numpy as np
np.seterr(divide='ignore', invalid='ignore') np.seterr(divide='ignore', invalid='ignore')
@@ -13,6 +12,7 @@ def plot_confusion_matrix(cm,
This function prints and plots the confusion matrix. This function prints and plots the confusion matrix.
Normalization can be applied by setting `normalize=True`. Normalization can be applied by setting `normalize=True`.
""" """
from matplotlib import pyplot as plt
if normalize: if normalize:
cm = cm.astype('float') / cm.sum(axis=1)[:, np.newaxis] cm = cm.astype('float') / cm.sum(axis=1)[:, np.newaxis]
print("Normalized confusion matrix") print("Normalized confusion matrix")
+2 -3
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.2', version='0.2.5.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,8 +19,7 @@ setup(
install_requires=[ install_requires=[
"torch>=1.1.0", "torch>=1.1.0",
"tqdm", "tqdm",
"test-tube>=0.653", "test-tube>=0.6.7.1",
"tensorflow>=1.14.0"
], ],
packages=find_packages(), packages=find_packages(),
long_description=open("README.md", encoding="utf-8").read(), long_description=open("README.md", encoding="utf-8").read(),