mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-09 11:32:07 +08:00
Deployed f0af138 with MkDocs version: 1.0.4
This commit is contained in:
+182
-65
@@ -602,9 +602,11 @@
|
||||
|
||||
<h3 id="template-model-definition">Template model definition</h3>
|
||||
<p>In 99% of cases you want to just copy <a href="https://github.com/williamFalcon/pytorch-lightning/blob/master/examples/new_project_templates/lightning_module_template.py">this template</a> to start a new lightningModule and change the core of what your model is actually trying to do.</p>
|
||||
<pre><code class="bash"># get a copy of the module template
|
||||
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
|
||||
2</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># get a copy of the module template</span>
|
||||
wget https://raw.githubusercontent.com/williamFalcon/pytorch-lightning/master/examples/new_project_templates/lightning_module_template.py
|
||||
</code></pre>
|
||||
</pre></div>
|
||||
</td></tr></table>
|
||||
|
||||
<hr />
|
||||
<h3 id="trainer-example">Trainer Example</h3>
|
||||
@@ -612,22 +614,37 @@ wget https://raw.githubusercontent.com/williamFalcon/pytorch-lightning/master/ex
|
||||
<p>Normally, we want to let the __main__ function start the training.
|
||||
Inside the main we parse training arguments with whatever hyperparameters we want. Your LightningModule will have a
|
||||
chance to add hyperparameters. </p>
|
||||
<pre><code class="python">from test_tube import HyperOptArgumentParser
|
||||
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre> 1
|
||||
2
|
||||
3
|
||||
4
|
||||
5
|
||||
6
|
||||
7
|
||||
8
|
||||
9
|
||||
10
|
||||
11
|
||||
12
|
||||
13
|
||||
14
|
||||
15</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="kn">from</span> <span class="nn">test_tube</span> <span class="kn">import</span> <span class="n">HyperOptArgumentParser</span>
|
||||
|
||||
if __name__ == '__main__':
|
||||
<span class="k">if</span> <span class="vm">__name__</span> <span class="o">==</span> <span class="s1">'__main__'</span><span class="p">:</span>
|
||||
|
||||
# use default args given by lightning
|
||||
root_dir = os.path.split(os.path.dirname(sys.modules['__main__'].__file__))[0]
|
||||
parent_parser = HyperOptArgumentParser(strategy='random_search', add_help=False)
|
||||
add_default_args(parent_parser, root_dir)
|
||||
<span class="c1"># use default args given by lightning</span>
|
||||
<span class="n">root_dir</span> <span class="o">=</span> <span class="n">os</span><span class="o">.</span><span class="n">path</span><span class="o">.</span><span class="n">split</span><span class="p">(</span><span class="n">os</span><span class="o">.</span><span class="n">path</span><span class="o">.</span><span class="n">dirname</span><span class="p">(</span><span class="n">sys</span><span class="o">.</span><span class="n">modules</span><span class="p">[</span><span class="s1">'__main__'</span><span class="p">]</span><span class="o">.</span><span class="vm">__file__</span><span class="p">))[</span><span class="mi">0</span><span class="p">]</span>
|
||||
<span class="n">parent_parser</span> <span class="o">=</span> <span class="n">HyperOptArgumentParser</span><span class="p">(</span><span class="n">strategy</span><span class="o">=</span><span class="s1">'random_search'</span><span class="p">,</span> <span class="n">add_help</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>
|
||||
<span class="n">add_default_args</span><span class="p">(</span><span class="n">parent_parser</span><span class="p">,</span> <span class="n">root_dir</span><span class="p">)</span>
|
||||
|
||||
# allow model to overwrite or extend args
|
||||
parser = ExampleModel.add_model_specific_args(parent_parser)
|
||||
hyperparams = parser.parse_args()
|
||||
<span class="c1"># allow model to overwrite or extend args</span>
|
||||
<span class="n">parser</span> <span class="o">=</span> <span class="n">ExampleModel</span><span class="o">.</span><span class="n">add_model_specific_args</span><span class="p">(</span><span class="n">parent_parser</span><span class="p">)</span>
|
||||
<span class="n">hyperparams</span> <span class="o">=</span> <span class="n">parser</span><span class="o">.</span><span class="n">parse_args</span><span class="p">()</span>
|
||||
|
||||
# train model
|
||||
main(hyperparams)
|
||||
</code></pre>
|
||||
<span class="c1"># train model</span>
|
||||
<span class="n">main</span><span class="p">(</span><span class="n">hyperparams</span><span class="p">)</span>
|
||||
</pre></div>
|
||||
</td></tr></table>
|
||||
|
||||
<p><strong>Main Function</strong> </p>
|
||||
<p>The main function is your entry into the program. This is where you init your model, checkpoint directory, and launch the training.
|
||||
@@ -635,7 +652,58 @@ The main function should have 3 arguments: <br />
|
||||
- hparams: a configuration of hyperparameters. <br />
|
||||
- slurm_manager: Slurm cluster manager object (can be None)
|
||||
- dict: for you to return any values you want (useful in meta-learning, otherwise set to _) </p>
|
||||
<pre><code>def main(hparams, cluster, results_dict):
|
||||
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre> 1
|
||||
2
|
||||
3
|
||||
4
|
||||
5
|
||||
6
|
||||
7
|
||||
8
|
||||
9
|
||||
10
|
||||
11
|
||||
12
|
||||
13
|
||||
14
|
||||
15
|
||||
16
|
||||
17
|
||||
18
|
||||
19
|
||||
20
|
||||
21
|
||||
22
|
||||
23
|
||||
24
|
||||
25
|
||||
26
|
||||
27
|
||||
28
|
||||
29
|
||||
30
|
||||
31
|
||||
32
|
||||
33
|
||||
34
|
||||
35
|
||||
36
|
||||
37
|
||||
38
|
||||
39
|
||||
40
|
||||
41
|
||||
42
|
||||
43
|
||||
44
|
||||
45
|
||||
46
|
||||
47
|
||||
48
|
||||
49
|
||||
50
|
||||
51
|
||||
52</pre></div></td><td class="code"><div class="codehilite"><pre><span></span>def main(hparams, cluster, results_dict):
|
||||
"""
|
||||
Main training routine specific for this project
|
||||
:param hparams:
|
||||
@@ -644,12 +712,12 @@ The main function should have 3 arguments: <br />
|
||||
# init experiment
|
||||
log_dir = os.path.dirname(os.path.realpath(__file__))
|
||||
exp = Experiment(
|
||||
name='test_tube_exp',
|
||||
name='test_tube_exp',
|
||||
debug=True,
|
||||
save_dir=log_dir,
|
||||
version=0,
|
||||
autosave=False,
|
||||
description='test demo'
|
||||
description='test demo'
|
||||
)
|
||||
|
||||
# set the hparams for the experiment
|
||||
@@ -667,7 +735,7 @@ The main function should have 3 arguments: <br />
|
||||
mode=hparams.early_stop_mode
|
||||
)
|
||||
|
||||
model_save_path = '{}/{}/{}'.format(hparams.model_save_path, exp.name, exp.version)
|
||||
model_save_path = '{}/{}/{}'.format(hparams.model_save_path, exp.name, exp.version)
|
||||
checkpoint = ModelCheckpoint(
|
||||
filepath=model_save_path,
|
||||
save_function=None,
|
||||
@@ -687,73 +755,122 @@ The main function should have 3 arguments: <br />
|
||||
|
||||
# train model
|
||||
trainer.fit(model)
|
||||
</code></pre>
|
||||
</pre></div>
|
||||
</td></tr></table>
|
||||
|
||||
<p>The <strong>main</strong> function will start training on your <strong>main</strong> function. If you use the HyperParameterOptimizer
|
||||
in hyper parameter optimization mode, this main function will get one set of hyperparameters. If you use it as a simple
|
||||
argument parser you get the default arguments in the argument parser.</p>
|
||||
<p>So, calling main(hyperparams) runs the model with the default argparse arguments. </p>
|
||||
<pre><code class="python">main(hyperparams)
|
||||
</code></pre>
|
||||
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="n">main</span><span class="p">(</span><span class="n">hyperparams</span><span class="p">)</span>
|
||||
</pre></div>
|
||||
</td></tr></table>
|
||||
|
||||
<hr />
|
||||
<h4 id="cpu-hyperparameter-search">CPU hyperparameter search</h4>
|
||||
<pre><code class="python"># run a grid search over 20 hyperparameter combinations.
|
||||
hyperparams.optimize_parallel_cpu(
|
||||
main_local,
|
||||
nb_trials=20,
|
||||
nb_workers=1
|
||||
)
|
||||
</code></pre>
|
||||
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
|
||||
2
|
||||
3
|
||||
4
|
||||
5
|
||||
6</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># run a grid search over 20 hyperparameter combinations.</span>
|
||||
<span class="n">hyperparams</span><span class="o">.</span><span class="n">optimize_parallel_cpu</span><span class="p">(</span>
|
||||
<span class="n">main_local</span><span class="p">,</span>
|
||||
<span class="n">nb_trials</span><span class="o">=</span><span class="mi">20</span><span class="p">,</span>
|
||||
<span class="n">nb_workers</span><span class="o">=</span><span class="mi">1</span>
|
||||
<span class="p">)</span>
|
||||
</pre></div>
|
||||
</td></tr></table>
|
||||
|
||||
<hr />
|
||||
<h4 id="hyperparameter-search-on-a-single-or-multiple-gpus">Hyperparameter search on a single or multiple GPUs</h4>
|
||||
<pre><code class="python"># run a grid search over 20 hyperparameter combinations.
|
||||
hyperparams.optimize_parallel_gpu(
|
||||
main_local,
|
||||
nb_trials=20,
|
||||
nb_workers=1,
|
||||
gpus=[0,1,2,3]
|
||||
)
|
||||
</code></pre>
|
||||
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
|
||||
2
|
||||
3
|
||||
4
|
||||
5
|
||||
6
|
||||
7</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># run a grid search over 20 hyperparameter combinations.</span>
|
||||
<span class="n">hyperparams</span><span class="o">.</span><span class="n">optimize_parallel_gpu</span><span class="p">(</span>
|
||||
<span class="n">main_local</span><span class="p">,</span>
|
||||
<span class="n">nb_trials</span><span class="o">=</span><span class="mi">20</span><span class="p">,</span>
|
||||
<span class="n">nb_workers</span><span class="o">=</span><span class="mi">1</span><span class="p">,</span>
|
||||
<span class="n">gpus</span><span class="o">=</span><span class="p">[</span><span class="mi">0</span><span class="p">,</span><span class="mi">1</span><span class="p">,</span><span class="mi">2</span><span class="p">,</span><span class="mi">3</span><span class="p">]</span>
|
||||
<span class="p">)</span>
|
||||
</pre></div>
|
||||
</td></tr></table>
|
||||
|
||||
<hr />
|
||||
<h4 id="hyperparameter-search-on-a-slurm-hpc-cluster">Hyperparameter search on a SLURM HPC cluster</h4>
|
||||
<pre><code class="python">def optimize_on_cluster(hyperparams):
|
||||
# enable cluster training
|
||||
cluster = SlurmCluster(
|
||||
hyperparam_optimizer=hyperparams,
|
||||
log_path=hyperparams.tt_save_path,
|
||||
test_tube_exp_name=hyperparams.tt_name
|
||||
)
|
||||
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre> 1
|
||||
2
|
||||
3
|
||||
4
|
||||
5
|
||||
6
|
||||
7
|
||||
8
|
||||
9
|
||||
10
|
||||
11
|
||||
12
|
||||
13
|
||||
14
|
||||
15
|
||||
16
|
||||
17
|
||||
18
|
||||
19
|
||||
20
|
||||
21
|
||||
22
|
||||
23
|
||||
24
|
||||
25
|
||||
26
|
||||
27
|
||||
28
|
||||
29
|
||||
30
|
||||
31
|
||||
32
|
||||
33
|
||||
34</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="k">def</span> <span class="nf">optimize_on_cluster</span><span class="p">(</span><span class="n">hyperparams</span><span class="p">):</span>
|
||||
<span class="c1"># enable cluster training</span>
|
||||
<span class="n">cluster</span> <span class="o">=</span> <span class="n">SlurmCluster</span><span class="p">(</span>
|
||||
<span class="n">hyperparam_optimizer</span><span class="o">=</span><span class="n">hyperparams</span><span class="p">,</span>
|
||||
<span class="n">log_path</span><span class="o">=</span><span class="n">hyperparams</span><span class="o">.</span><span class="n">tt_save_path</span><span class="p">,</span>
|
||||
<span class="n">test_tube_exp_name</span><span class="o">=</span><span class="n">hyperparams</span><span class="o">.</span><span class="n">tt_name</span>
|
||||
<span class="p">)</span>
|
||||
|
||||
# email for cluster coms
|
||||
cluster.notify_job_status(email='add_email_here', on_done=True, on_fail=True)
|
||||
<span class="c1"># email for cluster coms</span>
|
||||
<span class="n">cluster</span><span class="o">.</span><span class="n">notify_job_status</span><span class="p">(</span><span class="n">email</span><span class="o">=</span><span class="s1">'add_email_here'</span><span class="p">,</span> <span class="n">on_done</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span> <span class="n">on_fail</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>
|
||||
|
||||
# configure cluster
|
||||
cluster.per_experiment_nb_gpus = hyperparams.per_experiment_nb_gpus
|
||||
cluster.job_time = '48:00:00'
|
||||
cluster.gpu_type = '1080ti'
|
||||
cluster.memory_mb_per_node = 48000
|
||||
<span class="c1"># configure cluster</span>
|
||||
<span class="n">cluster</span><span class="o">.</span><span class="n">per_experiment_nb_gpus</span> <span class="o">=</span> <span class="n">hyperparams</span><span class="o">.</span><span class="n">per_experiment_nb_gpus</span>
|
||||
<span class="n">cluster</span><span class="o">.</span><span class="n">job_time</span> <span class="o">=</span> <span class="s1">'48:00:00'</span>
|
||||
<span class="n">cluster</span><span class="o">.</span><span class="n">gpu_type</span> <span class="o">=</span> <span class="s1">'1080ti'</span>
|
||||
<span class="n">cluster</span><span class="o">.</span><span class="n">memory_mb_per_node</span> <span class="o">=</span> <span class="mi">48000</span>
|
||||
|
||||
# any modules for code to run in env
|
||||
cluster.add_command('source activate pytorch_lightning')
|
||||
<span class="c1"># any modules for code to run in env</span>
|
||||
<span class="n">cluster</span><span class="o">.</span><span class="n">add_command</span><span class="p">(</span><span class="s1">'source activate pytorch_lightning'</span><span class="p">)</span>
|
||||
|
||||
# name of exp
|
||||
job_display_name = hyperparams.tt_name.split('_')[0]
|
||||
job_display_name = job_display_name[0:3]
|
||||
<span class="c1"># name of exp</span>
|
||||
<span class="n">job_display_name</span> <span class="o">=</span> <span class="n">hyperparams</span><span class="o">.</span><span class="n">tt_name</span><span class="o">.</span><span class="n">split</span><span class="p">(</span><span class="s1">'_'</span><span class="p">)[</span><span class="mi">0</span><span class="p">]</span>
|
||||
<span class="n">job_display_name</span> <span class="o">=</span> <span class="n">job_display_name</span><span class="p">[</span><span class="mi">0</span><span class="p">:</span><span class="mi">3</span><span class="p">]</span>
|
||||
|
||||
# run hopt
|
||||
print('submitting jobs...')
|
||||
cluster.optimize_parallel_cluster_gpu(
|
||||
main,
|
||||
nb_trials=hyperparams.nb_hopt_trials,
|
||||
job_name=job_display_name
|
||||
)
|
||||
<span class="c1"># run hopt</span>
|
||||
<span class="k">print</span><span class="p">(</span><span class="s1">'submitting jobs...'</span><span class="p">)</span>
|
||||
<span class="n">cluster</span><span class="o">.</span><span class="n">optimize_parallel_cluster_gpu</span><span class="p">(</span>
|
||||
<span class="n">main</span><span class="p">,</span>
|
||||
<span class="n">nb_trials</span><span class="o">=</span><span class="n">hyperparams</span><span class="o">.</span><span class="n">nb_hopt_trials</span><span class="p">,</span>
|
||||
<span class="n">job_name</span><span class="o">=</span><span class="n">job_display_name</span>
|
||||
<span class="p">)</span>
|
||||
|
||||
# run cluster hyperparameter search
|
||||
optimize_on_cluster(hyperparams)
|
||||
</code></pre>
|
||||
<span class="c1"># run cluster hyperparameter search </span>
|
||||
<span class="n">optimize_on_cluster</span><span class="p">(</span><span class="n">hyperparams</span><span class="p">)</span>
|
||||
</pre></div>
|
||||
</td></tr></table>
|
||||
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user