Compare commits

..
192 Commits
Author SHA1 Message Date
William Falcon 2c4052edb6 release v0.11 2019-06-26 18:44:59 -04:00
William Falcon a2f0f20674 finished data parallel 2019-06-26 18:29:38 -04:00
William Falcon a40b21bce0 removed self.model refs 2019-06-26 18:27:25 -04:00
William Falcon b58ec7ad5a removed self.model refs 2019-06-26 18:26:08 -04:00
William Falcon 301a4992f4 removed self.model refs 2019-06-26 18:24:47 -04:00
William Falcon 42fe76f794 removed self.model refs 2019-06-26 18:23:50 -04:00
William Falcon 8f9672603b removed self.model refs 2019-06-26 18:23:02 -04:00
William Falcon 11b4bc3fbc removed self.model refs 2019-06-26 18:21:17 -04:00
William Falcon 787f523a71 removed self.model refs 2019-06-26 18:19:11 -04:00
William Falcon 5c8875130b removed self.model refs 2019-06-26 18:17:40 -04:00
William Falcon bf0f5a5cbb removed self.model refs 2019-06-26 18:12:33 -04:00
William Falcon df4ac681ed removed self.model refs 2019-06-26 18:08:46 -04:00
William Falcon c1cbb1039a removed self.model refs 2019-06-26 18:05:48 -04:00
William Falcon bc0278252e removed self.model refs 2019-06-26 18:04:29 -04:00
William Falcon 12a0e98920 updated args 2019-06-26 17:54:59 -04:00
William Falcon 4a3c9de857 updated args 2019-06-26 17:53:05 -04:00
William Falcon 0b1e22ac51 updated args 2019-06-26 17:52:14 -04:00
William Falcon 808e86b17c updated args 2019-06-26 17:50:09 -04:00
William Falcon 71cd8f549d updated args 2019-06-26 17:49:58 -04:00
William Falcon 1ee6d21db2 updated args 2019-06-26 17:46:55 -04:00
William Falcon f8be24b09c updated args 2019-06-26 17:44:34 -04:00
William Falcon 1b497ac69a updated args 2019-06-25 20:32:20 -04:00
William Falcon d016431a3f updated args 2019-06-25 20:31:29 -04:00
William Falcon a2e4944f60 updated args 2019-06-25 20:31:10 -04:00
William Falcon 4d5123e379 updated args 2019-06-25 20:29:26 -04:00
William Falcon 5ce4e872de updated args 2019-06-25 20:28:33 -04:00
William Falcon 5eaaf82837 updated args 2019-06-25 20:27:17 -04:00
William Falcon 7527167f69 updated args 2019-06-25 20:25:34 -04:00
William Falcon 45331b396f updated args 2019-06-25 20:24:43 -04:00
William Falcon 440f47b864 updated args 2019-06-25 20:24:03 -04:00
William Falcon 88606c581f updated args 2019-06-25 20:22:59 -04:00
William Falcon f49c2f4c25 updated args 2019-06-25 20:22:21 -04:00
William Falcon 51305697c1 updated args 2019-06-25 20:21:11 -04:00
William Falcon 9b46f13230 updated args 2019-06-25 20:20:12 -04:00
William Falcon 078bbc5df5 updated args 2019-06-25 20:19:11 -04:00
William Falcon 4c556e9880 updated args 2019-06-25 20:19:02 -04:00
William Falcon 89a79a5d3c updated args 2019-06-25 20:18:19 -04:00
William Falcon e3f96d6f3a updated args 2019-06-25 20:17:50 -04:00
William Falcon d33048c67b updated args 2019-06-25 20:16:59 -04:00
William Falcon fea10fc792 updated args 2019-06-25 20:15:10 -04:00
William Falcon 7a7a9a9da0 updated args 2019-06-25 20:14:29 -04:00
William Falcon a76ae6bc48 updated args 2019-06-25 20:12:46 -04:00
William Falcon 0460821398 updated args 2019-06-25 20:12:41 -04:00
William Falcon 7cb6e34beb updated args 2019-06-25 20:10:23 -04:00
William Falcon 2ac5cce67a updated args 2019-06-25 20:09:40 -04:00
William Falcon bac0ef2d44 updated args 2019-06-25 20:08:32 -04:00
William Falcon d3b621dfd2 updated args 2019-06-25 20:04:27 -04:00
William Falcon ac88e3f832 updated args 2019-06-25 20:03:27 -04:00
William Falcon 117515db48 updated args 2019-06-25 20:00:43 -04:00
William Falcon 69be732b11 updated args 2019-06-25 19:56:47 -04:00
William Falcon b59af1813b updated args 2019-06-25 19:56:12 -04:00
William Falcon 7814b2d449 updated args 2019-06-25 19:54:28 -04:00
William Falcon c941649532 updated args 2019-06-25 19:52:26 -04:00
William Falcon 0795e4d51b updated args 2019-06-25 19:46:49 -04:00
William Falcon 158aca26e2 updated args 2019-06-25 19:45:31 -04:00
William Falcon cf57be9dca updated args 2019-06-25 19:43:25 -04:00
William Falcon 8df13035eb updated args 2019-06-25 19:42:15 -04:00
William Falcon c54dd94295 updated args 2019-06-25 19:35:11 -04:00
William Falcon e801914d1d updated args 2019-06-25 19:18:27 -04:00
William Falcon 89410e9090 updated args 2019-06-25 19:17:17 -04:00
William Falcon c4da914747 updated args 2019-06-25 19:06:39 -04:00
William Falcon 41a935185c updated args 2019-06-25 19:06:19 -04:00
William Falcon 73b4976500 updated args 2019-06-25 19:04:49 -04:00
William Falcon 4d42b1ed5f updated args 2019-06-25 19:00:38 -04:00
William Falcon 0fd4d5e7a1 updated args 2019-06-25 18:59:37 -04:00
William Falcon bf3b86ce4d updated args 2019-06-25 18:58:45 -04:00
William Falcon 684dfd0a38 updated args 2019-06-25 18:57:25 -04:00
William Falcon d4ca295762 updated args 2019-06-25 18:51:41 -04:00
William Falcon de0f7fc936 updated args 2019-06-25 18:47:11 -04:00
William Falcon 7b22de22a7 updated args 2019-06-25 18:45:19 -04:00
William Falcon d8cb739ab2 updated args 2019-06-25 18:44:50 -04:00
William Falcon c45a329df4 updated args 2019-06-25 18:44:11 -04:00
William Falcon 41f68861d5 updated args 2019-06-25 18:42:44 -04:00
William Falcon 156dc3e5ee updated args 2019-06-25 18:40:34 -04:00
William Falcon 8e10179214 updated args 2019-06-25 18:29:43 -04:00
William Falcon 775ca3736b updated args 2019-06-25 18:29:16 -04:00
William Falcon 242cccc234 updated args 2019-06-25 18:25:51 -04:00
William Falcon a00b8f7861 updated args 2019-06-25 18:25:19 -04:00
William Falcon 338f889e7c updated args 2019-06-25 18:23:29 -04:00
William Falcon e12c8ad21a updated args 2019-06-25 18:22:10 -04:00
William Falcon 35aa67df56 updated args 2019-06-25 18:20:24 -04:00
William Falcon e58cfafa74 updated args 2019-06-25 18:18:40 -04:00
William Falcon e58eee8d6a updated args 2019-06-25 18:18:20 -04:00
William Falcon b9d5397196 updated args 2019-06-25 18:14:48 -04:00
William Falcon 3f8e133303 updated args 2019-06-25 18:13:01 -04:00
William Falcon 983551653d fixed basic trainer 2019-06-25 18:11:13 -04:00
William Falcon 6c705a0525 adding framework level dp 2019-06-25 18:10:15 -04:00
William Falcon 516f441153 adding framework level dp 2019-06-25 18:09:29 -04:00
William Falcon cbc627459a adding framework level dp 2019-06-25 17:56:01 -04:00
William Falcon a519e0755b release v0.1.dev21 2019-06-14 10:05:03 -04:00
William Falcon 9bf3fcd45e adding support for interrupt signals 2019-06-14 09:59:28 -04:00
William Falcon 88ff860c90 adding support for interrupt signals 2019-06-14 09:46:41 -04:00
William Falcon edf03063a1 adding support for interrupt signals 2019-06-14 09:44:19 -04:00
William Falcon 32edc6d7b7 adding support for interrupt signals 2019-06-14 09:42:36 -04:00
William Falcon cd36b63167 adding support for interrupt signals 2019-06-14 09:39:52 -04:00
William Falcon 519d2e9321 adding support for interrupt signals 2019-06-14 09:28:23 -04:00
William Falcon 8cca02d652 adding support for interrupt signals 2019-06-14 09:25:46 -04:00
William Falcon 69274d304d adding support for interrupt signals 2019-06-14 09:24:51 -04:00
William Falcon d98e799404 adding dataparallel 2019-06-07 15:06:22 -04:00
William Falcon 931a45b760 dev2 release 2019-06-07 11:39:49 -04:00
William Falcon eb5b3cfee1 Update setup.py 2019-06-06 18:04:58 -04:00
William Falcon 15ca7a40a6 release v 2019-05-24 15:30:55 -04:00
William Falcon 96903c7910 added amp level option 2019-05-16 16:01:15 -04:00
William Falcon eb13bb8313 added amp level option 2019-05-16 15:58:58 -04:00
William Falcon d560fac104 added amp level option 2019-05-16 15:58:14 -04:00
William Falcon 2d3977046e added amp level option 2019-05-16 15:58:06 -04:00
William Falcon fa0a223ccb added amp level option 2019-05-16 15:55:29 -04:00
William Falcon 60d4b80322 added amp level option 2019-05-16 15:55:21 -04:00
William Falcon e052a3bc92 added amp level option 2019-05-16 15:52:00 -04:00
William Falcon 35ca80683e added amp level option 2019-05-16 15:47:21 -04:00
William Falcon b2ef6a6366 added amp level option 2019-05-16 15:46:17 -04:00
William Falcon 9d19ab5850 added amp level option 2019-05-16 15:45:56 -04:00
William Falcon 92f9b3e062 fixed alternating loss 2019-05-14 06:40:11 -04:00
William Falcon 5fa2a6a723 tng and val steps now have batch nbs 2019-05-14 06:37:56 -04:00
William Falcon 8531f33549 tng and val steps now have batch nbs 2019-05-14 06:36:26 -04:00
William Falcon 98b26c5c7e fixed error with shorter batch cycles 2019-05-14 06:11:52 -04:00
William Falcon c973245ba1 fixed error with shorter batch cycles 2019-05-14 06:11:16 -04:00
William Falcon ed787fb061 release v0.1.dev182 2019-05-14 05:53:58 -04:00
William Falcon 04681eeda9 release v0.1.dev18 2019-05-14 05:46:55 -04:00
William Falcon 6519c29119 added 16 bit training support with --use_amp flag 2019-05-14 05:44:33 -04:00
William Falcon 3b0fd7a6cb added option to change default tensor 2019-05-13 22:03:56 -04:00
William Falcon a8e57602d3 added option to change default tensor 2019-05-13 22:03:47 -04:00
William Falcon 8836f4f7a5 added option to change default tensor 2019-05-13 22:02:53 -04:00
William Falcon f246ae7fab added option to change default tensor 2019-05-13 21:55:57 -04:00
William Falcon 1c7d477d03 added option to change default tensor 2019-05-13 21:52:02 -04:00
William Falcon 90a460ec62 added option to change default tensor 2019-05-13 21:47:07 -04:00
William Falcon edd406f419 added option to change default tensor 2019-05-13 21:28:28 -04:00
William Falcon 8a68466710 added option to change default tensor 2019-05-13 21:27:01 -04:00
William Falcon 4dbf38093a added option to change default tensor 2019-05-13 21:22:50 -04:00
William Falcon 38717abcd4 added option to change default tensor 2019-05-13 21:19:37 -04:00
William Falcon 8e49fc6cf7 added option to change default tensor 2019-05-13 21:19:07 -04:00
William Falcon 5f0a71c414 added option to change default tensor 2019-05-13 21:18:17 -04:00
William Falcon 88fbf6cc4b added option to change default tensor 2019-05-13 20:44:25 -04:00
William Falcon 7002de1d4e added option to change default tensor 2019-05-13 20:43:26 -04:00
William Falcon fecd6a00cb added option to change default tensor 2019-05-13 20:41:23 -04:00
William Falcon 4693276494 added option to change default tensor 2019-05-13 20:40:07 -04:00
William Falcon f228e5ae66 added option to change default tensor 2019-05-13 19:39:56 -04:00
William Falcon e3425ec6a0 added option to change default tensor 2019-05-13 19:30:06 -04:00
William Falcon 5a7ad19403 fixed gpu map location 2019-05-13 05:32:18 -04:00
William Falcon d6bc203f05 release v0.1.dev16 2019-05-05 12:16:52 -04:00
William Falcon 12352f1949 fixed epoch continuation from checkpoint 2019-05-05 12:15:04 -04:00
William Falcon f881bf6750 added log saving when early epoch stop 2019-04-23 11:12:01 -04:00
William Falcon 0637d8e7a5 release v0.1.dev15 2019-04-23 09:08:06 -04:00
William Falcon 2514f62913 early epoch stopping 2019-04-23 08:57:58 -04:00
William Falcon 95aee7ff96 early epoch stopping 2019-04-23 08:46:20 -04:00
William Falcon ffd6dc678c early epoch stopping 2019-04-23 08:27:27 -04:00
William Falcon 1961a6abb2 early epoch stopping 2019-04-23 08:26:48 -04:00
William Falcon 676d76d839 pointer to trainer in model 2019-04-23 07:25:09 -04:00
William Falcon b625b293f4 running new CE then DDT 2019-04-21 14:46:33 -04:00
William Falcon 333f0fde9b fixed hooks 2019-04-21 14:16:54 -04:00
William Falcon 4b0b7e5ea3 if return -1 from a hook that loop stopps 2019-04-21 13:40:32 -04:00
William Falcon e89da15f18 if return -1 from a hook that loop stopps 2019-04-21 13:38:50 -04:00
William Falcon 004f015ee0 fixed imports 2019-04-21 13:13:09 -04:00
William Falcon 398b709b76 fixex imports 2019-04-21 13:12:42 -04:00
William Falcon e9bcbc2318 fixing setup 2019-04-21 13:09:06 -04:00
William Falcon ee51d7b7bc fixing setup 2019-04-21 13:05:29 -04:00
William Falcon bb75bdf87b fixing setup 2019-04-21 13:02:11 -04:00
William Falcon aeef648199 trainer updates 2019-04-21 12:42:44 -04:00
William Falcon cf110af384 added example and verified 2019-04-21 12:38:51 -04:00
William Falcon 76cc1c6eab added early epoch stopping hook 2019-04-21 12:30:54 -04:00
William Falcon efd750565e added early epoch stopping hook 2019-04-21 12:29:48 -04:00
William Falcon 86261b7404 added early epoch stopping hook 2019-04-21 12:26:35 -04:00
William Falcon 8eca3ffa41 Merge pull request #9 from Derek-Wds/master
Fix some link bugs in the README.md
2019-04-07 02:55:32 -04:00
Dingsu Wang 413c343d83 Update README.md 2019-04-05 16:27:45 -04:00
William Falcon ea2f50f1a4 Merge pull request #8 from shreyasbapat/further_changes
Some more fixes
2019-04-03 14:29:11 -04:00
Shreyas Bapat 4809de8765 Fix pip install too 2019-04-03 22:47:55 +05:30
Shreyas Bapat b79b011d5e Some more fixes 2019-04-03 22:31:22 +05:30
William Falcon 7d3399964b fixed os missing 2019-04-03 12:59:06 -04:00
William Falcon 64827b7029 removed bilstm 2019-04-03 12:55:45 -04:00
William Falcon 18eaa59c28 Merge pull request #6 from shreyasbapat/management
Add src, docs and other important folders
2019-04-03 12:53:11 -04:00
Shreyas Bapat 10b796b5c7 Fix merge conflicts 2019-04-03 22:18:49 +05:30
Shreyas Bapat 18b0c5a122 Add src, docs and other important folders 2019-04-03 22:16:02 +05:30
William Falcon f26488bd16 fixes #4 2019-04-03 11:27:01 -04:00
William Falcon bca1c4b594 Update embeddings.py 2019-04-03 11:21:16 -04:00
William Falcon a01e2ade25 Update embeddings.py 2019-04-03 11:18:51 -04:00
William Falcon 3e9f37a382 fixes #4 2019-04-03 09:07:20 -04:00
William Falcon 89be81863e fixes #5 2019-04-03 09:00:44 -04:00
William Falcon 7f00fa1409 Merge branch 'master' of https://github.com/williamFalcon/pytorch-lightning 2019-04-01 13:05:42 -04:00
William Falcon fc5583c8cd removed .egg 2019-04-01 13:05:35 -04:00
William Falcon d7d52ae2e7 Update README.md 2019-04-01 12:38:31 -04:00
William Falcon 4a2d21fc91 Update README.md 2019-04-01 12:34:38 -04:00
William Falcon 9fe005b0d9 Update README.md 2019-03-31 16:59:39 -04:00
William Falcon 461fed19b6 Update README.md 2019-03-31 16:59:24 -04:00
William Falcon ec57b3fe6d Update README.md 2019-03-31 16:51:00 -04:00
William Falcon e43b1d1d31 Update README.md 2019-03-31 16:50:32 -04:00
William Falcon 8e2e95e55d Update README.md 2019-03-31 16:47:15 -04:00
William Falcon 71113ca770 Update README.md 2019-03-31 16:46:00 -04:00
William Falcon 7e81a17c11 Update README.md 2019-03-31 16:36:29 -04:00
William Falcon 5943438316 Update README.md 2019-03-31 16:35:58 -04:00
William Falcon 72239b4419 Update README.md 2019-03-31 16:35:10 -04:00
William Falcon 9ff6108af1 added example and verified 2019-03-31 16:34:13 -04:00
William Falcon 9d56b1744f release v0.0.2 2019-03-31 16:31:48 -04:00
36 changed files with 489 additions and 466 deletions
+4
View File
@@ -7,6 +7,7 @@ test_tube_data/
datasets/
model_weights/
app/models/
pip-wheel-metadata/
# Byte-compiled / optimized / DLL files
__pycache__/
@@ -115,3 +116,6 @@ ENV/
# mypy
.mypy_cache/
# data
mnist/
View File
+9
View File
@@ -0,0 +1,9 @@
graft docs
include COPYING
include AUTHORS
recursive-include src/einsteinpy/tests *.py *.html
prune docs/source/examples/.ipynb_checkpoints
global-exclude *.py[cod] __pycache__ *.so *.dylib
+76 -61
View File
@@ -1,18 +1,18 @@
<p align="center">
<a href="https://williamfalcon.github.io/pytorch-lightning/">
<img alt="" src="https://github.com/williamFalcon/pytorch-lightning/blob/master/imgs/lightning_logo.png" width="50">
<img alt="" src="https://github.com/williamFalcon/pytorch-lightning/blob/master/docs/source/_static/lightning_logo.png" width="50">
</a>
</p>
<h3 align="center">
Pytorch Lightning
</h3>
<p align="center">
The Keras for ML-researchers in PyTorch. More control. Less boilerplate.
The Keras for ML researchers using PyTorch. More control. Less boilerplate.
</p>
<p align="center">
<a href="https://badge.fury.io/py/pytorch_lightning"><img src="https://badge.fury.io/py/pytorch_lightning.svg"></a>
<a href="https://travis-ci.org/williamFalcon/test-tube"><img src="https://travis-ci.org/williamFalcon/pytorch-lightning.svg?branch=master"></a>
<a href="https://github.com/williamFalcon/pytorch-lightning/blob/master/LICENSE"><img src="https://img.shields.io/badge/License-MIT-yellow.svg"></a>
<a href="https://badge.fury.io/py/pytorch-lightning"><img src="https://badge.fury.io/py/pytorch-lightning.svg" alt="PyPI version" height="18"></a>
<!-- <a href="https://travis-ci.org/williamFalcon/test-tube"><img src="https://travis-ci.org/williamFalcon/pytorch-lightning.svg?branch=master"></a> -->
<a href="https://github.com/williamFalcon/pytorch-lightning/blob/master/COPYING"><img src="https://img.shields.io/badge/License-MIT-yellow.svg"></a>
</p>
```bash
@@ -22,37 +22,66 @@ pip install pytorch-lightning
## Docs
In progress. Documenting now!
## What is it?
All you do is define the forward passes, your data and **lightning runs everything else for you:**
## Disclaimer
This is a research tool I built for myself internally while doing my PhD. The API is not 100% production quality, but my hope is that by open-sourcing, we can all get it there (I don't have too much time nowadays to write production-level code).
1. Running the training loop.
2. Running the validation loop.
3. Running the testing loop.
## What is it?
Keras is too abstract for researchers. Lightning makes it so you only have to define your model but still control all details of training if you need to.
Pytorch
<-- Lightning
Your model.
**Lightning will do the following for you:**
1. Run the training loop.
2. Run the validation loop.
3. Run the testing loop.
4. Early stopping.
5. Learning rate annealing.
5. Learning rate annealing.
6. Can train complex models like GANs or anything with multiple optimizers.
7. Weight checkpointing.
8. Model saving.
9. Model loading.
10. Logging training details (through test-tube).
11. Running training on multiple GPUs (through test-tube).
12. Running training on a GPU cluster managed by SLURM (through test-tube).
13. Distributing memory-bound models on multiple GPUs.
14. Gives your model hyperparameters parsed from the command line OR a JSON file.
15. Runs your model in a dev environment where nothing logs.
10. Log training details (through test-tube).
11. Run training on multiple GPUs (through test-tube).
12. Run training on a GPU cluster managed by SLURM (through test-tube).
13. Distribute memory-bound models on multiple GPUs.
14. Give your model hyperparameters parsed from the command line OR a JSON file.
15. Run your model in a dev environment where nothing logs.
## Usage
To use lightning do 2 things:
1. [Define a trainer](https://github.com/williamFalcon/pytorch-lightning/blob/master/pytorch_lightning/trainer_main.py) (which will run ALL your models).
2. [Define a model](https://github.com/williamFalcon/pytorch-lightning/blob/master/pytorch_lightning/models/sample_model_template/model_template.py).
1. [Define a trainer](https://github.com/williamFalcon/pytorch-lightning/blob/master/docs/source/examples/basic_trainer.py) (which will run ALL your models).
2. [Define a model](https://github.com/williamFalcon/pytorch-lightning/blob/master/docs/source/examples/example_model.py).
### Example:
#### Quick demo
Run the following demo to see how it works:
```bash
# install lightning
pip install pytorch-lightning
#### Define the trainer
# clone lightning for the demo
git clone https://github.com/williamFalcon/pytorch-lightning.git
cd pytorch-lightning/docs/source/examples
# run demo (on cpu)
python fully_featured_trainer.py
```
Without changing the model AT ALL, you can run the model on a single gpu, over multiple gpus, or over multiple nodes.
```bash
# run a grid search on two gpus
python fully_featured_trainer.py --gpus "0;1"
# run single model on multiple gpus
python fully_featured_trainer.py --gpus "0;1" --interactive
```
#### Basic trainer example
See [this demo](https://github.com/williamFalcon/pytorch-lightning/blob/master/docs/source/examples/fully_featured_trainer.py) for a more robust trainer example.
```python
# trainer.py
import os
import sys
@@ -82,35 +111,17 @@ def main(hparams):
exp.argparse(hparams)
exp.save()
model_save_path = '{}/{}/{}'.format(hparams.model_save_path, exp.name, exp.version)
# build model
print('loading model...')
model = ExampleModel(hparams)
print('model built')
# callbacks
early_stop = EarlyStopping(
monitor=hparams.early_stop_metric,
patience=hparams.early_stop_patience,
verbose=True,
mode=hparams.early_stop_mode
)
model_save_path = '{}/{}/{}'.format(hparams.model_save_path, exp.name, exp.version)
checkpoint = ModelCheckpoint(
filepath=model_save_path,
save_function=None,
save_best_only=True,
verbose=True,
monitor=hparams.model_save_monitor_value,
mode=hparams.model_save_monitor_mode
)
early_stop = EarlyStopping(monitor='val_acc', patience=3, mode='min', verbose=True)
checkpoint = ModelCheckpoint(filepath=model_save_path, save_function=None, save_best_only=True, verbose=True, monitor='val_acc', mode='min')
# configure trainer
trainer = Trainer(
experiment=exp,
checkpoint_callback=checkpoint,
early_stop_callback=early_stop,
)
trainer = Trainer(experiment=exp, checkpoint_callback=checkpoint, early_stop_callback=early_stop)
# train model
trainer.fit(model)
@@ -129,9 +140,11 @@ if __name__ == '__main__':
# train model
main(hyperparams)
```
#### Define the model
#### Basic model example
Here we only show the method signatures. It's up to you to define the content.
```python
from torch import nn
@@ -140,32 +153,32 @@ class My_Model(RootModule):
def __init__(self):
# define model
self.l1 = nn.Linear(200, 10)
# ---------------
# TRAINING
def training_step(self, data_batch):
x, y = data_batch
y_hat = self.l1(x)
loss = some_loss(y_hat)
return loss_val, {'train_loss': loss}
def validation_step(self, data_batch):
x, y = data_batch
y_hat = self.l1(x)
loss = some_loss(y_hat)
return loss_val, {'val_loss': loss}
def validation_end(self, outputs):
total_accs = []
for output in outputs:
total_accs.append(output['val_acc'].item())
# return a dict
return {'total_acc': np.mean(total_accs)}
# ---------------
# SAVING
def get_save_dict(self):
@@ -177,7 +190,7 @@ class My_Model(RootModule):
def load_model_specific(self, checkpoint):
# lightning loads for you. Here's your chance to say what you want to load
self.load_state_dict(checkpoint['state_dict'])
# ---------------
# TRAINING CONFIG
def configure_optimizers(self):
@@ -185,7 +198,7 @@ class My_Model(RootModule):
# lightning will call automatically
optimizer = self.choose_optimizer('adam', self.parameters(), {'lr': self.hparams.learning_rate}, 'optimizer')
return [optimizer]
@property
def tng_dataloader(self):
return pytorch_dataloader('train')
@@ -197,7 +210,7 @@ class My_Model(RootModule):
@property
def test_dataloader(self):
return pytorch_dataloader('test')
# ---------------
# MODIFY YOUR COMMAND LINE ARGS
@staticmethod
@@ -206,6 +219,8 @@ class My_Model(RootModule):
parser.add_argument('--out_features', default=20)
return parser
```
### Details
#### Model definition
@@ -214,7 +229,7 @@ class My_Model(RootModule):
| training_step | Called with a batch of data during training | data from your dataloaders | tuple: scalar, dict |
| validation_step | Called with a batch of data during validation | data from your dataloaders | tuple: scalar, dict |
| validation_end | Collate metrics from all validation steps | outputs: array where each item is the output of a validation step | dict: for logging |
| get_save_dict | called when your model needs to be saved (checkpoints, hpc save, etc...) | None | dict to be saved |
| get_save_dict | called when your model needs to be saved (checkpoints, hpc save, etc...) | None | dict to be saved |
#### Model training
| Name | Description | Input | Return |
@@ -230,7 +245,7 @@ class My_Model(RootModule):
|---|---|---|---|
| get_save_dict | called when your model needs to be saved (checkpoints, hpc save, etc...) | None | dict to be saved |
| load_model_specific | called when loading a model | checkpoint: dict you created in get_save_dict | dict: modified in whatever way you want |
## Optional model hooks.
Add these to the model whenever you want to configure training behavior.
Binary file not shown.
View File

Before

Width:  |  Height:  |  Size: 11 KiB

After

Width:  |  Height:  |  Size: 11 KiB

+1
View File
@@ -0,0 +1 @@
from .example_model import ExampleModel
@@ -5,7 +5,7 @@ from test_tube import HyperOptArgumentParser, Experiment
from pytorch_lightning.models.trainer import Trainer
from pytorch_lightning.utils.arg_parse import add_default_args
from pytorch_lightning.utils.pt_callbacks import EarlyStopping, ModelCheckpoint
from demo.example_model import ExampleModel
from docs.source.examples.example_model import ExampleModel
def main(hparams):
@@ -28,16 +28,14 @@ def main(hparams):
exp.save()
# build model
print('loading model...')
model = ExampleModel(hparams)
print('model built')
# callbacks
early_stop = EarlyStopping(
monitor=hparams.early_stop_metric,
patience=hparams.early_stop_patience,
monitor='val_acc',
patience=3,
mode='min',
verbose=True,
mode=hparams.early_stop_mode
)
model_save_path = '{}/{}/{}'.format(hparams.model_save_path, exp.name, exp.version)
@@ -46,8 +44,8 @@ def main(hparams):
save_function=None,
save_best_only=True,
verbose=True,
monitor=hparams.model_save_monitor_value,
mode=hparams.model_save_monitor_mode
monitor='val_acc',
mode='min'
)
# configure trainer
@@ -6,6 +6,8 @@ from torchvision.datasets import MNIST
import torchvision.transforms as transforms
import torch
import torch.nn.functional as F
import os, pdb
from collections import OrderedDict
class ExampleModel(RootModule):
@@ -40,8 +42,9 @@ class ExampleModel(RootModule):
# TRAINING
# ---------------------
def forward(self, x):
x = self.c_d1(x)
x = F.tanh(x)
x = torch.tanh(x)
x = self.c_d1_bn(x)
x = self.c_d1_drop(x)
@@ -54,7 +57,7 @@ class ExampleModel(RootModule):
nll = F.nll_loss(logits, labels)
return nll
def training_step(self, data_batch):
def training_step(self, data_batch, batch_i):
"""
Called inside the training loop
:param data_batch:
@@ -68,10 +71,16 @@ class ExampleModel(RootModule):
# calculate loss
loss_val = self.loss(y, y_hat)
tqdm_dic = {'tng_loss': loss_val.item()}
return loss_val, tqdm_dic
# tqdm_dic = {'tng_loss': loss_val.item()}
# return loss_val, tqdm_dic
def validation_step(self, data_batch):
output = OrderedDict({
'loss': loss_val,
'tqdm_metrics': {}
})
return output
def validation_step(self, data_batch, batch_i):
"""
Called inside the validation loop
:param data_batch:
@@ -87,9 +96,14 @@ class ExampleModel(RootModule):
labels_hat = torch.argmax(y_hat, dim=1)
val_acc = torch.sum(y == labels_hat).item() / (len(y) * 1.0)
output = {'y_hat': y_hat, 'val_loss': loss_val.item(), 'val_acc': val_acc}
# output = {'y_hat': y_hat, 'val_loss': loss_val.item(), 'val_acc': val_acc}
output = OrderedDict({
'val_loss': loss_val,
'val_acc': torch.tensor(val_acc),
})
return output
def validation_end(self, outputs):
"""
Called at the end of validation to aggregate outputs
@@ -97,13 +111,14 @@ class ExampleModel(RootModule):
:return:
"""
val_loss_mean = 0
accs = []
val_acc_mean = 0
for output in outputs:
val_loss_mean += output['val_loss']
accs.append(output['val_acc'])
val_acc_mean += output['val_acc']
val_loss_mean /= len(outputs)
tqdm_dic = {'val_loss': val_loss_mean, 'val_acc': np.mean(accs)}
val_acc_mean /= len(outputs)
tqdm_dic = {'val_loss': val_loss_mean.item(), 'val_acc': val_acc_mean.item()}
return tqdm_dic
def update_tng_log_metrics(self, logs):
@@ -177,7 +192,7 @@ class ExampleModel(RootModule):
return self._test_dataloader
@staticmethod
def add_model_specific_args(parent_parser):
def add_model_specific_args(parent_parser, root_dir):
parser = HyperOptArgumentParser(strategy=parent_parser.strategy, parents=[parent_parser])
# param overwrites
@@ -186,11 +201,11 @@ class ExampleModel(RootModule):
# network params
parser.opt_list('--drop_prob', default=0.2, options=[0.2, 0.5], type=float, tunable=False)
parser.add_argument('--in_features', default=28*28)
parser.add_argument('--hidden_dim', default=500)
parser.add_argument('--out_features', default=10)
parser.add_argument('--hidden_dim', default=50000) # use 500 for CPU, 50000 for GPU to see speed difference
# data
parser.add_argument('--data_root', default='/Users/williamfalcon/Developer/personal/research_lib/research_proj/datasets/mnist', type=str)
parser.add_argument('--data_root', default=os.path.join(root_dir, 'mnist'), type=str)
# training params (opt)
parser.opt_list('--learning_rate', default=0.001, type=float, options=[0.0001, 0.0005, 0.001, 0.005],
@@ -17,7 +17,7 @@ np.random.seed(SEED)
# ---------------------
# DEFINE MODEL HERE
# ---------------------
from demo.example_model import ExampleModel
from docs.source.examples.example_model import ExampleModel
# ---------------------
AVAILABLE_MODELS = {
@@ -27,7 +27,7 @@ AVAILABLE_MODELS = {
"""
Allows training by using command line arguments
Run by:
Run by:
# TYPE YOUR RUN COMMAND HERE
"""
@@ -42,9 +42,7 @@ def main(hparams, cluster, results_dict):
:param hparams:
:return:
"""
on_gpu = torch.cuda.is_available()
if hparams.disable_cuda:
on_gpu = False
on_gpu = hparams.gpus is not None and torch.cuda.is_available()
device = 'cuda' if on_gpu else 'cpu'
hparams.__setattr__('device', device)
@@ -57,13 +55,14 @@ def main(hparams, cluster, results_dict):
sleep(process_position + 1)
# init experiment
log_dir = os.path.dirname(os.path.realpath(__file__))
exp = Experiment(
name=hparams.tt_name,
debug=hparams.debug,
save_dir=hparams.tt_save_path,
version=hparams.hpc_exp_number,
name='test_tube_exp',
debug=True,
save_dir=log_dir,
version=0,
autosave=False,
description=hparams.tt_description
description='test demo'
)
exp.argparse(hparams)
@@ -92,12 +91,18 @@ def main(hparams, cluster, results_dict):
mode=hparams.model_save_monitor_mode
)
# gpus are ; separated for inside a node and , within nodes
gpu_list = None
if hparams.gpus is not None:
gpu_list = [int(x) for x in hparams.gpus.split(';')]
# configure trainer
trainer = Trainer(
experiment=exp,
cluster=cluster,
checkpoint_callback=checkpoint,
early_stop_callback=early_stop,
gpus=gpu_list
)
# train model
@@ -108,7 +113,7 @@ def get_default_parser(strategy, root_dir):
possible_model_names = list(AVAILABLE_MODELS.keys())
parser = HyperOptArgumentParser(strategy=strategy, add_help=False)
add_default_args(parser, root_dir, possible_model_names, SEED)
add_default_args(parser, root_dir, possible_model_names=possible_model_names, rand_seed=SEED)
return parser
@@ -158,35 +163,40 @@ if __name__ == '__main__':
model_name = 'model_template'
# use default args
root_dir = os.path.split(os.path.dirname(sys.modules['__main__'].__file__))[0]
root_dir = os.path.dirname(os.path.realpath(__file__))
parent_parser = get_default_parser(strategy='random_search', root_dir=root_dir)
# allow model to overwrite or extend args
TRAINING_MODEL = AVAILABLE_MODELS[model_name]
parser = TRAINING_MODEL.add_model_specific_args(parent_parser)
parser = TRAINING_MODEL.add_model_specific_args(parent_parser, root_dir)
hyperparams = parser.parse_args()
# format GPU layout
os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID"
gpu_ids = hyperparams.gpus.split(';')
# ---------------------
# RUN TRAINING
# ---------------------
# cluster and CPU
if hyperparams.on_cluster:
# Gets called when running via HPC cluster
# run on HPC cluster
print('RUNNING ON SLURM CLUSTER')
gpu_ids = hyperparams.gpus.split(';')
os.environ["CUDA_VISIBLE_DEVICES"] = ','.join(gpu_ids)
optimize_on_cluster(hyperparams)
elif hyperparams.single_run_gpu:
# run on 1 gpu
print(f'RUNNING 1 TRIAL ON GPU. gpu: {gpu_ids[0]}')
os.environ["CUDA_VISIBLE_DEVICES"] = gpu_ids[0]
elif hyperparams.gpus is None:
# run on cpu
print('RUNNING ON CPU')
main(hyperparams, None, None)
elif hyperparams.local or hyperparams.single_run:
# run 1 trial but on CPU
os.environ["CUDA_VISIBLE_DEVICES"] = '0'
print('RUNNING LOCALLY')
# single or multiple GPUs on same machine
gpu_ids = hyperparams.gpus.split(';')
if hyperparams.interactive:
# run on 1 gpu
print(f'RUNNING INTERACTIVE MODE ON GPUS. gpu ids: {gpu_ids}')
os.environ["CUDA_VISIBLE_DEVICES"] = ','.join(gpu_ids)
main(hyperparams, None, None)
else:
@@ -198,4 +208,3 @@ if __name__ == '__main__':
nb_trials=hyperparams.nb_hopt_trials,
nb_workers=len(gpu_ids)
)
BIN
View File
Binary file not shown.
+5
View File
@@ -0,0 +1,5 @@
[build-system]
requires = [
"setuptools",
"wheel",
]
-10
View File
@@ -1,10 +0,0 @@
Metadata-Version: 1.0
Name: pytorch-lightning
Version: 0.0.1
Summary: Rapid research framework
Home-page: https://github.com/williamFalcon/pytorch-lightning
Author: UNKNOWN
Author-email: UNKNOWN
License: UNKNOWN
Description: UNKNOWN
Platform: UNKNOWN
-32
View File
@@ -1,32 +0,0 @@
README.md
setup.py
notebooks/__init__.py
pytorch_lightning/__init__.py
pytorch_lightning/trainer_main.py
pytorch_lightning.egg-info/PKG-INFO
pytorch_lightning.egg-info/SOURCES.txt
pytorch_lightning.egg-info/dependency_links.txt
pytorch_lightning.egg-info/requires.txt
pytorch_lightning.egg-info/top_level.txt
pytorch_lightning/models/__init__.py
pytorch_lightning/models/trainer.py
pytorch_lightning/models/model_examples/__init__.py
pytorch_lightning/models/model_examples/bilstm.py
pytorch_lightning/models/sample_model_template/__init__.py
pytorch_lightning/models/sample_model_template/model_template.py
pytorch_lightning/root_module/__init__.py
pytorch_lightning/root_module/grads.py
pytorch_lightning/root_module/hooks.py
pytorch_lightning/root_module/memory.py
pytorch_lightning/root_module/model_saving.py
pytorch_lightning/root_module/optimization.py
pytorch_lightning/root_module/root_module.py
pytorch_lightning/utils/__init__.py
pytorch_lightning/utils/arg_parse.py
pytorch_lightning/utils/embeddings.py
pytorch_lightning/utils/plotting.py
pytorch_lightning/utils/pt_callbacks.py
tests/__init__.py
tests/research_proj/__init__.py
tests/research_proj/sample_model_template/__init__.py
tests/research_proj/sample_model_template/model_template_test.py
@@ -1 +0,0 @@
-3
View File
@@ -1,3 +0,0 @@
notebooks
pytorch_lightning
tests
@@ -1,167 +0,0 @@
import torch.nn as nn
import numpy as np
from test_tube import HyperOptArgumentParser
import torch
from torch.autograd import Variable
from sklearn.metrics import confusion_matrix, f1_score
from torch.nn import functional as F
class BiLSTMPack(nn.Module):
"""
Sample model to show how to define a template
"""
def __init__(self, hparams):
# init superclass
super(BiLSTMPack, self).__init__(hparams)
self.hidden = None
# trigger tag building
self.ner_tagset = {'O': 0, 'I-Bio': 1}
self.nb_tags = len(self.ner_tagset)
# build model
print('building model...')
if hparams.model_load_weights_path is None:
self.__build_model()
print('model built')
else:
self = BiLSTMPack.load(hparams.model_load_weights_path, hparams.on_gpu, hparams)
print('model loaded from: {}'.format(hparams.model_load_weights_path))
def __build_model(self):
"""
Layout model
:return:
"""
# design the number of final units
self.output_dim = self.hparams.nb_lstm_units
# when it's bidirectional our weights double
if self.hparams.bidirectional:
self.output_dim *= 2
# total number of words
total_words = len(self.tng_dataloader.dataset.words_token_to_idx)
# word embeddings
self.word_embedding = nn.Embedding(
num_embeddings=total_words + 1,
embedding_dim=self.hparams.embedding_dim,
padding_idx=0
)
# design the LSTM
self.lstm = nn.LSTM(
self.hparams.embedding_dim,
self.hparams.nb_lstm_units,
num_layers=self.hparams.nb_lstm_layers,
bidirectional=self.hparams.bidirectional,
dropout=self.hparams.drop_prob,
batch_first=True,
)
# map to tag space
self.fc_out = nn.Linear(self.output_dim, self.out_dim)
self.hidden_to_tag = nn.Linear(self.output_dim, self.nb_tags)
def init_hidden(self, batch_size):
# the weights are of the form (nb_layers * 2 if bidirectional, batch_size, nb_lstm_units)
mult = 2 if self.hparams.bidirectional else 1
hidden_a = torch.randn(self.hparams.nb_layers * mult, batch_size, self.nb_rnn_units)
hidden_b = torch.randn(self.hparams.nb_layers * mult, batch_size, self.nb_rnn_units)
if self.hparams.on_gpu:
hidden_a = hidden_a.cuda()
hidden_b = hidden_b.cuda()
hidden_a = Variable(hidden_a)
hidden_b = Variable(hidden_b)
return (hidden_a, hidden_b)
def forward(self, model_in):
# layout data (expand it, etc...)
# x = sequences
x, seq_lengths = model_in
batch_size, seq_len = x.size()
# reset RNN hidden state
self.hidden = self.init_hidden(batch_size)
# embed
x = self.word_embedding(x)
# run through rnn using packed sequences
x = torch.nn.utils.rnn.pack_padded_sequence(x, seq_lengths, batch_first=True)
x, self.hidden = self.lstm(x, self.hidden)
x, _ = torch.nn.utils.rnn.pad_packed_sequence(x, batch_first=True)
# if asked for only last state, use the h_n which is the same as out(t=n)
if not self.return_sequence:
# pull out hidden states
# h_n = (nb_directions * nb_layers, batch_size, emb_size)
nb_directions = 2 if self.bidirectional else 1
(h_n, _) = self.hidden
# reshape to make indexing easier
# forward = 0, backward = 1 (of nb_directions)
h_n = h_n.view(self.nb_layers, nb_directions, batch_size, self.nb_rnn_units)
# pull out last forward
forward_h_n = h_n[-1, 0, :, :]
x = forward_h_n
# if bidirectional, also pull out the last hidden of backward network
if self.bidirectional:
backward_h_n = h_n[-1, 1, :, :]
x = torch.cat([forward_h_n, backward_h_n], dim=1)
# project to tag space
x = x.contiguous()
x = x.view(-1, self.output_dim)
x = self.hidden_to_tag(x)
return x
def loss(self, model_out):
# cross entropy loss
logits, y = model_out
y, y_lens = y
# flatten y and logits
y = y.view(-1)
logits = logits.view(-1, self.nb_tags)
# calculate a mask to remove padding tokens
mask = (y >= 0).float()
# count how many tokens we have
num_tokens = int(torch.sum(mask).data[0])
# pick the correct values and mask out
logits = logits[range(logits.shape[0]), y] * mask
# compute the ce loss
ce_loss = -torch.sum(logits)/num_tokens
return ce_loss
def pull_out_last_embedding(self, x, seq_lengths, batch_size, on_gpu):
# grab only the last activations from the non-padded ouput
x_last = torch.zeros([batch_size, 1, x.size(-1)])
for i, seq_len in enumerate(seq_lengths):
x_last[i, :, :] = x[i, seq_len-1, :]
# put on gpu when requested
if on_gpu:
x_last = x_last.cuda()
# turn into torch var
x_last = Variable(x_last)
return x_last
+131 -36
View File
@@ -5,6 +5,27 @@ from pytorch_lightning.root_module.memory import get_gpu_memory_map
import traceback
from pytorch_lightning.root_module.model_saving import TrainerIO
from torch.optim.lr_scheduler import MultiStepLR
from pytorch_lightning.pt_overrides.override_data_parallel import LightningDataParallel
import pdb
try:
from apex import amp
APEX_AVAILABLE = True
except ModuleNotFoundError:
APEX_AVAILABLE = False
def reduce_distributed_output(output, nb_gpus):
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):
@@ -15,7 +36,7 @@ class Trainer(TrainerIO):
cluster=None,
process_position=0,
current_gpu_name=0,
on_gpu=False,
gpus=None,
enable_tqdm=True,
overfit_pct=0.0,
track_grad_norm=-1,
@@ -26,6 +47,9 @@ class Trainer(TrainerIO):
train_percent_check=1.0, val_percent_check=1.0, test_percent_check=1.0, val_check_interval=0.95,
log_save_interval=1, add_log_row_interval=1,
lr_scheduler_milestones=None,
use_amp=False,
check_grad_nans=False,
amp_level='O2',
nb_sanity_val_steps=5):
# Transfer params
@@ -33,7 +57,7 @@ class Trainer(TrainerIO):
self.enable_early_stop = enable_early_stop
self.track_grad_norm = track_grad_norm
self.fast_dev_run = fast_dev_run
self.on_gpu = on_gpu
self.on_gpu = gpus is not None and torch.cuda.is_available()
self.enable_tqdm = enable_tqdm
self.experiment = experiment
self.exp_save_path = experiment.get_data_path(experiment.name, experiment.version)
@@ -51,6 +75,10 @@ class Trainer(TrainerIO):
self.nb_sanity_val_steps = nb_sanity_val_steps
self.lr_scheduler_milestones = [] if lr_scheduler_milestones is None else [int(x.strip()) for x in lr_scheduler_milestones.split(',')]
self.lr_schedulers = []
self.amp_level = amp_level
self.check_grad_nans = check_grad_nans
self.data_parallel_device_ids = gpus
self.data_parallel = gpus is not None and len(gpus) > 0
# training state
self.optimizers = None
@@ -73,6 +101,11 @@ class Trainer(TrainerIO):
self.__determine_data_use_amount(train_percent_check, val_percent_check, test_percent_check, overfit_pct)
print('gpu available: {}, used: {}'.format(torch.cuda.is_available(), self.on_gpu))
# apex test
self.use_amp = use_amp and APEX_AVAILABLE
if self.use_amp:
print('using 16bit precision')
def __determine_data_use_amount(self, train_percent_check, val_percent_check, test_percent_check, overfit_pct):
"""
Use less data for debugging purposes
@@ -86,7 +119,7 @@ class Trainer(TrainerIO):
self.test_percent_check = overfit_pct
def __is_function_implemented(self, f_name):
f_op = getattr(self, f_name, None)
f_op = getattr(self.model, f_name, None)
return callable(f_op)
@property
@@ -101,7 +134,7 @@ class Trainer(TrainerIO):
tqdm_dic.update(self.tqdm_metrics)
return tqdm_dic
def __layout_bookeeping(self):
def __layout_bookeeping(self, model):
# training bookeeping
self.total_batch_nb = 0
self.running_loss = []
@@ -110,21 +143,21 @@ class Trainer(TrainerIO):
self.tqdm_metrics = {}
# determine number of training batches
nb_tng_batches = self.model.nb_batches(self.tng_dataloader)
self.nb_tng_batches = int(nb_tng_batches * self.train_percent_check)
self.nb_tng_batches = model.nb_batches(self.tng_dataloader)
self.nb_tng_batches = int(self.nb_tng_batches * self.train_percent_check)
# determine number of validation batches
nb_val_batches = self.model.nb_batches(self.val_dataloader)
nb_val_batches = int(nb_val_batches * self.val_percent_check)
nb_val_batches = max(1, nb_val_batches)
self.nb_val_batches = nb_val_batches
self.nb_val_batches = model.nb_batches(self.val_dataloader)
self.nb_val_batches = int(self.nb_val_batches * self.val_percent_check)
self.nb_val_batches = max(1, self.nb_val_batches)
self.nb_val_batches = self.nb_val_batches
# determine number of test batches
nb_test_batches = self.model.nb_batches(self.test_dataloader)
self.nb_test_batches = int(nb_test_batches * self.test_percent_check)
self.nb_test_batches = model.nb_batches(self.test_dataloader)
self.nb_test_batches = int(self.nb_test_batches * self.test_percent_check)
# determine when to check validation
self.val_check_batch = int(nb_tng_batches * self.val_check_interval)
self.val_check_batch = int(self.nb_tng_batches * self.val_check_interval)
def __add_tqdm_metrics(self, metrics):
for k, v in metrics.items():
@@ -143,6 +176,7 @@ class Trainer(TrainerIO):
# enable eval mode
model.zero_grad()
model.eval()
model.from_lightning = True
# disable gradients to save memory
torch.set_grad_enabled(False)
@@ -151,19 +185,24 @@ class Trainer(TrainerIO):
outputs = []
# run training
for i, data_batch in enumerate(dataloader):
for batch_i, data_batch in enumerate(dataloader):
if data_batch is None:
continue
# stop short when on fast dev run
if max_batches is not None and i >= max_batches:
if max_batches is not None and batch_i >= max_batches:
break
# -----------------
# RUN VALIDATION STEP
# -----------------
output = model.validation_step(data_batch)
if self.data_parallel:
output = model(data_batch, batch_i)
output = reduce_distributed_output(output, len(self.data_parallel_device_ids))
else:
output = model.validation_step(data_batch, batch_i)
outputs.append(output)
# batch done
@@ -171,13 +210,17 @@ class Trainer(TrainerIO):
self.prog_bar.update(1)
# give model a chance to do something with the outputs
val_results = model.validation_end(outputs)
if self.data_parallel:
val_results = model.module.validation_end(outputs)
else:
val_results = model.validation_end(outputs)
# enable train mode again
model.train()
# enable gradients to save memory
torch.set_grad_enabled(True)
return val_results
def __get_dataloaders(self, model):
@@ -194,18 +237,27 @@ class Trainer(TrainerIO):
# MODEL TRAINING
# -----------------------------
def fit(self, model):
self.model = model
model.trainer = self
# transfer data loaders from model
self.__get_dataloaders(model)
# init training constants
self.__layout_bookeeping()
self.__layout_bookeeping(model)
# CHOOSE OPTIMIZER
# filter out the weights that were done on gpu so we can load on good old cpus
self.optimizers = model.configure_optimizers()
if self.use_amp:
# An example
model, optimizer = amp.initialize(
model, self.optimizers[0], opt_level=self.amp_level,
)
self.optimizers[0] = optimizer
model.trainer = self
# add lr schedulers
if self.lr_scheduler_milestones is not None:
for optimizer in self.optimizers:
@@ -217,7 +269,7 @@ class Trainer(TrainerIO):
# put on gpu if needed
if self.on_gpu:
model = model.cuda()
model = LightningDataParallel(model, device_ids=self.data_parallel_device_ids)
# run tiny validation to make sure program won't crash during val
_ = self.validate(model, self.val_dataloader, max_batches=self.nb_sanity_val_steps)
@@ -232,6 +284,7 @@ class Trainer(TrainerIO):
# ---------------------------
# CORE TRAINING LOOP
# ---------------------------
self.model = model
self.__train()
def __train(self):
@@ -241,11 +294,13 @@ class Trainer(TrainerIO):
for lr_scheduler in self.lr_schedulers:
lr_scheduler.step()
self.model.current_epoch = epoch_nb
model = self.model.module if self.data_parallel else self.model
model.current_epoch = epoch_nb
# hook
if self.__is_function_implemented('on_epoch_start'):
self.model.on_epoch_start()
model = self.model.module if self.data_parallel else self.model
model.on_epoch_start()
self.current_epoch = epoch_nb
self.total_batches = self.nb_tng_batches + self.nb_val_batches
@@ -258,7 +313,9 @@ class Trainer(TrainerIO):
for batch_nb, data_batch in enumerate(self.tng_dataloader):
self.batch_nb = batch_nb
self.global_step += 1
self.model.global_step = self.global_step
model = self.model.module if self.data_parallel else self.model
model.global_step = self.global_step
# stop when the flag is changed or we've gone past the amount requested in the batches
self.total_batch_nb += 1
@@ -269,25 +326,29 @@ class Trainer(TrainerIO):
# ---------------
# RUN TRAIN STEP
# ---------------
self.__run_tng_batch(data_batch)
batch_result = self.__run_tng_batch(data_batch, batch_nb)
early_stop_epoch = batch_result == -1
# ---------------
# RUN VAL STEP
# ---------------
is_val_check_batch = (batch_nb + 1) % self.val_check_batch == 0
if self.fast_dev_run or is_val_check_batch:
if self.fast_dev_run or is_val_check_batch or early_stop_epoch:
self.__run_validation()
# when batch should be saved
if (batch_nb + 1) % self.log_save_interval == 0:
if (batch_nb + 1) % self.log_save_interval == 0 or early_stop_epoch:
self.experiment.save()
# when metrics should be logged
if batch_nb % self.add_log_row_interval == 0:
if batch_nb % self.add_log_row_interval == 0 or early_stop_epoch:
# count items in memory
# nb_params, nb_tensors = count_mem_items()
metrics = self.model.update_tng_log_metrics(self.__tng_tqdm_dic)
if self.data_parallel:
metrics = self.model.module.update_tng_log_metrics(self.__tng_tqdm_dic)
else:
metrics = self.model.update_tng_log_metrics(self.__tng_tqdm_dic)
# add gpu memory
if self.on_gpu:
@@ -296,7 +357,9 @@ class Trainer(TrainerIO):
# add norms
if self.track_grad_norm > 0:
grad_norm_dic = self.model.grad_norm(self.track_grad_norm)
model = self.model.module if self.data_parallel else self.model
grad_norm_dic = model.grad_norm(self.track_grad_norm)
metrics.update(grad_norm_dic)
# log metrics
@@ -305,11 +368,17 @@ class Trainer(TrainerIO):
# hook
if self.__is_function_implemented('on_batch_end'):
self.model.on_batch_end()
model = self.model.module if self.data_parallel else self.model
model.on_batch_end()
# end epoch early
if early_stop_epoch:
break
# hook
if self.__is_function_implemented('on_epoch_end'):
self.model.on_epoch_end()
model = self.model.module if self.data_parallel else self.model
model.on_epoch_end()
# early stopping
if self.enable_early_stop:
@@ -321,24 +390,48 @@ class Trainer(TrainerIO):
if stop:
return
def __run_tng_batch(self, data_batch):
def __run_tng_batch(self, data_batch, batch_nb):
if data_batch is None:
return
return 0
# hook
if self.__is_function_implemented('on_batch_start'):
self.model.on_batch_start()
model = self.model.module if self.data_parallel else self.model
response = model.on_batch_start(data_batch)
if response == -1:
return -1
if self.enable_tqdm:
self.prog_bar.update(1)
# forward pass
# return a scalar value and a dic with tqdm metrics
loss, model_specific_tqdm_metrics_dic = self.model.training_step(data_batch)
if self.data_parallel:
output = self.model(data_batch, batch_nb)
output = reduce_distributed_output(output, len(self.data_parallel_device_ids))
else:
output = self.model.training_step(data_batch, batch_nb)
model_specific_tqdm_metrics_dic = output['tqdm_metrics']
loss = output['loss']
self.__add_tqdm_metrics(model_specific_tqdm_metrics_dic)
# backward pass
loss.backward()
if self.use_amp:
for optimizer in self.optimizers:
with amp.scale_loss(loss, optimizer) as scaled_loss:
scaled_loss.backward()
else:
loss.backward()
if self.check_grad_nans:
model = self.model.module if self.data_parallel else self.model
for param in model.parameters():
print(param.grad.float().sum())
self.batch_loss_value += loss.item()
# gradient update with accumulated gradients
@@ -369,6 +462,8 @@ class Trainer(TrainerIO):
if self.__is_function_implemented('on_batch_end'):
self.model.on_batch_end()
return 0
def __run_validation(self):
# decide if can check epochs
can_check_epoch = (self.current_epoch + 1) % self.check_val_every_n_epoch == 0
@@ -0,0 +1,105 @@
from torch.nn import DataParallel
import threading
import torch
from torch.cuda._utils import _get_device_index
import pdb
def get_a_var(obj):
if isinstance(obj, torch.Tensor):
return obj
if isinstance(obj, list) or isinstance(obj, tuple):
for result in map(get_a_var, obj):
if isinstance(result, torch.Tensor):
return result
if isinstance(obj, dict):
for result in map(get_a_var, obj.items()):
if isinstance(result, torch.Tensor):
return result
return None
class LightningDataParallel(DataParallel):
"""
Override the forward call in lightning so it goes to training and validation step respectively
"""
def parallel_apply(self, replicas, inputs, kwargs):
return parallel_apply(replicas, inputs, kwargs, self.device_ids[:len(replicas)])
def parallel_apply(modules, inputs, kwargs_tup=None, devices=None):
r"""Applies each `module` in :attr:`modules` in parallel on arguments
contained in :attr:`inputs` (positional) and :attr:`kwargs_tup` (keyword)
on each of :attr:`devices`.
Args:
modules (Module): modules to be parallelized
inputs (tensor): inputs to the modules
devices (list of int or torch.device): CUDA devices
:attr:`modules`, :attr:`inputs`, :attr:`kwargs_tup` (if given), and
:attr:`devices` (if given) should all have same length. Moreover, each
element of :attr:`inputs` can either be a single object as the only argument
to a module, or a collection of positional arguments.
"""
assert len(modules) == len(inputs)
if kwargs_tup is not None:
assert len(modules) == len(kwargs_tup)
else:
kwargs_tup = ({},) * len(modules)
if devices is not None:
assert len(modules) == len(devices)
else:
devices = [None] * len(modules)
devices = list(map(lambda x: _get_device_index(x, True), devices))
lock = threading.Lock()
results = {}
grad_enabled = torch.is_grad_enabled()
def _worker(i, module, input, kwargs, device=None):
torch.set_grad_enabled(grad_enabled)
if device is None:
device = get_a_var(input).get_device()
try:
with torch.cuda.device(device):
# this also avoids accidental slicing of `input` if it is a Tensor
if not isinstance(input, (list, tuple)):
input = (input,)
# ---------------
# CHANGE
if module.training:
output = module.training_step(*input, **kwargs)
else:
output = module.validation_step(*input, **kwargs)
# ---------------
with lock:
results[i] = output
except Exception as e:
with lock:
results[i] = e
if len(modules) > 1:
threads = [threading.Thread(target=_worker,
args=(i, module, input, kwargs, device))
for i, (module, input, kwargs, device) in
enumerate(zip(modules, inputs, kwargs_tup, devices))]
for thread in threads:
thread.start()
for thread in threads:
thread.join()
else:
_worker(0, modules[0], inputs[0], kwargs_tup[0], devices[0])
outputs = []
for i in range(len(inputs)):
output = results[i]
if isinstance(output, Exception):
raise output
outputs.append(output)
return outputs
+2 -1
View File
@@ -1,7 +1,7 @@
import torch
class ModelHooks(torch.nn.Module):
def on_batch_start(self):
def on_batch_start(self, data_batch):
pass
def on_batch_end(self):
@@ -18,3 +18,4 @@ class ModelHooks(torch.nn.Module):
def on_post_performance_check(self):
pass
+14 -3
View File
@@ -1,7 +1,8 @@
import torch
import os
import re
import pdb
from pytorch_lightning.pt_overrides.override_data_parallel import LightningDataParallel
class ModelIO(object):
@@ -48,7 +49,8 @@ class TrainerIO(object):
checkpoint['optimizer_states'] = optimizer_states
# request what to save from the model
checkpoint_dict = self.model.get_save_dict()
model = self.model.module if type(self.model) is LightningDataParallel else self.model
checkpoint_dict = model.get_save_dict()
# merge trainer and model saving items
checkpoint.update(checkpoint_dict)
@@ -88,6 +90,7 @@ class TrainerIO(object):
self.early_stop_callback.wait = checkpoint['early_stop_callback_wait']
self.early_stop_callback.patience = checkpoint['early_stop_callback_patience']
self.global_step = checkpoint['global_step']
self.current_epoch = checkpoint['epoch']
# restore the optimizers
optimizer_states = checkpoint['optimizer_states']
@@ -98,6 +101,9 @@ class TrainerIO(object):
# PRIVATE OPS
# ----------------------------------
def hpc_save(self, folderpath, experiment):
# make sure the checkpoint folder exists
os.makedirs(folderpath, exist_ok=True)
# save exp to make sure we get all the metrics
experiment.save()
@@ -125,10 +131,15 @@ class TrainerIO(object):
self.restore_training_state(checkpoint)
# load model state
self.model.load_model_specific(checkpoint)
model = self.model.module if type(self.model) is LightningDataParallel else self.model
model.load_model_specific(checkpoint)
def max_ckpt_in_folder(self, path):
files = os.listdir(path)
files = [x for x in files if 'ckpt_' in x]
if len(files) == 0:
return 0
ckpt_vs = []
for name in files:
name = name.split('ckpt_')[-1]
+11 -6
View File
@@ -24,6 +24,8 @@ class RootModule(GradInformation, ModelIO, OptimizerConfig, ModelHooks):
self.overfit = hparams.overfit
self.gradient_clip = hparams.gradient_clip
self.num = 2
self.trainer = None
self.from_lightning = True
# track if gpu was requested for checkpointing
self.on_gpu = False
@@ -39,8 +41,7 @@ class RootModule(GradInformation, ModelIO, OptimizerConfig, ModelHooks):
if self.on_gpu:
print('running on gpu...')
self.dtype = torch.cuda.FloatTensor
torch.set_default_tensor_type('torch.cuda.FloatTensor')
torch.set_default_tensor_type(hparams.default_tensor_type)
def forward(self, *args, **kwargs):
"""
@@ -51,7 +52,7 @@ class RootModule(GradInformation, ModelIO, OptimizerConfig, ModelHooks):
"""
raise NotImplementedError
def validation_step(self, data_batch):
def validation_step(self, data_batch, batch_nb):
"""
return whatever outputs will need to be aggregated in validation_end
:param data_batch:
@@ -67,7 +68,7 @@ class RootModule(GradInformation, ModelIO, OptimizerConfig, ModelHooks):
"""
raise NotImplementedError
def training_step(self, data_batch):
def training_step(self, data_batch, batch_nb):
"""
return loss, dict with metrics for tqdm
:param data_batch:
@@ -150,19 +151,23 @@ class RootModule(GradInformation, ModelIO, OptimizerConfig, ModelHooks):
return 0, 0
@classmethod
def load_from_metrics(cls, weights_path, tags_csv, on_gpu):
def load_from_metrics(cls, weights_path, tags_csv, on_gpu, map_location=None):
"""
Primary way of loading model from csv weights path
:param weights_path:
:param tags_csv:
:param on_gpu:
:param map_location: dic for mapping storage {'cuda:1':'cuda:0'}
:return:
"""
hparams = load_hparams_from_tags_csv(tags_csv)
hparams.__setattr__('on_gpu', on_gpu)
if on_gpu:
checkpoint = torch.load(weights_path)
if map_location is not None:
checkpoint = torch.load(weights_path, map_location=map_location)
else:
checkpoint = torch.load(weights_path)
else:
checkpoint = torch.load(weights_path, map_location=lambda storage, loc: storage)
-1
View File
@@ -9,7 +9,6 @@ from pytorch_lightning.utils.arg_parse import add_default_args
from time import sleep
from pytorch_lightning.utils.pt_callbacks import EarlyStopping, ModelCheckpoint
SEED = 2334
torch.manual_seed(SEED)
np.random.seed(SEED)
+11 -6
View File
@@ -1,3 +1,5 @@
import pdb
def add_default_args(parser, root_dir, rand_seed=None, possible_model_names=None):
# tng, test, val check intervals
@@ -45,10 +47,13 @@ def add_default_args(parser, root_dir, rand_seed=None, possible_model_names=None
parser.add_argument('--log_stdout', dest='log_stdout', action='store_true')
# GPU
parser.add_argument('--per_experiment_nb_gpus', default=1, type=int)
parser.add_argument('--gpus', default='0', type=str)
parser.add_argument('--gpus', default=None, type=str)
parser.add_argument('--single_run_gpu', dest='single_run_gpu', action='store_true')
parser.add_argument('--disable_cuda', dest='disable_cuda', action='store_true')
parser.add_argument('--default_tensor_type', default='torch.cuda.FloatTensor', type=str)
parser.add_argument('--use_amp', dest='use_amp', action='store_true')
parser.add_argument('--check_grad_nans', dest='check_grad_nans', action='store_true')
parser.add_argument('--amp_level', default='O2',type=str)
# run on hpc
parser.add_argument('--on_cluster', dest='on_cluster', action='store_true')
@@ -63,9 +68,9 @@ def add_default_args(parser, root_dir, rand_seed=None, possible_model_names=None
if rand_seed is not None:
parser.add_argument('--random_seed', default=rand_seed, type=int)
parser.add_argument('--live', dest='live', action='store_true', help='runs on gpu without cluster')
parser.add_argument('--enable_debug', dest='debug', action='store_true', help='enables/disables test tube')
parser.add_argument('--enable_local', dest='local', action='store_true', help='enables local tng')
parser.add_argument('--interactive', dest='interactive', action='store_true', help='runs on gpu without cluster')
parser.add_argument('--debug', dest='debug', action='store_true', help='enables/disables test tube')
parser.add_argument('--local', dest='local', action='store_true', help='enables local tng')
# optimizer
parser.add_argument('--lr_scheduler_milestones', default=None, type=str)
+5 -8
View File
@@ -13,11 +13,7 @@ class PretrainedEmbedding(torch.nn.Embedding):
>>> emb = PretrainedEmbedding(embedding_path='glove.840B.300d.txt',embedding_dim=300, task_vocab={'hello': 1, 'world': 2})
>>> data = torch.Tensor([[0, 1], [0, 2]]).long()
>>> embedded = emb(data)
tensor([[[ 0.0000, 0.0000, 0.0000, ..., 0.0000, 0.0000, 0.0000],
[ 0.2523, 0.1018, -0.6748, ..., 0.1787, -0.5192, 0.3359]],
[[ 0.0000, 0.0000, 0.0000, ..., 0.0000, 0.0000, 0.0000],
[-0.0067, 0.2224, 0.2771, ..., 0.0594, 0.0014, 0.0987]]])
:param embedding_path:
@@ -37,7 +33,8 @@ class PretrainedEmbedding(torch.nn.Embedding):
self.weight = new_emb.weight
# apply freeze
self.weight.requires_grad = not freeze
should_freeze = not freeze
self.weight.requires_grad = should_freeze
def __load_task_specific_embeddings(self, vocab_words, embedding_path, emb_dim, freeze):
"""
@@ -97,11 +94,11 @@ class PretrainedEmbedding(torch.nn.Embedding):
if __name__ == '__main__':
emb = PretrainedEmbedding(
embedding_path='/Users/waf/Developer/NGV/research-fermat/fermat/.vector_cache/glove.840B.300d.txt',
embedding_path='/Users/waf/Developer',
embedding_dim=300,
task_vocab={'hello': 1, 'world': 2}
)
data = torch.Tensor([[0, 1], [0, 2]]).long()
embedded = emb(data)
print(embedded)
print(embedded)
+3
View File
@@ -1,5 +1,6 @@
import numpy as np
import os, shutil
from pytorch_lightning.pt_overrides.override_data_parallel import LightningDataParallel
class Callback(object):
@@ -33,6 +34,8 @@ class Callback(object):
self.params = params
def set_model(self, model):
if type(model) is LightningDataParallel:
model = model.module
self.model = model
def on_epoch_begin(self, epoch, logs=None):
+21
View File
@@ -0,0 +1,21 @@
[tool:pytest]
norecursedirs =
.git
dist
build
python_files =
test_*.py
doctest_plus = disabled
addopts = --strict
markers =
slow
remote_data
filterwarnings
[pycodestyle]
ignore = E731,W504
max-line-length = 120
[flake8]
ignore = E731,W504,F401,F841
max-line-length = 120
+25 -9
View File
@@ -2,12 +2,28 @@
from setuptools import setup, find_packages
setup(name='pytorch-lightning',
version='0.0.2',
description='Rapid research framework',
author='',
author_email='',
url='https://github.com/williamFalcon/pytorch-lightning',
install_requires=['test-tube', 'torch', 'tqdm'],
packages=find_packages()
)
# https://packaging.python.org/guides/single-sourcing-package-version/
# http://blog.ionelmc.ro/2014/05/25/python-packaging/
setup(
name="pytorch-lightning",
version='0.11',
description="The Keras for ML researchers using PyTorch",
author="William Falcon",
author_email="waf2107@columbia.edu",
url="https://github.com/williamFalcon/pytorch-lightning",
download_url="https://github.com/williamFalcon/pytorch-lightning",
license="MIT",
keywords=["deep learning", "pytorch", "AI"],
python_requires=">=3.5",
install_requires=[
"torch>=1.0.0",
"tqdm",
"test-tube",
],
packages=find_packages(),
long_description=open("README.md", encoding="utf-8").read(),
long_description_content_type='text/markdown',
include_package_data=True,
zip_safe=False,
)
-65
View File
@@ -1,65 +0,0 @@
# Testing setup
## A. Enable CircleCI for your project
1. Integrate CircleCI by clicking "Set up Project" at [this link](https://circleci.com/add-projects/gh/NextGenVest).
## B. Add your own tests
1. In the /tests, emulate exactly the folder structure for your module found under /bot_seed
2. To create a test for file ```/bot_seed/folder/example.py```:
- create the file ```/tests/folder/example_test.py```
- notice the **_test**
- notice the mirror path under **/tests**
3. Your ```example_test.py``` file should have these main components
```python
# example.py
def function_i_want_to_test(x):
return x*2
def square(x):
return x*x
```
```python
# example_test.py
import pytest
# do whatever imports you need
from app.bot_seed.folder.example import function_i_want_to_test, square
def test_function_i_want_to_test():
answer = function_i_want_to_test(4)
assert answer == 8
# -----------------------------------
# Your function must start with test_
# -----------------------------------
def test_square():
answer = square(3)
assert answer == 9
# -----------------------------------
# boilerplate (link this file to pytest)
# -----------------------------------
if __name__ == '__main__':
pytest.main([__file__])
```
## C. Add build passing badge
1. Create a CircleCI status token:
- Go here: https://circleci.com/gh/NextGenVest/your-project-name/edit#api
- Click create token
- Select status
- Type "badge status"
2. Get a copy of the markdown code:
- Go here: https://circleci.com/gh/NextGenVest/your-project-name/edit#badges
- Select master
- Select "badge status" token
- Select image URL
- Copy the image url link and change the html at the top of the root README.md file for your project
View File
View File
@@ -1,13 +0,0 @@
import pytest
"""
Example test to show how to add a test for anything in the project.
Look at the README for more instructions
"""
def test_cube():
assert 27 == 27
if __name__ == '__main__':
pytest.main([__file__])