mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-12 12:40:20 +08:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
9757841e67 | ||
|
|
0ac7a8590b | ||
|
|
6e12431e6b | ||
|
|
c2e2298586 | ||
|
|
319feb7da5 | ||
|
|
5195124d4e | ||
|
|
4e67983f23 | ||
|
|
c02b6c4c88 | ||
|
|
ad44d9168b | ||
|
|
53a1a6d462 | ||
|
|
59d60eaf18 | ||
|
|
112be99b19 | ||
|
|
0e67773d2e | ||
|
|
394cdeeb8b | ||
|
|
d0a8292e02 | ||
|
|
d7409afed9 | ||
|
|
f01cb63234 | ||
|
|
8be7480f31 | ||
|
|
751bc7c695 | ||
|
|
3be26dbb95 | ||
|
|
2ca0864ce8 | ||
|
|
b1041220ac | ||
|
|
da842c0cd6 | ||
|
|
c4971e8432 | ||
|
|
e81dbce38c | ||
|
|
0d992689d5 | ||
|
|
4085b3fa69 | ||
|
|
b684bb55c5 | ||
|
|
f98f88ff08 | ||
|
|
f0955df4f0 | ||
|
|
7744c7117d | ||
|
|
22f4d6e26e | ||
|
|
d49a83dec0 | ||
|
|
6d1d5ef68e | ||
|
|
3a1525222d | ||
|
|
f650253cae | ||
|
|
c67c84b443 | ||
|
|
4db32984c6 | ||
|
|
81d39786d9 | ||
|
|
63de076765 | ||
|
|
256ca62a3c | ||
|
|
39d04eb795 | ||
|
|
e02857fcce | ||
|
|
c163caf8cb | ||
|
|
2096a0aa84 | ||
|
|
c253f96c53 | ||
|
|
551daca047 | ||
|
|
ded0abead7 | ||
|
|
e86b191691 | ||
|
|
3321e8c541 | ||
|
|
bc3a805202 | ||
|
|
162b9f4f27 | ||
|
|
e5bc3ea5b4 | ||
|
|
baa139f97a | ||
|
|
470f3e6d29 | ||
|
|
c12a0b57da | ||
|
|
e7ecfa15f8 | ||
|
|
9051eb0039 | ||
|
|
bb8dbfca09 | ||
|
|
0240c70780 | ||
|
|
a41abad5b2 | ||
|
|
a83588b14e | ||
|
|
80192752b7 | ||
|
|
fbd3873a0f | ||
|
|
28cfddbe65 | ||
|
|
b4bdb283ce | ||
|
|
967e57f071 | ||
|
|
d12f6b7dd8 | ||
|
|
182c025c88 | ||
|
|
58e6199ce8 | ||
|
|
6a33f0d483 | ||
|
|
dd230a93e8 | ||
|
|
3aa9cfc18e | ||
|
|
e57f461323 | ||
|
|
ab00514ef6 | ||
|
|
1dd58b4687 | ||
|
|
b4b8a3dfde | ||
|
|
d8782c7b90 | ||
|
|
d5878e9a72 | ||
|
|
ad24bef1c9 | ||
|
|
50246a5066 | ||
|
|
21914cb1c1 | ||
|
|
904935cf98 | ||
|
|
468e75c180 | ||
|
|
849f52b7a6 | ||
|
|
e520297781 | ||
|
|
cefc27112d | ||
|
|
6876f60098 | ||
|
|
fc1653e337 | ||
|
|
e9f5913dac | ||
|
|
7da82c2560 | ||
|
|
a2639c6894 | ||
|
|
eb05fa316f | ||
|
|
6d55adb0d8 | ||
|
|
cff0500a63 | ||
|
|
f3ca184fb6 | ||
|
|
3239c9fdf8 | ||
|
|
7e37f68a5b | ||
|
|
960937ebe9 | ||
|
|
a87784b4c5 | ||
|
|
5812efcf24 | ||
|
|
e82014ec6c | ||
|
|
dc87a4fc91 | ||
|
|
4f5eef2e78 | ||
|
|
6c02afefca | ||
|
|
4696e12641 | ||
|
|
7c688fbf2e | ||
|
|
9ccfc7bd33 | ||
|
|
52a98d76d8 | ||
|
|
8b0cda84e7 | ||
|
|
9f41a9e8b7 | ||
|
|
b7baa96186 | ||
|
|
faa2d4fa8b | ||
|
|
4f5da45fae | ||
|
|
7e54ad3f7c | ||
|
|
3bf366bcd8 | ||
|
|
6219f24a03 | ||
|
|
0bd81db538 | ||
|
|
c84700814d | ||
|
|
c244599ae8 | ||
|
|
d99b121379 | ||
|
|
91b869d043 | ||
|
|
08e1ab64b5 | ||
|
|
c1b21fb1e4 | ||
|
|
8451bb7745 | ||
|
|
1a1771cfd8 | ||
|
|
1952e9be49 | ||
|
|
19391b1df1 | ||
|
|
369174c4d3 | ||
|
|
5ba0a2ed48 | ||
|
|
88061b2284 | ||
|
|
ba38037917 | ||
|
|
a7bb731a1d | ||
|
|
58531888e0 | ||
|
|
56ac885f03 | ||
|
|
5e033fd97a | ||
|
|
5d14b97aa6 | ||
|
|
0b0addbcbe | ||
|
|
098d518398 | ||
|
|
ba111e681e | ||
|
|
3de053c903 | ||
|
|
ac1bd57b8b | ||
|
|
3f0fab9160 | ||
|
|
24c13aadc0 | ||
|
|
885bad3555 | ||
|
|
6dde1d7ae3 | ||
|
|
c223960edb | ||
|
|
32646cf2ee | ||
|
|
415ee4903b | ||
|
|
a21dc5a187 | ||
|
|
0929908229 | ||
|
|
cc12a1c8fa | ||
|
|
91b3a0aac6 | ||
|
|
ed35f4e076 | ||
|
|
c4781cb415 | ||
|
|
730a06640b | ||
|
|
6eb25edb31 |
@@ -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/">
|
||||||
|
|||||||
@@ -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)
|
||||||
@@ -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)])
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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,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")
|
||||||
|
|||||||
@@ -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(),
|
||||||
|
|||||||
Reference in New Issue
Block a user