Deployed f0af138 with MkDocs version: 1.0.4

This commit is contained in:
William Falcon
2019-08-13 15:21:24 -05:00
parent d32c548cf8
commit 9b8a06bfdb
15 changed files with 1240 additions and 598 deletions
@@ -988,61 +988,115 @@
</ul>
<hr />
<h3 id="minimal-example">Minimal example</h3>
<pre><code class="python">import os
import torch
from torch.nn import functional as F
from torch.utils.data import DataLoader
from torchvision.datasets import MNIST
import torchvision.transforms as transforms
<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
53
54</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="kn">import</span> <span class="nn">os</span>
<span class="kn">import</span> <span class="nn">torch</span>
<span class="kn">from</span> <span class="nn">torch.nn</span> <span class="kn">import</span> <span class="n">functional</span> <span class="k">as</span> <span class="n">F</span>
<span class="kn">from</span> <span class="nn">torch.utils.data</span> <span class="kn">import</span> <span class="n">DataLoader</span>
<span class="kn">from</span> <span class="nn">torchvision.datasets</span> <span class="kn">import</span> <span class="n">MNIST</span>
<span class="kn">import</span> <span class="nn">torchvision.transforms</span> <span class="kn">as</span> <span class="nn">transforms</span>
import pytorch_lightning as pl
<span class="kn">import</span> <span class="nn">pytorch_lightning</span> <span class="kn">as</span> <span class="nn">pl</span>
class CoolModel(pl.LightningModule):
<span class="k">class</span> <span class="nc">CoolModel</span><span class="p">(</span><span class="n">pl</span><span class="o">.</span><span class="n">LightningModule</span><span class="p">):</span>
def __init__(self):
super(CoolModel, self).__init__()
# not the best model...
self.l1 = torch.nn.Linear(28 * 28, 10)
<span class="k">def</span> <span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
<span class="nb">super</span><span class="p">(</span><span class="n">CoolModel</span><span class="p">,</span> <span class="bp">self</span><span class="p">)</span><span class="o">.</span><span class="fm">__init__</span><span class="p">()</span>
<span class="c1"># not the best model...</span>
<span class="bp">self</span><span class="o">.</span><span class="n">l1</span> <span class="o">=</span> <span class="n">torch</span><span class="o">.</span><span class="n">nn</span><span class="o">.</span><span class="n">Linear</span><span class="p">(</span><span class="mi">28</span> <span class="o">*</span> <span class="mi">28</span><span class="p">,</span> <span class="mi">10</span><span class="p">)</span>
def forward(self, x):
return torch.relu(self.l1(x.view(x.size(0), -1)))
<span class="k">def</span> <span class="nf">forward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">x</span><span class="p">):</span>
<span class="k">return</span> <span class="n">torch</span><span class="o">.</span><span class="n">relu</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">l1</span><span class="p">(</span><span class="n">x</span><span class="o">.</span><span class="n">view</span><span class="p">(</span><span class="n">x</span><span class="o">.</span><span class="n">size</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span> <span class="o">-</span><span class="mi">1</span><span class="p">)))</span>
def training_step(self, batch, batch_nb):
# REQUIRED
x, y = batch
y_hat = self.forward(x)
return {'loss': F.cross_entropy(y_hat, y)(y_hat, y)}
<span class="k">def</span> <span class="nf">training_step</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">batch</span><span class="p">,</span> <span class="n">batch_nb</span><span class="p">):</span>
<span class="c1"># REQUIRED</span>
<span class="n">x</span><span class="p">,</span> <span class="n">y</span> <span class="o">=</span> <span class="n">batch</span>
<span class="n">y_hat</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">forward</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
<span class="k">return</span> <span class="p">{</span><span class="s1">&#39;loss&#39;</span><span class="p">:</span> <span class="n">F</span><span class="o">.</span><span class="n">cross_entropy</span><span class="p">(</span><span class="n">y_hat</span><span class="p">,</span> <span class="n">y</span><span class="p">)(</span><span class="n">y_hat</span><span class="p">,</span> <span class="n">y</span><span class="p">)}</span>
def validation_step(self, batch, batch_nb):
# OPTIONAL
x, y = batch
y_hat = self.forward(x)
return {'val_loss': F.cross_entropy(y_hat, y)(y_hat, y)}
<span class="k">def</span> <span class="nf">validation_step</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">batch</span><span class="p">,</span> <span class="n">batch_nb</span><span class="p">):</span>
<span class="c1"># OPTIONAL</span>
<span class="n">x</span><span class="p">,</span> <span class="n">y</span> <span class="o">=</span> <span class="n">batch</span>
<span class="n">y_hat</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">forward</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
<span class="k">return</span> <span class="p">{</span><span class="s1">&#39;val_loss&#39;</span><span class="p">:</span> <span class="n">F</span><span class="o">.</span><span class="n">cross_entropy</span><span class="p">(</span><span class="n">y_hat</span><span class="p">,</span> <span class="n">y</span><span class="p">)(</span><span class="n">y_hat</span><span class="p">,</span> <span class="n">y</span><span class="p">)}</span>
def validation_end(self, outputs):
# OPTIONAL
avg_loss = torch.stack([x['val_loss'] for x in outputs]).mean()
return {'avg_val_loss': avg_loss}
<span class="k">def</span> <span class="nf">validation_end</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">outputs</span><span class="p">):</span>
<span class="c1"># OPTIONAL</span>
<span class="n">avg_loss</span> <span class="o">=</span> <span class="n">torch</span><span class="o">.</span><span class="n">stack</span><span class="p">([</span><span class="n">x</span><span class="p">[</span><span class="s1">&#39;val_loss&#39;</span><span class="p">]</span> <span class="k">for</span> <span class="n">x</span> <span class="ow">in</span> <span class="n">outputs</span><span class="p">])</span><span class="o">.</span><span class="n">mean</span><span class="p">()</span>
<span class="k">return</span> <span class="p">{</span><span class="s1">&#39;avg_val_loss&#39;</span><span class="p">:</span> <span class="n">avg_loss</span><span class="p">}</span>
def configure_optimizers(self):
# REQUIRED
return [torch.optim.Adam(self.parameters(), lr=0.02)]
<span class="k">def</span> <span class="nf">configure_optimizers</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
<span class="c1"># REQUIRED</span>
<span class="k">return</span> <span class="p">[</span><span class="n">torch</span><span class="o">.</span><span class="n">optim</span><span class="o">.</span><span class="n">Adam</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">parameters</span><span class="p">(),</span> <span class="n">lr</span><span class="o">=</span><span class="mf">0.02</span><span class="p">)]</span>
@pl.data_loader
def tng_dataloader(self):
return DataLoader(MNIST(os.getcwd(), train=True, download=True, transform=transforms.ToTensor()), batch_size=32)
<span class="nd">@pl.data_loader</span>
<span class="k">def</span> <span class="nf">tng_dataloader</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
<span class="k">return</span> <span class="n">DataLoader</span><span class="p">(</span><span class="n">MNIST</span><span class="p">(</span><span class="n">os</span><span class="o">.</span><span class="n">getcwd</span><span class="p">(),</span> <span class="n">train</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span> <span class="n">download</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span> <span class="n">transform</span><span class="o">=</span><span class="n">transforms</span><span class="o">.</span><span class="n">ToTensor</span><span class="p">()),</span> <span class="n">batch_size</span><span class="o">=</span><span class="mi">32</span><span class="p">)</span>
@pl.data_loader
def val_dataloader(self):
# OPTIONAL
# can also return a list of val dataloaders
return DataLoader(MNIST(os.getcwd(), train=True, download=True, transform=transforms.ToTensor()), batch_size=32)
<span class="nd">@pl.data_loader</span>
<span class="k">def</span> <span class="nf">val_dataloader</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
<span class="c1"># OPTIONAL</span>
<span class="c1"># can also return a list of val dataloaders</span>
<span class="k">return</span> <span class="n">DataLoader</span><span class="p">(</span><span class="n">MNIST</span><span class="p">(</span><span class="n">os</span><span class="o">.</span><span class="n">getcwd</span><span class="p">(),</span> <span class="n">train</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span> <span class="n">download</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span> <span class="n">transform</span><span class="o">=</span><span class="n">transforms</span><span class="o">.</span><span class="n">ToTensor</span><span class="p">()),</span> <span class="n">batch_size</span><span class="o">=</span><span class="mi">32</span><span class="p">)</span>
@pl.data_loader
def test_dataloader(self):
# OPTIONAL
return DataLoader(MNIST(os.getcwd(), train=True, download=True, transform=transforms.ToTensor()), batch_size=32)
</code></pre>
<span class="nd">@pl.data_loader</span>
<span class="k">def</span> <span class="nf">test_dataloader</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
<span class="c1"># OPTIONAL</span>
<span class="k">return</span> <span class="n">DataLoader</span><span class="p">(</span><span class="n">MNIST</span><span class="p">(</span><span class="n">os</span><span class="o">.</span><span class="n">getcwd</span><span class="p">(),</span> <span class="n">train</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span> <span class="n">download</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span> <span class="n">transform</span><span class="o">=</span><span class="n">transforms</span><span class="o">.</span><span class="n">ToTensor</span><span class="p">()),</span> <span class="n">batch_size</span><span class="o">=</span><span class="mi">32</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<hr />
<h3 id="how-do-these-methods-fit-into-the-broader-training">How do these methods fit into the broader training?</h3>
@@ -1055,8 +1109,9 @@ class CoolModel(pl.LightningModule):
<h2 id="required-methods">Required Methods</h2>
<h3 id="training_step">training_step</h3>
<pre><code class="python">def training_step(self, data_batch, batch_nb)
</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="k">def</span> <span class="nf">training_step</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">data_batch</span><span class="p">,</span> <span class="n">batch_nb</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<p>In this step you'd normally do the forward pass and calculate the loss for a batch. You can also do fancier things like multiple forward passes or something specific to your model.</p>
<p><strong>Params</strong> </p>
@@ -1102,57 +1157,90 @@ class CoolModel(pl.LightningModule):
</tbody>
</table>
<p><strong>Example</strong></p>
<pre><code class="python">def training_step(self, data_batch, batch_nb):
x, y, z = data_batch
<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</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="k">def</span> <span class="nf">training_step</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">data_batch</span><span class="p">,</span> <span class="n">batch_nb</span><span class="p">):</span>
<span class="n">x</span><span class="p">,</span> <span class="n">y</span><span class="p">,</span> <span class="n">z</span> <span class="o">=</span> <span class="n">data_batch</span>
# implement your own
out = self.forward(x)
loss = self.loss(out, x)
<span class="c1"># implement your own</span>
<span class="n">out</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">forward</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
<span class="n">loss</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">loss</span><span class="p">(</span><span class="n">out</span><span class="p">,</span> <span class="n">x</span><span class="p">)</span>
output = {
'loss': loss, # required
'prog': {'tng_loss': loss, 'batch_nb': batch_nb} # optional
}
<span class="n">output</span> <span class="o">=</span> <span class="p">{</span>
<span class="s1">&#39;loss&#39;</span><span class="p">:</span> <span class="n">loss</span><span class="p">,</span> <span class="c1"># required</span>
<span class="s1">&#39;prog&#39;</span><span class="p">:</span> <span class="p">{</span><span class="s1">&#39;tng_loss&#39;</span><span class="p">:</span> <span class="n">loss</span><span class="p">,</span> <span class="s1">&#39;batch_nb&#39;</span><span class="p">:</span> <span class="n">batch_nb</span><span class="p">}</span> <span class="c1"># optional</span>
<span class="p">}</span>
# return a dict
return output
</code></pre>
<span class="c1"># return a dict</span>
<span class="k">return</span> <span class="n">output</span>
</pre></div>
</td></tr></table>
<p>If you define multiple optimizers, this step will also be called with an additional <code>optimizer_idx</code> param. </p>
<pre><code class="python"># Multiple optimizers (ie: GANs)
def training_step(self, data_batch, batch_nb, optimizer_idx):
if optimizer_idx == 0:
# do training_step with encoder
if optimizer_idx == 1:
# do training_step with decoder
</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"># Multiple optimizers (ie: GANs) </span>
<span class="k">def</span> <span class="nf">training_step</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">data_batch</span><span class="p">,</span> <span class="n">batch_nb</span><span class="p">,</span> <span class="n">optimizer_idx</span><span class="p">):</span>
<span class="k">if</span> <span class="n">optimizer_idx</span> <span class="o">==</span> <span class="mi">0</span><span class="p">:</span>
<span class="c1"># do training_step with encoder</span>
<span class="k">if</span> <span class="n">optimizer_idx</span> <span class="o">==</span> <span class="mi">1</span><span class="p">:</span>
<span class="c1"># do training_step with decoder </span>
</pre></div>
</td></tr></table>
<hr />
<h3 id="tng_dataloader">tng_dataloader</h3>
<pre><code class="python">@pl.data_loader
def tng_dataloader(self)
</code></pre>
<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="nd">@pl.data_loader</span>
<span class="k">def</span> <span class="nf">tng_dataloader</span><span class="p">(</span><span class="bp">self</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<p>Called by lightning during training loop. Make sure to use the @pl.data_loader decorator, this ensures not calling this function until the data are needed.</p>
<h5 id="return">Return</h5>
<p>PyTorch DataLoader</p>
<p><strong>Example</strong></p>
<pre><code class="python">@pl.data_loader
def tng_dataloader(self):
transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,), (1.0,))])
dataset = MNIST(root='/path/to/mnist/', train=True, transform=transform, download=True)
loader = torch.utils.data.DataLoader(
dataset=dataset,
batch_size=self.hparams.batch_size,
shuffle=True
)
return loader
</code></pre>
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre> 1
2
3
4
5
6
7
8
9
10</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="nd">@pl.data_loader</span>
<span class="k">def</span> <span class="nf">tng_dataloader</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
<span class="n">transform</span> <span class="o">=</span> <span class="n">transforms</span><span class="o">.</span><span class="n">Compose</span><span class="p">([</span><span class="n">transforms</span><span class="o">.</span><span class="n">ToTensor</span><span class="p">(),</span> <span class="n">transforms</span><span class="o">.</span><span class="n">Normalize</span><span class="p">((</span><span class="mf">0.5</span><span class="p">,),</span> <span class="p">(</span><span class="mf">1.0</span><span class="p">,))])</span>
<span class="n">dataset</span> <span class="o">=</span> <span class="n">MNIST</span><span class="p">(</span><span class="n">root</span><span class="o">=</span><span class="s1">&#39;/path/to/mnist/&#39;</span><span class="p">,</span> <span class="n">train</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span> <span class="n">transform</span><span class="o">=</span><span class="n">transform</span><span class="p">,</span> <span class="n">download</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>
<span class="n">loader</span> <span class="o">=</span> <span class="n">torch</span><span class="o">.</span><span class="n">utils</span><span class="o">.</span><span class="n">data</span><span class="o">.</span><span class="n">DataLoader</span><span class="p">(</span>
<span class="n">dataset</span><span class="o">=</span><span class="n">dataset</span><span class="p">,</span>
<span class="n">batch_size</span><span class="o">=</span><span class="bp">self</span><span class="o">.</span><span class="n">hparams</span><span class="o">.</span><span class="n">batch_size</span><span class="p">,</span>
<span class="n">shuffle</span><span class="o">=</span><span class="bp">True</span>
<span class="p">)</span>
<span class="k">return</span> <span class="n">loader</span>
</pre></div>
</td></tr></table>
<hr />
<h3 id="configure_optimizers">configure_optimizers</h3>
<pre><code class="python">def configure_optimizers(self)
</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="k">def</span> <span class="nf">configure_optimizers</span><span class="p">(</span><span class="bp">self</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<p>Set up as many optimizers and (optionally) learning rate schedulers as you need. Normally you'd need one. But in the case of GANs or something more esoteric you might have multiple.
Lightning will call .backward() and .step() on each one in every epoch. If you use 16 bit precision it will also handle that.</p>
@@ -1160,28 +1248,43 @@ Lightning will call .backward() and .step() on each one in every epoch. If you
<h5 id="return_1">Return</h5>
<p>List or Tuple - List of optimizers with an optional second list of learning-rate schedulers</p>
<p><strong>Example</strong></p>
<pre><code class="python"># most cases
def configure_optimizers(self):
opt = Adam(self.parameters(), lr=0.01)
return [opt]
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre> 1
2
3
4
5
6
7
8
9
10
11</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># most cases</span>
<span class="k">def</span> <span class="nf">configure_optimizers</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
<span class="n">opt</span> <span class="o">=</span> <span class="n">Adam</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">parameters</span><span class="p">(),</span> <span class="n">lr</span><span class="o">=</span><span class="mf">0.01</span><span class="p">)</span>
<span class="k">return</span> <span class="p">[</span><span class="n">opt</span><span class="p">]</span>
# gan example, with scheduler for discriminator
def configure_optimizers(self):
generator_opt = Adam(self.model_gen.parameters(), lr=0.01)
disriminator_opt = Adam(self.model_disc.parameters(), lr=0.02)
discriminator_sched = CosineAnnealing(discriminator_opt, T_max=10)
return [generator_opt, disriminator_opt], [discriminator_sched]
</code></pre>
<span class="c1"># gan example, with scheduler for discriminator</span>
<span class="k">def</span> <span class="nf">configure_optimizers</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
<span class="n">generator_opt</span> <span class="o">=</span> <span class="n">Adam</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">model_gen</span><span class="o">.</span><span class="n">parameters</span><span class="p">(),</span> <span class="n">lr</span><span class="o">=</span><span class="mf">0.01</span><span class="p">)</span>
<span class="n">disriminator_opt</span> <span class="o">=</span> <span class="n">Adam</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">model_disc</span><span class="o">.</span><span class="n">parameters</span><span class="p">(),</span> <span class="n">lr</span><span class="o">=</span><span class="mf">0.02</span><span class="p">)</span>
<span class="n">discriminator_sched</span> <span class="o">=</span> <span class="n">CosineAnnealing</span><span class="p">(</span><span class="n">discriminator_opt</span><span class="p">,</span> <span class="n">T_max</span><span class="o">=</span><span class="mi">10</span><span class="p">)</span>
<span class="k">return</span> <span class="p">[</span><span class="n">generator_opt</span><span class="p">,</span> <span class="n">disriminator_opt</span><span class="p">],</span> <span class="p">[</span><span class="n">discriminator_sched</span><span class="p">]</span>
</pre></div>
</td></tr></table>
<p>If you need to control how often those optimizers step or override the default .step() schedule, override
the <a href="https://williamfalcon.github.io/pytorch-lightning/Trainer/hooks/#optimizer_step">optimizer_step</a> hook. </p>
<h2 id="optional-methods">Optional Methods</h2>
<h3 id="validation_step">validation_step</h3>
<pre><code class="python">def validation_step(self, data_batch, batch_nb)
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
2
3
4</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="k">def</span> <span class="nf">validation_step</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">data_batch</span><span class="p">,</span> <span class="n">batch_nb</span><span class="p">)</span>
# if have multiple val dataloaders:
def validation_step(self, data_batch, batch_nb, dataloader_idx)
</code></pre>
<span class="c1"># if have multiple val dataloaders: </span>
<span class="k">def</span> <span class="nf">validation_step</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">data_batch</span><span class="p">,</span> <span class="n">batch_nb</span><span class="p">,</span> <span class="n">dataloader_idx</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<p><strong>OPTIONAL</strong> <br />
If you don't need to validate you don't need to implement this method. </p>
@@ -1228,40 +1331,65 @@ If you don't need to validate you don't need to implement this method. </p>
</tbody>
</table>
<p><strong>Example</strong></p>
<pre><code class="python"># CASE 1: A single validation dataset
def validation_step(self, data_batch, batch_nb):
x, y, z = data_batch
<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</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># CASE 1: A single validation dataset</span>
<span class="k">def</span> <span class="nf">validation_step</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">data_batch</span><span class="p">,</span> <span class="n">batch_nb</span><span class="p">):</span>
<span class="n">x</span><span class="p">,</span> <span class="n">y</span><span class="p">,</span> <span class="n">z</span> <span class="o">=</span> <span class="n">data_batch</span>
# implement your own
out = self.forward(x)
loss = self.loss(out, x)
<span class="c1"># implement your own</span>
<span class="n">out</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">forward</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
<span class="n">loss</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">loss</span><span class="p">(</span><span class="n">out</span><span class="p">,</span> <span class="n">x</span><span class="p">)</span>
# calculate acc
labels_hat = torch.argmax(out, dim=1)
val_acc = torch.sum(y == labels_hat).item() / (len(y) * 1.0)
<span class="c1"># calculate acc</span>
<span class="n">labels_hat</span> <span class="o">=</span> <span class="n">torch</span><span class="o">.</span><span class="n">argmax</span><span class="p">(</span><span class="n">out</span><span class="p">,</span> <span class="n">dim</span><span class="o">=</span><span class="mi">1</span><span class="p">)</span>
<span class="n">val_acc</span> <span class="o">=</span> <span class="n">torch</span><span class="o">.</span><span class="n">sum</span><span class="p">(</span><span class="n">y</span> <span class="o">==</span> <span class="n">labels_hat</span><span class="p">)</span><span class="o">.</span><span class="n">item</span><span class="p">()</span> <span class="o">/</span> <span class="p">(</span><span class="nb">len</span><span class="p">(</span><span class="n">y</span><span class="p">)</span> <span class="o">*</span> <span class="mf">1.0</span><span class="p">)</span>
# all optional...
# return whatever you need for the collation function validation_end
output = OrderedDict({
'val_loss': loss_val,
'val_acc': torch.tensor(val_acc), # everything must be a tensor
})
<span class="c1"># all optional...</span>
<span class="c1"># return whatever you need for the collation function validation_end</span>
<span class="n">output</span> <span class="o">=</span> <span class="n">OrderedDict</span><span class="p">({</span>
<span class="s1">&#39;val_loss&#39;</span><span class="p">:</span> <span class="n">loss_val</span><span class="p">,</span>
<span class="s1">&#39;val_acc&#39;</span><span class="p">:</span> <span class="n">torch</span><span class="o">.</span><span class="n">tensor</span><span class="p">(</span><span class="n">val_acc</span><span class="p">),</span> <span class="c1"># everything must be a tensor</span>
<span class="p">})</span>
# return an optional dict
return output
</code></pre>
<span class="c1"># return an optional dict</span>
<span class="k">return</span> <span class="n">output</span>
</pre></div>
</td></tr></table>
<p>If you pass in multiple validation datasets, validation_step will have an additional argument.</p>
<pre><code class="python"># CASE 2: multiple validation datasets
def validation_step(self, data_batch, batch_nb, dataset_idx):
# dataset_idx tells you which dataset this is.
</code></pre>
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
2
3</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># CASE 2: multiple validation datasets</span>
<span class="k">def</span> <span class="nf">validation_step</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">data_batch</span><span class="p">,</span> <span class="n">batch_nb</span><span class="p">,</span> <span class="n">dataset_idx</span><span class="p">):</span>
<span class="c1"># dataset_idx tells you which dataset this is. </span>
</pre></div>
</td></tr></table>
<p>The <code>dataset_idx</code> corresponds to the order of datasets returned in <code>val_dataloader</code>. </p>
<hr />
<h3 id="validation_end">validation_end</h3>
<pre><code class="python">def validation_end(self, outputs)
</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="k">def</span> <span class="nf">validation_end</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">outputs</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<p>If you didn't define a validation_step, this won't be called. </p>
<p>Called at the end of the validation loop with the output of each validation_step. Called once per validation dataset. </p>
@@ -1299,28 +1427,45 @@ def validation_step(self, data_batch, batch_nb, dataset_idx):
</tbody>
</table>
<p><strong>Example</strong></p>
<pre><code class="python">def validation_end(self, outputs):
&quot;&quot;&quot;
Called at the end of validation to aggregate outputs
:param outputs: list of individual outputs of each validation step
:return:
&quot;&quot;&quot;
val_loss_mean = 0
val_acc_mean = 0
for output in outputs:
val_loss_mean += output['val_loss']
val_acc_mean += output['val_acc']
<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</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="k">def</span> <span class="nf">validation_end</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">outputs</span><span class="p">):</span>
<span class="sd">&quot;&quot;&quot;</span>
<span class="sd"> Called at the end of validation to aggregate outputs</span>
<span class="sd"> :param outputs: list of individual outputs of each validation step</span>
<span class="sd"> :return:</span>
<span class="sd"> &quot;&quot;&quot;</span>
<span class="n">val_loss_mean</span> <span class="o">=</span> <span class="mi">0</span>
<span class="n">val_acc_mean</span> <span class="o">=</span> <span class="mi">0</span>
<span class="k">for</span> <span class="n">output</span> <span class="ow">in</span> <span class="n">outputs</span><span class="p">:</span>
<span class="n">val_loss_mean</span> <span class="o">+=</span> <span class="n">output</span><span class="p">[</span><span class="s1">&#39;val_loss&#39;</span><span class="p">]</span>
<span class="n">val_acc_mean</span> <span class="o">+=</span> <span class="n">output</span><span class="p">[</span><span class="s1">&#39;val_acc&#39;</span><span class="p">]</span>
val_loss_mean /= len(outputs)
val_acc_mean /= len(outputs)
tqdm_dic = {'val_loss': val_loss_mean.item(), 'val_acc': val_acc_mean.item()}
return tqdm_dic
</code></pre>
<span class="n">val_loss_mean</span> <span class="o">/=</span> <span class="nb">len</span><span class="p">(</span><span class="n">outputs</span><span class="p">)</span>
<span class="n">val_acc_mean</span> <span class="o">/=</span> <span class="nb">len</span><span class="p">(</span><span class="n">outputs</span><span class="p">)</span>
<span class="n">tqdm_dic</span> <span class="o">=</span> <span class="p">{</span><span class="s1">&#39;val_loss&#39;</span><span class="p">:</span> <span class="n">val_loss_mean</span><span class="o">.</span><span class="n">item</span><span class="p">(),</span> <span class="s1">&#39;val_acc&#39;</span><span class="p">:</span> <span class="n">val_acc_mean</span><span class="o">.</span><span class="n">item</span><span class="p">()}</span>
<span class="k">return</span> <span class="n">tqdm_dic</span>
</pre></div>
</td></tr></table>
<hr />
<h3 id="on_save_checkpoint">on_save_checkpoint</h3>
<pre><code class="python">def on_save_checkpoint(self, checkpoint)
</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="k">def</span> <span class="nf">on_save_checkpoint</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">checkpoint</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<p>Called by lightning to checkpoint your model. Lightning saves the training state (current epoch, global_step, etc)
and also saves the model state_dict. If you want to save anything else, use this method to add your own
@@ -1328,15 +1473,19 @@ key-value pair.</p>
<h5 id="return_2">Return</h5>
<p>Nothing</p>
<p><strong>Example</strong></p>
<pre><code class="python">def on_save_checkpoint(self, checkpoint):
# 99% of use cases you don't need to implement this method
checkpoint['something_cool_i_want_to_save'] = my_cool_pickable_object
</code></pre>
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
2
3</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="k">def</span> <span class="nf">on_save_checkpoint</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">checkpoint</span><span class="p">):</span>
<span class="c1"># 99% of use cases you don&#39;t need to implement this method </span>
<span class="n">checkpoint</span><span class="p">[</span><span class="s1">&#39;something_cool_i_want_to_save&#39;</span><span class="p">]</span> <span class="o">=</span> <span class="n">my_cool_pickable_object</span>
</pre></div>
</td></tr></table>
<hr />
<h3 id="on_load_checkpoint">on_load_checkpoint</h3>
<pre><code class="python">def on_load_checkpoint(self, checkpoint)
</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="k">def</span> <span class="nf">on_load_checkpoint</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">checkpoint</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<p>Called by lightning to restore your model. Lighting auto-restores global step, epoch, etc...
It also restores the model state_dict.
@@ -1344,16 +1493,21 @@ If you saved something with <strong>on_save_checkpoint</strong> this is your cha
<h5 id="return_3">Return</h5>
<p>Nothing </p>
<p><strong>Example</strong></p>
<pre><code class="python">def on_load_checkpoint(self, checkpoint):
# 99% of the time you don't need to implement this method
self.something_cool_i_want_to_save = checkpoint['something_cool_i_want_to_save']
</code></pre>
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
2
3</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="k">def</span> <span class="nf">on_load_checkpoint</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">checkpoint</span><span class="p">):</span>
<span class="c1"># 99% of the time you don&#39;t need to implement this method</span>
<span class="bp">self</span><span class="o">.</span><span class="n">something_cool_i_want_to_save</span> <span class="o">=</span> <span class="n">checkpoint</span><span class="p">[</span><span class="s1">&#39;something_cool_i_want_to_save&#39;</span><span class="p">]</span>
</pre></div>
</td></tr></table>
<hr />
<h3 id="val_dataloader">val_dataloader</h3>
<pre><code class="python">@pl.data_loader
def tng_dataloader(self)
</code></pre>
<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="nd">@pl.data_loader</span>
<span class="k">def</span> <span class="nf">tng_dataloader</span><span class="p">(</span><span class="bp">self</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<p><strong>OPTIONAL</strong> <br />
If you don't need a validation dataset and a validation_step, you don't need to implement this method. </p>
@@ -1361,31 +1515,49 @@ If you don't need a validation dataset and a validation_step, you don't need to
<h5 id="return_4">Return</h5>
<p>PyTorch DataLoader or list of PyTorch Dataloaders. </p>
<p><strong>Example</strong></p>
<pre><code class="python">@pl.data_loader
def val_dataloader(self):
transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,), (1.0,))])
dataset = MNIST(root='/path/to/mnist/', train=False, transform=transform, download=True)
loader = torch.utils.data.DataLoader(
dataset=dataset,
batch_size=self.hparams.batch_size,
shuffle=True
)
<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</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="nd">@pl.data_loader</span>
<span class="k">def</span> <span class="nf">val_dataloader</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
<span class="n">transform</span> <span class="o">=</span> <span class="n">transforms</span><span class="o">.</span><span class="n">Compose</span><span class="p">([</span><span class="n">transforms</span><span class="o">.</span><span class="n">ToTensor</span><span class="p">(),</span> <span class="n">transforms</span><span class="o">.</span><span class="n">Normalize</span><span class="p">((</span><span class="mf">0.5</span><span class="p">,),</span> <span class="p">(</span><span class="mf">1.0</span><span class="p">,))])</span>
<span class="n">dataset</span> <span class="o">=</span> <span class="n">MNIST</span><span class="p">(</span><span class="n">root</span><span class="o">=</span><span class="s1">&#39;/path/to/mnist/&#39;</span><span class="p">,</span> <span class="n">train</span><span class="o">=</span><span class="bp">False</span><span class="p">,</span> <span class="n">transform</span><span class="o">=</span><span class="n">transform</span><span class="p">,</span> <span class="n">download</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>
<span class="n">loader</span> <span class="o">=</span> <span class="n">torch</span><span class="o">.</span><span class="n">utils</span><span class="o">.</span><span class="n">data</span><span class="o">.</span><span class="n">DataLoader</span><span class="p">(</span>
<span class="n">dataset</span><span class="o">=</span><span class="n">dataset</span><span class="p">,</span>
<span class="n">batch_size</span><span class="o">=</span><span class="bp">self</span><span class="o">.</span><span class="n">hparams</span><span class="o">.</span><span class="n">batch_size</span><span class="p">,</span>
<span class="n">shuffle</span><span class="o">=</span><span class="bp">True</span>
<span class="p">)</span>
return loader
<span class="k">return</span> <span class="n">loader</span>
# can also return multiple dataloaders
@pl.data_loader
def val_dataloader(self):
return [loader_a, loader_b, ..., loader_n]
</code></pre>
<span class="c1"># can also return multiple dataloaders </span>
<span class="nd">@pl.data_loader</span>
<span class="k">def</span> <span class="nf">val_dataloader</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
<span class="k">return</span> <span class="p">[</span><span class="n">loader_a</span><span class="p">,</span> <span class="n">loader_b</span><span class="p">,</span> <span class="o">...</span><span class="p">,</span> <span class="n">loader_n</span><span class="p">]</span>
</pre></div>
</td></tr></table>
<p>In the case where you return multiple val_dataloaders, the validation_step will have an arguement <code>dataset_idx</code>
which matches the order here. </p>
<hr />
<h3 id="test_dataloader">test_dataloader</h3>
<pre><code class="python">@pl.data_loader
def test_dataloader(self)
</code></pre>
<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="nd">@pl.data_loader</span>
<span class="k">def</span> <span class="nf">test_dataloader</span><span class="p">(</span><span class="bp">self</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<p><strong>OPTIONAL</strong> <br />
If you don't need a test dataset and a test_step, you don't need to implement this method. </p>
@@ -1393,39 +1565,56 @@ If you don't need a test dataset and a test_step, you don't need to implement th
<h5 id="return_5">Return</h5>
<p>PyTorch DataLoader</p>
<p><strong>Example</strong></p>
<pre><code class="python">@pl.data_loader
def test_dataloader(self):
transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,), (1.0,))])
dataset = MNIST(root='/path/to/mnist/', train=False, transform=transform, download=True)
loader = torch.utils.data.DataLoader(
dataset=dataset,
batch_size=self.hparams.batch_size,
shuffle=True
)
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre> 1
2
3
4
5
6
7
8
9
10
11</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="nd">@pl.data_loader</span>
<span class="k">def</span> <span class="nf">test_dataloader</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
<span class="n">transform</span> <span class="o">=</span> <span class="n">transforms</span><span class="o">.</span><span class="n">Compose</span><span class="p">([</span><span class="n">transforms</span><span class="o">.</span><span class="n">ToTensor</span><span class="p">(),</span> <span class="n">transforms</span><span class="o">.</span><span class="n">Normalize</span><span class="p">((</span><span class="mf">0.5</span><span class="p">,),</span> <span class="p">(</span><span class="mf">1.0</span><span class="p">,))])</span>
<span class="n">dataset</span> <span class="o">=</span> <span class="n">MNIST</span><span class="p">(</span><span class="n">root</span><span class="o">=</span><span class="s1">&#39;/path/to/mnist/&#39;</span><span class="p">,</span> <span class="n">train</span><span class="o">=</span><span class="bp">False</span><span class="p">,</span> <span class="n">transform</span><span class="o">=</span><span class="n">transform</span><span class="p">,</span> <span class="n">download</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>
<span class="n">loader</span> <span class="o">=</span> <span class="n">torch</span><span class="o">.</span><span class="n">utils</span><span class="o">.</span><span class="n">data</span><span class="o">.</span><span class="n">DataLoader</span><span class="p">(</span>
<span class="n">dataset</span><span class="o">=</span><span class="n">dataset</span><span class="p">,</span>
<span class="n">batch_size</span><span class="o">=</span><span class="bp">self</span><span class="o">.</span><span class="n">hparams</span><span class="o">.</span><span class="n">batch_size</span><span class="p">,</span>
<span class="n">shuffle</span><span class="o">=</span><span class="bp">True</span>
<span class="p">)</span>
return loader
</code></pre>
<span class="k">return</span> <span class="n">loader</span>
</pre></div>
</td></tr></table>
<hr />
<h3 id="update_tng_log_metrics">update_tng_log_metrics</h3>
<pre><code class="python">def update_tng_log_metrics(self, logs)
</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="k">def</span> <span class="nf">update_tng_log_metrics</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">logs</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<p>Called by lightning right before it logs metrics for this batch.
This is a chance to ammend or add to the metrics about to be logged.</p>
<h5 id="return_6">Return</h5>
<p>Dict </p>
<p><strong>Example</strong></p>
<pre><code class="python">def update_tng_log_metrics(self, logs):
# modify or add to logs
return logs
</code></pre>
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
2
3</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="k">def</span> <span class="nf">update_tng_log_metrics</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">logs</span><span class="p">):</span>
<span class="c1"># modify or add to logs</span>
<span class="k">return</span> <span class="n">logs</span>
</pre></div>
</td></tr></table>
<hr />
<h3 id="add_model_specific_args">add_model_specific_args</h3>
<pre><code class="python">@staticmethod
def add_model_specific_args(parent_parser, root_dir)
</code></pre>
<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="nd">@staticmethod</span>
<span class="k">def</span> <span class="nf">add_model_specific_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>
</pre></div>
</td></tr></table>
<p>Lightning has a list of default argparse commands.
This method is your chance to add or modify commands specific to your model.
@@ -1433,29 +1622,51 @@ The <a href="https://williamfalcon.github.io/test-tube/hyperparameter_optimizati
<h5 id="return_7">Return</h5>
<p>An argument parser</p>
<p><strong>Example</strong></p>
<pre><code class="python">@staticmethod
def add_model_specific_args(parent_parser, root_dir):
parser = HyperOptArgumentParser(strategy=parent_parser.strategy, parents=[parent_parser])
<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</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="nd">@staticmethod</span>
<span class="k">def</span> <span class="nf">add_model_specific_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>
<span class="n">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="n">parent_parser</span><span class="o">.</span><span class="n">strategy</span><span class="p">,</span> <span class="n">parents</span><span class="o">=</span><span class="p">[</span><span class="n">parent_parser</span><span class="p">])</span>
# param overwrites
# parser.set_defaults(gradient_clip=5.0)
<span class="c1"># param overwrites</span>
<span class="c1"># parser.set_defaults(gradient_clip=5.0)</span>
# 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('--out_features', default=10)
parser.add_argument('--hidden_dim', default=50000) # use 500 for CPU, 50000 for GPU to see speed difference
<span class="c1"># network params</span>
<span class="n">parser</span><span class="o">.</span><span class="n">opt_list</span><span class="p">(</span><span class="s1">&#39;--drop_prob&#39;</span><span class="p">,</span> <span class="n">default</span><span class="o">=</span><span class="mf">0.2</span><span class="p">,</span> <span class="n">options</span><span class="o">=</span><span class="p">[</span><span class="mf">0.2</span><span class="p">,</span> <span class="mf">0.5</span><span class="p">],</span> <span class="nb">type</span><span class="o">=</span><span class="nb">float</span><span class="p">,</span> <span class="n">tunable</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>
<span class="n">parser</span><span class="o">.</span><span class="n">add_argument</span><span class="p">(</span><span class="s1">&#39;--in_features&#39;</span><span class="p">,</span> <span class="n">default</span><span class="o">=</span><span class="mi">28</span><span class="o">*</span><span class="mi">28</span><span class="p">)</span>
<span class="n">parser</span><span class="o">.</span><span class="n">add_argument</span><span class="p">(</span><span class="s1">&#39;--out_features&#39;</span><span class="p">,</span> <span class="n">default</span><span class="o">=</span><span class="mi">10</span><span class="p">)</span>
<span class="n">parser</span><span class="o">.</span><span class="n">add_argument</span><span class="p">(</span><span class="s1">&#39;--hidden_dim&#39;</span><span class="p">,</span> <span class="n">default</span><span class="o">=</span><span class="mi">50000</span><span class="p">)</span> <span class="c1"># use 500 for CPU, 50000 for GPU to see speed difference</span>
# data
parser.add_argument('--data_root', default=os.path.join(root_dir, 'mnist'), type=str)
<span class="c1"># data</span>
<span class="n">parser</span><span class="o">.</span><span class="n">add_argument</span><span class="p">(</span><span class="s1">&#39;--data_root&#39;</span><span class="p">,</span> <span class="n">default</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">join</span><span class="p">(</span><span class="n">root_dir</span><span class="p">,</span> <span class="s1">&#39;mnist&#39;</span><span class="p">),</span> <span class="nb">type</span><span class="o">=</span><span class="nb">str</span><span class="p">)</span>
# training params (opt)
parser.opt_list('--learning_rate', default=0.001, type=float, options=[0.0001, 0.0005, 0.001, 0.005],
tunable=False)
parser.opt_list('--batch_size', default=256, type=int, options=[32, 64, 128, 256], tunable=False)
parser.opt_list('--optimizer_name', default='adam', type=str, options=['adam'], tunable=False)
return parser
</code></pre>
<span class="c1"># training params (opt)</span>
<span class="n">parser</span><span class="o">.</span><span class="n">opt_list</span><span class="p">(</span><span class="s1">&#39;--learning_rate&#39;</span><span class="p">,</span> <span class="n">default</span><span class="o">=</span><span class="mf">0.001</span><span class="p">,</span> <span class="nb">type</span><span class="o">=</span><span class="nb">float</span><span class="p">,</span> <span class="n">options</span><span class="o">=</span><span class="p">[</span><span class="mf">0.0001</span><span class="p">,</span> <span class="mf">0.0005</span><span class="p">,</span> <span class="mf">0.001</span><span class="p">,</span> <span class="mf">0.005</span><span class="p">],</span>
<span class="n">tunable</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>
<span class="n">parser</span><span class="o">.</span><span class="n">opt_list</span><span class="p">(</span><span class="s1">&#39;--batch_size&#39;</span><span class="p">,</span> <span class="n">default</span><span class="o">=</span><span class="mi">256</span><span class="p">,</span> <span class="nb">type</span><span class="o">=</span><span class="nb">int</span><span class="p">,</span> <span class="n">options</span><span class="o">=</span><span class="p">[</span><span class="mi">32</span><span class="p">,</span> <span class="mi">64</span><span class="p">,</span> <span class="mi">128</span><span class="p">,</span> <span class="mi">256</span><span class="p">],</span> <span class="n">tunable</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>
<span class="n">parser</span><span class="o">.</span><span class="n">opt_list</span><span class="p">(</span><span class="s1">&#39;--optimizer_name&#39;</span><span class="p">,</span> <span class="n">default</span><span class="o">=</span><span class="s1">&#39;adam&#39;</span><span class="p">,</span> <span class="nb">type</span><span class="o">=</span><span class="nb">str</span><span class="p">,</span> <span class="n">options</span><span class="o">=</span><span class="p">[</span><span class="s1">&#39;adam&#39;</span><span class="p">],</span> <span class="n">tunable</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>
<span class="k">return</span> <span class="n">parser</span>
</pre></div>
</td></tr></table>
+32 -17
View File
@@ -564,26 +564,39 @@
<hr />
<h3 id="freeze">freeze</h3>
<p>Freeze all params for inference</p>
<pre><code class="python">model = MyLightningModule(...)
model.freeze()
</code></pre>
<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="n">model</span> <span class="o">=</span> <span class="n">MyLightningModule</span><span class="p">(</span><span class="o">...</span><span class="p">)</span>
<span class="n">model</span><span class="o">.</span><span class="n">freeze</span><span class="p">()</span>
</pre></div>
</td></tr></table>
<hr />
<h3 id="load_from_metrics">load_from_metrics</h3>
<p>This is the easiest/fastest way which uses the meta_tags.csv file from test-tube to rebuild the model.
The meta_tags.csv file can be found in the test-tube experiment save_dir. </p>
<pre><code class="python">pretrained_model = MyLightningModule.load_from_metrics(
weights_path='/path/to/pytorch_checkpoint.ckpt',
tags_csv='/path/to/test_tube/experiment/version/meta_tags.csv',
on_gpu=True,
map_location=None
)
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre> 1
2
3
4
5
6
7
8
9
10
11</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="n">pretrained_model</span> <span class="o">=</span> <span class="n">MyLightningModule</span><span class="o">.</span><span class="n">load_from_metrics</span><span class="p">(</span>
<span class="n">weights_path</span><span class="o">=</span><span class="s1">&#39;/path/to/pytorch_checkpoint.ckpt&#39;</span><span class="p">,</span>
<span class="n">tags_csv</span><span class="o">=</span><span class="s1">&#39;/path/to/test_tube/experiment/version/meta_tags.csv&#39;</span><span class="p">,</span>
<span class="n">on_gpu</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span>
<span class="n">map_location</span><span class="o">=</span><span class="bp">None</span>
<span class="p">)</span>
# predict
pretrained_model.eval()
pretrained_model.freeze()
y_hat = pretrained_model(x)
</code></pre>
<span class="c1"># predict</span>
<span class="n">pretrained_model</span><span class="o">.</span><span class="n">eval</span><span class="p">()</span>
<span class="n">pretrained_model</span><span class="o">.</span><span class="n">freeze</span><span class="p">()</span>
<span class="n">y_hat</span> <span class="o">=</span> <span class="n">pretrained_model</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<p><strong>Params</strong> </p>
<table>
@@ -617,9 +630,11 @@ y_hat = pretrained_model(x)
<hr />
<h3 id="unfreeze">unfreeze</h3>
<p>Unfreeze all params for inference</p>
<pre><code class="python">model = MyLightningModule(...)
model.unfreeze()
</code></pre>
<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="n">model</span> <span class="o">=</span> <span class="n">MyLightningModule</span><span class="p">(</span><span class="o">...</span><span class="p">)</span>
<span class="n">model</span><span class="o">.</span><span class="n">unfreeze</span><span class="p">()</span>
</pre></div>
</td></tr></table>
+21 -12
View File
@@ -666,10 +666,13 @@
<hr />
<h4 id="experiment">experiment</h4>
<p>An instance of test-tube Experiment which you can use to log anything for tensorboarX. </p>
<pre><code class="python">self.experiment.add_embedding(...)
self.experiment.log({'val_loss': 0.9})
self.experiment.add_scalars(...)
</code></pre>
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
2
3</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="bp">self</span><span class="o">.</span><span class="n">experiment</span><span class="o">.</span><span class="n">add_embedding</span><span class="p">(</span><span class="o">...</span><span class="p">)</span>
<span class="bp">self</span><span class="o">.</span><span class="n">experiment</span><span class="o">.</span><span class="n">log</span><span class="p">({</span><span class="s1">&#39;val_loss&#39;</span><span class="p">:</span> <span class="mf">0.9</span><span class="p">})</span>
<span class="bp">self</span><span class="o">.</span><span class="n">experiment</span><span class="o">.</span><span class="n">add_scalars</span><span class="p">(</span><span class="o">...</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="global_step">global_step</h4>
@@ -683,10 +686,13 @@ self.experiment.add_scalars(...)
<hr />
<h4 id="trainer">trainer</h4>
<p>Last resort access to any state the trainer has. Changing certain properties here could affect your training run.</p>
<pre><code class="python">self.trainer.optimizers
self.trainer.current_epoch
...
</code></pre>
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
2
3</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="bp">self</span><span class="o">.</span><span class="n">trainer</span><span class="o">.</span><span class="n">optimizers</span>
<span class="bp">self</span><span class="o">.</span><span class="n">trainer</span><span class="o">.</span><span class="n">current_epoch</span>
<span class="o">...</span>
</pre></div>
</td></tr></table>
<h2 id="debugging">Debugging</h2>
<p>The LightningModule also offers these tricks to help debug. </p>
@@ -694,10 +700,13 @@ self.trainer.current_epoch
<h4 id="example_input_array">example_input_array</h4>
<p>In the LightningModule init, you can set a dummy tensor for this property
to get a print out of sizes coming into and out of every layer. </p>
<pre><code class="python">def __init__(self):
# put the dimensions of the first input to your system
self.example_input_array = torch.rand(5, 28 * 28)
</code></pre>
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
2
3</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="k">def</span> <span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
<span class="c1"># put the dimensions of the first input to your system</span>
<span class="bp">self</span><span class="o">.</span><span class="n">example_input_array</span> <span class="o">=</span> <span class="n">torch</span><span class="o">.</span><span class="n">rand</span><span class="p">(</span><span class="mi">5</span><span class="p">,</span> <span class="mi">28</span> <span class="o">*</span> <span class="mi">28</span><span class="p">)</span>
</pre></div>
</td></tr></table>
+66 -32
View File
@@ -550,18 +550,29 @@
<hr />
<h3 id="model-saving">Model saving</h3>
<p>To enable checkpointing, define the checkpoint callback and give it to the trainer.</p>
<pre><code class="python">from pytorch_lightning.callbacks import ModelCheckpoint
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre> 1
2
3
4
5
6
7
8
9
10
11</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="kn">from</span> <span class="nn">pytorch_lightning.callbacks</span> <span class="kn">import</span> <span class="n">ModelCheckpoint</span>
checkpoint_callback = ModelCheckpoint(
filepath='/path/to/store/weights.ckpt',
save_best_only=True,
verbose=True,
monitor='val_loss',
mode='min'
)
<span class="n">checkpoint_callback</span> <span class="o">=</span> <span class="n">ModelCheckpoint</span><span class="p">(</span>
<span class="n">filepath</span><span class="o">=</span><span class="s1">&#39;/path/to/store/weights.ckpt&#39;</span><span class="p">,</span>
<span class="n">save_best_only</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span>
<span class="n">verbose</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span>
<span class="n">monitor</span><span class="o">=</span><span class="s1">&#39;val_loss&#39;</span><span class="p">,</span>
<span class="n">mode</span><span class="o">=</span><span class="s1">&#39;min&#39;</span>
<span class="p">)</span>
trainer = Trainer(checkpoint_callback=checkpoint_callback)
</code></pre>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</span><span class="p">(</span><span class="n">checkpoint_callback</span><span class="o">=</span><span class="n">checkpoint_callback</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<hr />
<h3 id="restoring-training-session">Restoring training session</h3>
@@ -569,16 +580,25 @@ trainer = Trainer(checkpoint_callback=checkpoint_callback)
restore the trainer state as well. This will continue from the epoch and global step you last left off.<br />
However, the dataloaders will start from the first batch again (if you shuffled it shouldn't matter). </p>
<p>Lightning will restore the session if you pass an experiment with the same version and there's a saved checkpoint. </p>
<pre><code class="python">from test_tube import Experiment
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
2
3
4
5
6
7
8
9</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">Experiment</span>
exp = Experiment(version=a_previous_version_with_a_saved_checkpoint)
trainer = Trainer(experiment=exp)
<span class="n">exp</span> <span class="o">=</span> <span class="n">Experiment</span><span class="p">(</span><span class="n">version</span><span class="o">=</span><span class="n">a_previous_version_with_a_saved_checkpoint</span><span class="p">)</span>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</span><span class="p">(</span><span class="n">experiment</span><span class="o">=</span><span class="n">exp</span><span class="p">)</span>
# this fit call loads model weights and trainer state
# the trainer continues seamlessly from where you left off
# without having to do anything else.
trainer.fit(model)
</code></pre>
<span class="c1"># this fit call loads model weights and trainer state</span>
<span class="c1"># the trainer continues seamlessly from where you left off</span>
<span class="c1"># without having to do anything else.</span>
<span class="n">trainer</span><span class="o">.</span><span class="n">fit</span><span class="p">(</span><span class="n">model</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<p>The trainer restores:<br />
- global_step <br />
@@ -589,23 +609,37 @@ trainer.fit(model)
<p>You can even change the logic of your model as long as the weights and "architecture" of
the system isn't different. If you add a layer, for instance, it might not work. </p>
<p>At a rough level, here's <a href="https://github.com/williamFalcon/pytorch-lightning/blob/master/pytorch_lightning/root_module/model_saving.py#L63">what happens inside Trainer</a>: </p>
<pre><code class="python">
self.global_step = checkpoint['global_step']
self.current_epoch = checkpoint['epoch']
<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="bp">self</span><span class="o">.</span><span class="n">global_step</span> <span class="o">=</span> <span class="n">checkpoint</span><span class="p">[</span><span class="s1">&#39;global_step&#39;</span><span class="p">]</span>
<span class="bp">self</span><span class="o">.</span><span class="n">current_epoch</span> <span class="o">=</span> <span class="n">checkpoint</span><span class="p">[</span><span class="s1">&#39;epoch&#39;</span><span class="p">]</span>
# restore the optimizers
optimizer_states = checkpoint['optimizer_states']
for optimizer, opt_state in zip(self.optimizers, optimizer_states):
optimizer.load_state_dict(opt_state)
<span class="c1"># restore the optimizers</span>
<span class="n">optimizer_states</span> <span class="o">=</span> <span class="n">checkpoint</span><span class="p">[</span><span class="s1">&#39;optimizer_states&#39;</span><span class="p">]</span>
<span class="k">for</span> <span class="n">optimizer</span><span class="p">,</span> <span class="n">opt_state</span> <span class="ow">in</span> <span class="nb">zip</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">optimizers</span><span class="p">,</span> <span class="n">optimizer_states</span><span class="p">):</span>
<span class="n">optimizer</span><span class="o">.</span><span class="n">load_state_dict</span><span class="p">(</span><span class="n">opt_state</span><span class="p">)</span>
# restore the lr schedulers
lr_schedulers = checkpoint['lr_schedulers']
for scheduler, lrs_state in zip(self.lr_schedulers, lr_schedulers):
scheduler.load_state_dict(lrs_state)
<span class="c1"># restore the lr schedulers</span>
<span class="n">lr_schedulers</span> <span class="o">=</span> <span class="n">checkpoint</span><span class="p">[</span><span class="s1">&#39;lr_schedulers&#39;</span><span class="p">]</span>
<span class="k">for</span> <span class="n">scheduler</span><span class="p">,</span> <span class="n">lrs_state</span> <span class="ow">in</span> <span class="nb">zip</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">lr_schedulers</span><span class="p">,</span> <span class="n">lr_schedulers</span><span class="p">):</span>
<span class="n">scheduler</span><span class="o">.</span><span class="n">load_state_dict</span><span class="p">(</span><span class="n">lrs_state</span><span class="p">)</span>
# uses the model you passed into trainer
model.load_state_dict(checkpoint['state_dict'])
</code></pre>
<span class="c1"># uses the model you passed into trainer </span>
<span class="n">model</span><span class="o">.</span><span class="n">load_state_dict</span><span class="p">(</span><span class="n">checkpoint</span><span class="p">[</span><span class="s1">&#39;state_dict&#39;</span><span class="p">])</span>
</pre></div>
</td></tr></table>
+108 -54
View File
@@ -638,12 +638,17 @@ None of the flags below require changing anything about your lightningModel defi
<p>Lightning supports two backends. DataParallel and DistributedDataParallel. Both can be used for single-node multi-GPU training.
For multi-node training you must use DistributedDataParallel. </p>
<p>You can toggle between each mode by setting this flag.</p>
<pre><code class="python"># DEFAULT uses DataParallel
trainer = Trainer(distributed_backend='dp')
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
2
3
4
5</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># DEFAULT uses DataParallel</span>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</span><span class="p">(</span><span class="n">distributed_backend</span><span class="o">=</span><span class="s1">&#39;dp&#39;</span><span class="p">)</span>
# change to distributed data parallel
trainer = Trainer(distributed_backend='ddp')
</code></pre>
<span class="c1"># change to distributed data parallel</span>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</span><span class="p">(</span><span class="n">distributed_backend</span><span class="o">=</span><span class="s1">&#39;ddp&#39;</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<p>If you request multiple nodes, the back-end will auto-switch to ddp.
We recommend you use DistributedDataparallel even for single-node multi-GPU training. It is MUCH faster than DP but <em>may</em>
@@ -712,88 +717,137 @@ not allow 16-bit and DP training. We tried to get this to work, but it's an issu
<h4 id="cuda-flags">CUDA flags</h4>
<p>CUDA flags make certain GPUs visible to your script.
Lightning sets these for you automatically, there's NO NEED to do this yourself.</p>
<pre><code class="python"># lightning will set according to what you give the trainer
# os.environ[&quot;CUDA_DEVICE_ORDER&quot;] = &quot;PCI_BUS_ID&quot;
# os.environ[&quot;CUDA_VISIBLE_DEVICES&quot;] = &quot;0&quot;
</code></pre>
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
2
3</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># lightning will set according to what you give the trainer</span>
<span class="c1"># os.environ[&quot;CUDA_DEVICE_ORDER&quot;] = &quot;PCI_BUS_ID&quot;</span>
<span class="c1"># os.environ[&quot;CUDA_VISIBLE_DEVICES&quot;] = &quot;0&quot;</span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="16-bit-mixed-precision">16-bit mixed precision</h4>
<p>16 bit precision can cut your memory footprint by half. If using volta architecture GPUs it can give a dramatic training speed-up as well. <br />
First, install apex (if install fails, look <a href="https://github.com/NVIDIA/apex">here</a>):</p>
<pre><code class="bash">$ git clone https://github.com/NVIDIA/apex
$ cd apex
$ pip install -v --no-cache-dir --global-option=&quot;--cpp_ext&quot; --global-option=&quot;--cuda_ext&quot; ./
</code></pre>
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
2
3</pre></div></td><td class="code"><div class="codehilite"><pre><span></span>$ git clone https://github.com/NVIDIA/apex
$ <span class="nb">cd</span> apex
$ pip install -v --no-cache-dir --global-option<span class="o">=</span><span class="s2">&quot;--cpp_ext&quot;</span> --global-option<span class="o">=</span><span class="s2">&quot;--cuda_ext&quot;</span> ./
</pre></div>
</td></tr></table>
<p>then set this use_amp to True.</p>
<pre><code class="python"># DEFAULT
trainer = Trainer(amp_level='O2', use_amp=False)
</code></pre>
<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"># DEFAULT</span>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</span><span class="p">(</span><span class="n">amp_level</span><span class="o">=</span><span class="s1">&#39;O2&#39;</span><span class="p">,</span> <span class="n">use_amp</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="single-gpu">Single-gpu</h4>
<p>Make sure you're on a GPU machine. </p>
<pre><code class="python"># DEFAULT
trainer = Trainer(gpus=[0])
</code></pre>
<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"># DEFAULT</span>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</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>
</pre></div>
</td></tr></table>
<hr />
<h4 id="multi-gpu">multi-gpu</h4>
<p>Make sure you're on a GPU machine. You can set as many GPUs as you want.
In this setting, the model will run on all 8 GPUs at once using DataParallel under the hood.</p>
<pre><code class="python"># to use DataParallel (default)
trainer = Trainer(gpus=[0,1,2,3,4,5,6,7], distributed_backend='dp')
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
2
3
4
5</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># to use DataParallel (default)</span>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</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="mi">4</span><span class="p">,</span><span class="mi">5</span><span class="p">,</span><span class="mi">6</span><span class="p">,</span><span class="mi">7</span><span class="p">],</span> <span class="n">distributed_backend</span><span class="o">=</span><span class="s1">&#39;dp&#39;</span><span class="p">)</span>
# RECOMMENDED use DistributedDataParallel
trainer = Trainer(gpus=[0,1,2,3,4,5,6,7], distributed_backend='ddp')
</code></pre>
<span class="c1"># RECOMMENDED use DistributedDataParallel</span>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</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="mi">4</span><span class="p">,</span><span class="mi">5</span><span class="p">,</span><span class="mi">6</span><span class="p">,</span><span class="mi">7</span><span class="p">],</span> <span class="n">distributed_backend</span><span class="o">=</span><span class="s1">&#39;ddp&#39;</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="multi-node">Multi-node</h4>
<p>Multi-node training is easily done by specifying these flags.</p>
<pre><code class="python"># train on 12*8 GPUs
trainer = Trainer(gpus=[0,1,2,3,4,5,6,7], nb_gpu_nodes=12)
</code></pre>
<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"># train on 12*8 GPUs</span>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</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="mi">4</span><span class="p">,</span><span class="mi">5</span><span class="p">,</span><span class="mi">6</span><span class="p">,</span><span class="mi">7</span><span class="p">],</span> <span class="n">nb_gpu_nodes</span><span class="o">=</span><span class="mi">12</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<p>In addition, make sure to set up your SLURM job correctly via the <a href="https://williamfalcon.github.io/test-tube/hpc/SlurmCluster/">SlurmClusterObject</a>. In particular, specify the number of tasks per node correctly.</p>
<pre><code class="python">cluster = SlurmCluster(
hyperparam_optimizer=test_tube.HyperOptArgumentParser(),
log_path='/some/path/to/save',
)
<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</pre></div></td><td class="code"><div class="codehilite"><pre><span></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">test_tube</span><span class="o">.</span><span class="n">HyperOptArgumentParser</span><span class="p">(),</span>
<span class="n">log_path</span><span class="o">=</span><span class="s1">&#39;/some/path/to/save&#39;</span><span class="p">,</span>
<span class="p">)</span>
# OPTIONAL FLAGS WHICH MAY BE CLUSTER DEPENDENT
# which interface your nodes use for communication
cluster.add_command('export NCCL_SOCKET_IFNAME=^docker0,lo')
<span class="c1"># OPTIONAL FLAGS WHICH MAY BE CLUSTER DEPENDENT</span>
<span class="c1"># which interface your nodes use for communication</span>
<span class="n">cluster</span><span class="o">.</span><span class="n">add_command</span><span class="p">(</span><span class="s1">&#39;export NCCL_SOCKET_IFNAME=^docker0,lo&#39;</span><span class="p">)</span>
# see output of the NCCL connection process
# NCCL is how the nodes talk to each other
cluster.add_command('export NCCL_DEBUG=INFO')
<span class="c1"># see output of the NCCL connection process</span>
<span class="c1"># NCCL is how the nodes talk to each other</span>
<span class="n">cluster</span><span class="o">.</span><span class="n">add_command</span><span class="p">(</span><span class="s1">&#39;export NCCL_DEBUG=INFO&#39;</span><span class="p">)</span>
# setting a master port here is a good idea.
cluster.add_command('export MASTER_PORT=%r' % PORT)
<span class="c1"># setting a master port here is a good idea.</span>
<span class="n">cluster</span><span class="o">.</span><span class="n">add_command</span><span class="p">(</span><span class="s1">&#39;export MASTER_PORT=</span><span class="si">%r</span><span class="s1">&#39;</span> <span class="o">%</span> <span class="n">PORT</span><span class="p">)</span>
# good to load the latest NCCL version
cluster.load_modules(['NCCL/2.4.7-1-cuda.10.0'])
<span class="c1"># good to load the latest NCCL version</span>
<span class="n">cluster</span><span class="o">.</span><span class="n">load_modules</span><span class="p">([</span><span class="s1">&#39;NCCL/2.4.7-1-cuda.10.0&#39;</span><span class="p">])</span>
# configure cluster
cluster.per_experiment_nb_nodes = 12
cluster.per_experiment_nb_gpus = 8
<span class="c1"># configure cluster</span>
<span class="n">cluster</span><span class="o">.</span><span class="n">per_experiment_nb_nodes</span> <span class="o">=</span> <span class="mi">12</span>
<span class="n">cluster</span><span class="o">.</span><span class="n">per_experiment_nb_gpus</span> <span class="o">=</span> <span class="mi">8</span>
cluster.add_slurm_cmd(cmd='ntasks-per-node', value=8, comment='1 task per gpu')
</code></pre>
<span class="n">cluster</span><span class="o">.</span><span class="n">add_slurm_cmd</span><span class="p">(</span><span class="n">cmd</span><span class="o">=</span><span class="s1">&#39;ntasks-per-node&#39;</span><span class="p">,</span> <span class="n">value</span><span class="o">=</span><span class="mi">8</span><span class="p">,</span> <span class="n">comment</span><span class="o">=</span><span class="s1">&#39;1 task per gpu&#39;</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<p>Finally, make sure to add a distributed sampler to your dataset. The distributed sampler copies a
portion of your dataset onto each GPU. (World_size = gpus_per_node * nb_nodes). </p>
<pre><code class="python"># ie: this:
dataset = myDataset()
dataloader = Dataloader(dataset)
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
2
3
4
5
6
7
8</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># ie: this:</span>
<span class="n">dataset</span> <span class="o">=</span> <span class="n">myDataset</span><span class="p">()</span>
<span class="n">dataloader</span> <span class="o">=</span> <span class="n">Dataloader</span><span class="p">(</span><span class="n">dataset</span><span class="p">)</span>
# becomes:
dataset = myDataset()
dist_sampler = torch.utils.data.distributed.DistributedSampler(dataset)
dataloader = Dataloader(dataset, sampler=dist_sampler)
</code></pre>
<span class="c1"># becomes:</span>
<span class="n">dataset</span> <span class="o">=</span> <span class="n">myDataset</span><span class="p">()</span>
<span class="n">dist_sampler</span> <span class="o">=</span> <span class="n">torch</span><span class="o">.</span><span class="n">utils</span><span class="o">.</span><span class="n">data</span><span class="o">.</span><span class="n">distributed</span><span class="o">.</span><span class="n">DistributedSampler</span><span class="p">(</span><span class="n">dataset</span><span class="p">)</span>
<span class="n">dataloader</span> <span class="o">=</span> <span class="n">Dataloader</span><span class="p">(</span><span class="n">dataset</span><span class="p">,</span> <span class="n">sampler</span><span class="o">=</span><span class="n">dist_sampler</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="self-balancing-architecture">Self-balancing architecture</h4>
+87 -43
View File
@@ -631,89 +631,133 @@
<p>Lighting offers a few options for logging information about model, gpu usage, etc (via test-tube). It also offers printing options for training monitoring.</p>
<hr />
<h4 id="display-metrics-in-progress-bar">Display metrics in progress bar</h4>
<pre><code class="python"># DEFAULT
trainer = Trainer(progress_bar=True)
</code></pre>
<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"># DEFAULT</span>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</span><span class="p">(</span><span class="n">progress_bar</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="log-metric-row-every-k-batches">Log metric row every k batches</h4>
<p>Every k batches lightning will make an entry in the metrics log</p>
<pre><code class="python"># DEFAULT (ie: save a .csv log file every 10 batches)
trainer = Trainer(add_log_row_interval=10)
</code></pre>
<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"># DEFAULT (ie: save a .csv log file every 10 batches)</span>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</span><span class="p">(</span><span class="n">add_log_row_interval</span><span class="o">=</span><span class="mi">10</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="process-position">Process position</h4>
<p>When running multiple models on the same machine we want to decide which progress bar to use.
Lightning will stack progress bars according to this value. </p>
<pre><code class="python"># DEFAULT
trainer = Trainer(process_position=0)
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
2
3
4
5</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># DEFAULT</span>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</span><span class="p">(</span><span class="n">process_position</span><span class="o">=</span><span class="mi">0</span><span class="p">)</span>
# if this is the second model on the node, show the second progress bar below
trainer = Trainer(process_position=1)
</code></pre>
<span class="c1"># if this is the second model on the node, show the second progress bar below</span>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</span><span class="p">(</span><span class="n">process_position</span><span class="o">=</span><span class="mi">1</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="save-a-snapshot-of-all-hyperparameters">Save a snapshot of all hyperparameters</h4>
<p>Whenever you call .save() on the test-tube experiment it logs all the hyperparameters in current use.
Give lightning a test-tube Experiment object to automate this for you.</p>
<pre><code class="python">from test_tube import Experiment
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
2
3
4</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">Experiment</span>
exp = Experiment(...)
Trainer(experiment=exp)
</code></pre>
<span class="n">exp</span> <span class="o">=</span> <span class="n">Experiment</span><span class="p">(</span><span class="o">...</span><span class="p">)</span>
<span class="n">Trainer</span><span class="p">(</span><span class="n">experiment</span><span class="o">=</span><span class="n">exp</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="snapshot-code-for-a-training-run">Snapshot code for a training run</h4>
<p>Whenever you call .save() on the test-tube experiment it snapshows all code and pushes to a git tag.
Give lightning a test-tube Experiment object to automate this for you.</p>
<pre><code class="python">from test_tube import Experiment
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
2
3
4</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">Experiment</span>
exp = Experiment(create_git_tag=True)
Trainer(experiment=exp)
</code></pre>
<span class="n">exp</span> <span class="o">=</span> <span class="n">Experiment</span><span class="p">(</span><span class="n">create_git_tag</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>
<span class="n">Trainer</span><span class="p">(</span><span class="n">experiment</span><span class="o">=</span><span class="n">exp</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<hr />
<h3 id="tensorboard-support">Tensorboard support</h3>
<p>In the LightningModule you can access the experiment logger by doing:</p>
<pre><code class="python">self.experiment
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
2
3
4
5</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="bp">self</span><span class="o">.</span><span class="n">experiment</span>
# add image
# Look at PyTorch SummaryWriter docs for what you can do.
self.experiment.add_image(...)
</code></pre>
<span class="c1"># add image</span>
<span class="c1"># Look at PyTorch SummaryWriter docs for what you can do. </span>
<span class="bp">self</span><span class="o">.</span><span class="n">experiment</span><span class="o">.</span><span class="n">add_image</span><span class="p">(</span><span class="o">...</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<p>The experiment object is a strict subclass of PyTorch SummaryWriter. However, this class
also snapshots every detail about the experiment (data folder paths, code, hyperparams),
and allows you to visualize it using tensorboard.</p>
<pre><code class="python">from test_tube import Experiment, 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
16
17
18
19
20</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">Experiment</span><span class="p">,</span> <span class="n">HyperOptArgumentParser</span>
# exp hyperparams
args = HyperOptArgumentParser()
hparams = args.parse_args()
<span class="c1"># exp hyperparams</span>
<span class="n">args</span> <span class="o">=</span> <span class="n">HyperOptArgumentParser</span><span class="p">()</span>
<span class="n">hparams</span> <span class="o">=</span> <span class="n">args</span><span class="o">.</span><span class="n">parse_args</span><span class="p">()</span>
# this is a summaryWriter with nicer logging structure
exp = Experiment(save_dir='/some/path', create_git_tag=True)
<span class="c1"># this is a summaryWriter with nicer logging structure</span>
<span class="n">exp</span> <span class="o">=</span> <span class="n">Experiment</span><span class="p">(</span><span class="n">save_dir</span><span class="o">=</span><span class="s1">&#39;/some/path&#39;</span><span class="p">,</span> <span class="n">create_git_tag</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>
# track experiment details (must be ArgumentParser or HyperOptArgumentParser).
# each option in the parser is tracked
exp.argparse(hparams)
exp.tag({'description': 'running demo'})
<span class="c1"># track experiment details (must be ArgumentParser or HyperOptArgumentParser).</span>
<span class="c1"># each option in the parser is tracked</span>
<span class="n">exp</span><span class="o">.</span><span class="n">argparse</span><span class="p">(</span><span class="n">hparams</span><span class="p">)</span>
<span class="n">exp</span><span class="o">.</span><span class="n">tag</span><span class="p">({</span><span class="s1">&#39;description&#39;</span><span class="p">:</span> <span class="s1">&#39;running demo&#39;</span><span class="p">})</span>
# trainer uses the exp object to log exp data
trainer = Trainer(experiment=exp)
trainer.fit(model)
<span class="c1"># trainer uses the exp object to log exp data</span>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</span><span class="p">(</span><span class="n">experiment</span><span class="o">=</span><span class="n">exp</span><span class="p">)</span>
<span class="n">trainer</span><span class="o">.</span><span class="n">fit</span><span class="p">(</span><span class="n">model</span><span class="p">)</span>
# view logs at:
# tensorboard --logdir /some/path
</code></pre>
<span class="c1"># view logs at:</span>
<span class="c1"># tensorboard --logdir /some/path </span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="write-logs-file-to-csv-every-k-batches">Write logs file to csv every k batches</h4>
<p>Every k batches, lightning will write the new logs to disk</p>
<pre><code class="python"># DEFAULT (ie: save a .csv log file every 100 batches)
trainer = Trainer(log_save_interval=100)
</code></pre>
<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"># DEFAULT (ie: save a .csv log file every 100 batches)</span>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</span><span class="p">(</span><span class="n">log_save_interval</span><span class="o">=</span><span class="mi">100</span><span class="p">)</span>
</pre></div>
</td></tr></table>
+101 -48
View File
@@ -556,78 +556,131 @@
<h4 id="running-grid-search-on-a-cluster">Running grid search on a cluster</h4>
<p>To use lightning to run a hyperparameter search (grid-search or random-search) on a cluster do 4 things: </p>
<p>(1). Define the parameters for the grid search </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</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>
# subclass of argparse
parser = HyperOptArgumentParser(strategy='random_search')
parser.add_argument('--learning_rate', default=0.002, type=float, help='the learning rate')
<span class="c1"># subclass of argparse</span>
<span class="n">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">&#39;random_search&#39;</span><span class="p">)</span>
<span class="n">parser</span><span class="o">.</span><span class="n">add_argument</span><span class="p">(</span><span class="s1">&#39;--learning_rate&#39;</span><span class="p">,</span> <span class="n">default</span><span class="o">=</span><span class="mf">0.002</span><span class="p">,</span> <span class="nb">type</span><span class="o">=</span><span class="nb">float</span><span class="p">,</span> <span class="n">help</span><span class="o">=</span><span class="s1">&#39;the learning rate&#39;</span><span class="p">)</span>
# let's enable optimizing over the number of layers in the network
parser.opt_list('--nb_layers', default=2, type=int, tunable=True, options=[2, 4, 8])
<span class="c1"># let&#39;s enable optimizing over the number of layers in the network</span>
<span class="n">parser</span><span class="o">.</span><span class="n">opt_list</span><span class="p">(</span><span class="s1">&#39;--nb_layers&#39;</span><span class="p">,</span> <span class="n">default</span><span class="o">=</span><span class="mi">2</span><span class="p">,</span> <span class="nb">type</span><span class="o">=</span><span class="nb">int</span><span class="p">,</span> <span class="n">tunable</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span> <span class="n">options</span><span class="o">=</span><span class="p">[</span><span class="mi">2</span><span class="p">,</span> <span class="mi">4</span><span class="p">,</span> <span class="mi">8</span><span class="p">])</span>
hparams = parser.parse_args()
</code></pre>
<span class="n">hparams</span> <span class="o">=</span> <span class="n">parser</span><span class="o">.</span><span class="n">parse_args</span><span class="p">()</span>
</pre></div>
</td></tr></table>
<p>(2). Define the cluster options in the <a href="https://williamfalcon.github.io/test-tube/hpc/SlurmCluster/">SlurmCluster object</a> (over 5 nodes and 8 gpus) </p>
<pre><code class="python">from test_tube.hpc import SlurmCluster
<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</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="kn">from</span> <span class="nn">test_tube.hpc</span> <span class="kn">import</span> <span class="n">SlurmCluster</span>
# hyperparameters is a test-tube hyper params object
# see https://williamfalcon.github.io/test-tube/hyperparameter_optimization/HyperOptArgumentParser/
hyperparams = args.parse()
<span class="c1"># hyperparameters is a test-tube hyper params object</span>
<span class="c1"># see https://williamfalcon.github.io/test-tube/hyperparameter_optimization/HyperOptArgumentParser/</span>
<span class="n">hyperparams</span> <span class="o">=</span> <span class="n">args</span><span class="o">.</span><span class="n">parse</span><span class="p">()</span>
# init cluster
cluster = SlurmCluster(
hyperparam_optimizer=hyperparams,
log_path='/path/to/log/results/to',
python_cmd='python3'
)
<span class="c1"># init cluster</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="s1">&#39;/path/to/log/results/to&#39;</span><span class="p">,</span>
<span class="n">python_cmd</span><span class="o">=</span><span class="s1">&#39;python3&#39;</span>
<span class="p">)</span>
# let the cluster know where to email for a change in job status (ie: complete, fail, etc...)
cluster.notify_job_status(email='some@email.com', on_done=True, on_fail=True)
<span class="c1"># let the cluster know where to email for a change in job status (ie: complete, fail, etc...)</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">&#39;some@email.com&#39;</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>
# set the job options. In this instance, we'll run 20 different models
# each with its own set of hyperparameters giving each one 1 GPU (ie: taking up 20 GPUs)
cluster.per_experiment_nb_gpus = 8
cluster.per_experiment_nb_nodes = 5
<span class="c1"># set the job options. In this instance, we&#39;ll run 20 different models</span>
<span class="c1"># each with its own set of hyperparameters giving each one 1 GPU (ie: taking up 20 GPUs)</span>
<span class="n">cluster</span><span class="o">.</span><span class="n">per_experiment_nb_gpus</span> <span class="o">=</span> <span class="mi">8</span>
<span class="n">cluster</span><span class="o">.</span><span class="n">per_experiment_nb_nodes</span> <span class="o">=</span> <span class="mi">5</span>
# we'll request 10GB of memory per node
cluster.memory_mb_per_node = 10000
<span class="c1"># we&#39;ll request 10GB of memory per node</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">10000</span>
# set a walltime of 10 minues
cluster.job_time = '10:00'
</code></pre>
<span class="c1"># set a walltime of 10 minues</span>
<span class="n">cluster</span><span class="o">.</span><span class="n">job_time</span> <span class="o">=</span> <span class="s1">&#39;10:00&#39;</span>
</pre></div>
</td></tr></table>
<p>(3). Give trainer the cluster_manager in your main function: </p>
<pre><code class="python">from pytorch_lightning import Trainer
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre> 1
2
3
4
5
6
7
8
9
10</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="kn">from</span> <span class="nn">pytorch_lightning</span> <span class="kn">import</span> <span class="n">Trainer</span>
def train_fx(trial_hparams, cluster_manager, _):
# hparams has a specific set of hyperparams
<span class="k">def</span> <span class="nf">train_fx</span><span class="p">(</span><span class="n">trial_hparams</span><span class="p">,</span> <span class="n">cluster_manager</span><span class="p">,</span> <span class="n">_</span><span class="p">):</span>
<span class="c1"># hparams has a specific set of hyperparams</span>
my_model = MyLightningModel()
<span class="n">my_model</span> <span class="o">=</span> <span class="n">MyLightningModel</span><span class="p">()</span>
# give the trainer the cluster object
trainer = Trainer(cluster=cluster_manager)
trainer.fit(my_model)
</code></pre>
<span class="c1"># give the trainer the cluster object</span>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</span><span class="p">(</span><span class="n">cluster</span><span class="o">=</span><span class="n">cluster_manager</span><span class="p">)</span>
<span class="n">trainer</span><span class="o">.</span><span class="n">fit</span><span class="p">(</span><span class="n">my_model</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<p>(4). Start the grid search </p>
<pre><code class="python"># run the models on the cluster
cluster.optimize_parallel_cluster_gpu(
train_fx,
nb_trials=20,
job_name='my_grid_search_exp_name',
job_display_name='my_exp')
</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 the models on the cluster</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">train_fx</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">job_name</span><span class="o">=</span><span class="s1">&#39;my_grid_search_exp_name&#39;</span><span class="p">,</span>
<span class="n">job_display_name</span><span class="o">=</span><span class="s1">&#39;my_exp&#39;</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<p>That's it! The SlurmCluster object will automatically checkpoint the lightning model and resubmit if it runs into the walltime!</p>
<hr />
<h4 id="walltime-auto-resubmit">Walltime auto-resubmit</h4>
<p>Lightning automatically resubmits jobs when they reach the walltime. You get this behavior for free if you give lightning
a slurm cluster object.</p>
<pre><code class="python">def my_main_fx(hparams, slurm_manager, _):
trainer = Trainer(cluster=slurm_manager)
</code></pre>
<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="k">def</span> <span class="nf">my_main_fx</span><span class="p">(</span><span class="n">hparams</span><span class="p">,</span> <span class="n">slurm_manager</span><span class="p">,</span> <span class="n">_</span><span class="p">):</span>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</span><span class="p">(</span><span class="n">cluster</span><span class="o">=</span><span class="n">slurm_manager</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<p>(See the grid search example above for cluster configuration).
With this feature lightning will: </p>
+45 -24
View File
@@ -607,54 +607,75 @@
<hr />
<h4 id="accumulated-gradients">Accumulated gradients</h4>
<p>Accumulated gradients runs K small batches of size N before doing a backwards pass. The effect is a large effective batch size of size KxN. </p>
<pre><code class="python"># DEFAULT (ie: no accumulated grads)
trainer = Trainer(accumulate_grad_batches=1)
</code></pre>
<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"># DEFAULT (ie: no accumulated grads)</span>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</span><span class="p">(</span><span class="n">accumulate_grad_batches</span><span class="o">=</span><span class="mi">1</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="force-training-for-min-or-max-epochs">Force training for min or max epochs</h4>
<p>It can be useful to force training for a minimum number of epochs or limit to a max number</p>
<pre><code class="python"># DEFAULT
trainer = Trainer(min_nb_epochs=1, max_nb_epochs=1000)
</code></pre>
<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"># DEFAULT</span>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</span><span class="p">(</span><span class="n">min_nb_epochs</span><span class="o">=</span><span class="mi">1</span><span class="p">,</span> <span class="n">max_nb_epochs</span><span class="o">=</span><span class="mi">1000</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="force-disable-early-stop">Force disable early stop</h4>
<p>Use this to turn off early stopping and run training to the <a href="#force-training-for-min-or-max-epochs">max_epoch</a></p>
<pre><code class="python"># DEFAULT
trainer = Trainer(enable_early_stop=True)
</code></pre>
<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"># DEFAULT</span>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</span><span class="p">(</span><span class="n">enable_early_stop</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="gradient-clipping">Gradient Clipping</h4>
<p>Gradient clipping may be enabled to avoid exploding gradients.
Specifically, this will <a href="https://pytorch.org/docs/stable/nn.html#torch.nn.utils.clip_grad_norm_">clip the gradient norm computed over all model parameters <em>together</em></a>.</p>
<pre><code class="python"># DEFAULT (ie: don't clip)
trainer = Trainer(gradient_clip=0)
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
2
3
4
5</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># DEFAULT (ie: don&#39;t clip)</span>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</span><span class="p">(</span><span class="n">gradient_clip</span><span class="o">=</span><span class="mi">0</span><span class="p">)</span>
# clip gradients with norm above 0.5
trainer = Trainer(gradient_clip=0.5)
</code></pre>
<span class="c1"># clip gradients with norm above 0.5</span>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</span><span class="p">(</span><span class="n">gradient_clip</span><span class="o">=</span><span class="mf">0.5</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="inspect-gradient-norms">Inspect gradient norms</h4>
<p>Looking at grad norms can help you figure out where training might be going wrong.</p>
<pre><code class="python"># DEFAULT (-1 doesn't track norms)
trainer = Trainer(track_grad_norm=-1)
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
2
3
4
5</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># DEFAULT (-1 doesn&#39;t track norms)</span>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</span><span class="p">(</span><span class="n">track_grad_norm</span><span class="o">=-</span><span class="mi">1</span><span class="p">)</span>
# track the LP norm (P=2 here)
trainer = Trainer(track_grad_norm=2)
</code></pre>
<span class="c1"># track the LP norm (P=2 here)</span>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</span><span class="p">(</span><span class="n">track_grad_norm</span><span class="o">=</span><span class="mi">2</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="set-how-much-of-the-training-set-to-check">Set how much of the training set to check</h4>
<p>If you don't want to check 100% of the training set (for debugging or if it's huge), set this flag</p>
<pre><code class="python"># DEFAULT
trainer = Trainer(train_percent_check=1.0)
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
2
3
4
5</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># DEFAULT</span>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</span><span class="p">(</span><span class="n">train_percent_check</span><span class="o">=</span><span class="mf">1.0</span><span class="p">)</span>
# check 10% only
trainer = Trainer(train_percent_check=0.1)
</code></pre>
<span class="c1"># check 10% only</span>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</span><span class="p">(</span><span class="n">train_percent_check</span><span class="o">=</span><span class="mf">0.1</span><span class="p">)</span>
</pre></div>
</td></tr></table>
+40 -21
View File
@@ -595,46 +595,65 @@ Lightning will run 5 steps of validation in the beginning of training as a sanit
<hr />
<h4 id="check-validation-every-n-epochs">Check validation every n epochs</h4>
<p>If you have a small dataset you might want to check validation every n epochs</p>
<pre><code class="python"># DEFAULT
trainer = Trainer(check_val_every_n_epoch=1)
</code></pre>
<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"># DEFAULT</span>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</span><span class="p">(</span><span class="n">check_val_every_n_epoch</span><span class="o">=</span><span class="mi">1</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="set-how-much-of-the-validation-set-to-check">Set how much of the validation set to check</h4>
<p>If you don't want to check 100% of the validation set (for debugging or if it's huge), set this flag</p>
<pre><code class="python"># DEFAULT
trainer = Trainer(val_percent_check=1.0)
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
2
3
4
5</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># DEFAULT</span>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</span><span class="p">(</span><span class="n">val_percent_check</span><span class="o">=</span><span class="mf">1.0</span><span class="p">)</span>
# check 10% only
trainer = Trainer(val_percent_check=0.1)
</code></pre>
<span class="c1"># check 10% only</span>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</span><span class="p">(</span><span class="n">val_percent_check</span><span class="o">=</span><span class="mf">0.1</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="set-how-much-of-the-test-set-to-check">Set how much of the test set to check</h4>
<p>If you don't want to check 100% of the test set (for debugging or if it's huge), set this flag</p>
<pre><code class="python"># DEFAULT
trainer = Trainer(test_percent_check=1.0)
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
2
3
4
5</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># DEFAULT</span>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</span><span class="p">(</span><span class="n">test_percent_check</span><span class="o">=</span><span class="mf">1.0</span><span class="p">)</span>
# check 10% only
trainer = Trainer(test_percent_check=0.1)
</code></pre>
<span class="c1"># check 10% only</span>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</span><span class="p">(</span><span class="n">test_percent_check</span><span class="o">=</span><span class="mf">0.1</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="set-validation-check-frequency-within-1-training-epoch">Set validation check frequency within 1 training epoch</h4>
<p>For large datasets it's often desirable to check validation multiple times within a training loop</p>
<pre><code class="python"># DEFAULT
trainer = Trainer(val_check_interval=0.95)
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
2
3
4
5</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># DEFAULT</span>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</span><span class="p">(</span><span class="n">val_check_interval</span><span class="o">=</span><span class="mf">0.95</span><span class="p">)</span>
# check every .25 of an epoch
trainer = Trainer(val_check_interval=0.25)
</code></pre>
<span class="c1"># check every .25 of an epoch </span>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</span><span class="p">(</span><span class="n">val_check_interval</span><span class="o">=</span><span class="mf">0.25</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="set-the-number-of-validation-sanity-steps">Set the number of validation sanity steps</h4>
<p>Lightning runs a few steps of validation in the beginning of training. This avoids crashing in the validation loop sometime deep into a lengthy training loop.</p>
<pre><code class="python"># DEFAULT
trainer = Trainer(nb_sanity_val_steps=5)
</code></pre>
<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"># DEFAULT</span>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</span><span class="p">(</span><span class="n">nb_sanity_val_steps</span><span class="o">=</span><span class="mi">5</span><span class="p">)</span>
</pre></div>
</td></tr></table>
+30 -16
View File
@@ -607,29 +607,41 @@
<h4 id="fast-dev-run">Fast dev run</h4>
<p>This flag is meant for debugging a full train/val/test loop. It'll activate callbacks, everything but only with 1 training and 1 validation batch.
Use this to debug a full run of your program quickly</p>
<pre><code class="python"># DEFAULT
trainer = Trainer(fast_dev_run=False)
</code></pre>
<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"># DEFAULT</span>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</span><span class="p">(</span><span class="n">fast_dev_run</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="inspect-gradient-norms">Inspect gradient norms</h4>
<p>Looking at grad norms can help you figure out where training might be going wrong.</p>
<pre><code class="python"># DEFAULT (-1 doesn't track norms)
trainer = Trainer(track_grad_norm=-1)
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
2
3
4
5</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># DEFAULT (-1 doesn&#39;t track norms)</span>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</span><span class="p">(</span><span class="n">track_grad_norm</span><span class="o">=-</span><span class="mi">1</span><span class="p">)</span>
# track the LP norm (P=2 here)
trainer = Trainer(track_grad_norm=2)
</code></pre>
<span class="c1"># track the LP norm (P=2 here)</span>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</span><span class="p">(</span><span class="n">track_grad_norm</span><span class="o">=</span><span class="mi">2</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="make-model-overfit-on-subset-of-data">Make model overfit on subset of data</h4>
<p>A useful debugging trick is to make your model overfit a tiny fraction of the data.</p>
<pre><code class="python"># DEFAULT don't overfit (ie: normal training)
trainer = Trainer(overfit_pct=0.0)
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
2
3
4
5</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># DEFAULT don&#39;t overfit (ie: normal training)</span>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</span><span class="p">(</span><span class="n">overfit_pct</span><span class="o">=</span><span class="mf">0.0</span><span class="p">)</span>
# overfit on 1% of data
trainer = Trainer(overfit_pct=0.01)
</code></pre>
<span class="c1"># overfit on 1% of data </span>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</span><span class="p">(</span><span class="n">overfit_pct</span><span class="o">=</span><span class="mf">0.01</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="print-the-parameter-count-by-layer">Print the parameter count by layer</h4>
@@ -637,9 +649,11 @@ trainer = Trainer(overfit_pct=0.01)
<hr />
<h4 id="print-which-gradients-are-nan">Print which gradients are nan</h4>
<p>This option prints a list of tensors with nan gradients.</p>
<pre><code class="python"># DEFAULT
trainer = Trainer(print_nan_grads=False)
</code></pre>
<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"># DEFAULT</span>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</span><span class="p">(</span><span class="n">print_nan_grads</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="log-gpu-usage">Log GPU usage</h4>
+97 -52
View File
@@ -673,104 +673,149 @@ To enable a hook, simply override the method in your LightningModule and the tra
<hr />
<h4 id="on_epoch_start">on_epoch_start</h4>
<p>Called in the training loop at the very beginning of the epoch. </p>
<pre><code class="python">def on_epoch_start(self):
# do something when the epoch starts
</code></pre>
<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="k">def</span> <span class="nf">on_epoch_start</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
<span class="c1"># do something when the epoch starts</span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="on_epoch_end">on_epoch_end</h4>
<p>Called in the training loop at the very end of the epoch. </p>
<pre><code class="python">def on_epoch_end(self):
# do something when the epoch ends
</code></pre>
<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="k">def</span> <span class="nf">on_epoch_end</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
<span class="c1"># do something when the epoch ends </span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="on_batch_start">on_batch_start</h4>
<p>Called in the training loop before anything happens for that batch. </p>
<pre><code class="python">def on_batch_start(self):
# do something when the batch starts
</code></pre>
<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="k">def</span> <span class="nf">on_batch_start</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
<span class="c1"># do something when the batch starts</span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="on_batch_end">on_batch_end</h4>
<p>Called in the training loop after the batch. </p>
<pre><code class="python">def on_batch_end(self):
# do something when the batch ends
</code></pre>
<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="k">def</span> <span class="nf">on_batch_end</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
<span class="c1"># do something when the batch ends </span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="on_pre_performance_check">on_pre_performance_check</h4>
<p>Called at the very beginning of the validation loop. </p>
<pre><code class="python">def on_pre_performance_check(self):
# do something before validation starts
</code></pre>
<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="k">def</span> <span class="nf">on_pre_performance_check</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
<span class="c1"># do something before validation starts </span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="on_post_performance_check">on_post_performance_check</h4>
<p>Called at the very end of the validation loop. </p>
<pre><code class="python">def on_post_performance_check(self):
# do something before validation end
</code></pre>
<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="k">def</span> <span class="nf">on_post_performance_check</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
<span class="c1"># do something before validation end</span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="on_tng_metrics">on_tng_metrics</h4>
<p>Called in the training loop, right before metrics are logged.
Although you can log at any time by using self.experiment, you can use
this callback to modify what will be logged.</p>
<pre><code class="python">def on_tng_metrics(self, metrics):
# do something before validation end
</code></pre>
<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="k">def</span> <span class="nf">on_tng_metrics</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">metrics</span><span class="p">):</span>
<span class="c1"># do something before validation end</span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="optimizer_step">optimizer_step</h4>
<p>Calls .step() and .zero_grad for each optimizer.<br />
You can override this method to adjust how you do the optimizer step for each optimizer</p>
<p>Called once per optimizer</p>
<pre><code class="python"># DEFAULT
def optimizer_step(self, current_epoch, batch_nb, optimizer, optimizer_i):
optimizer.step()
optimizer.zero_grad()
<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</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># DEFAULT</span>
<span class="k">def</span> <span class="nf">optimizer_step</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">current_epoch</span><span class="p">,</span> <span class="n">batch_nb</span><span class="p">,</span> <span class="n">optimizer</span><span class="p">,</span> <span class="n">optimizer_i</span><span class="p">):</span>
<span class="n">optimizer</span><span class="o">.</span><span class="n">step</span><span class="p">()</span>
<span class="n">optimizer</span><span class="o">.</span><span class="n">zero_grad</span><span class="p">()</span>
# Alternating schedule for optimizer steps (ie: GANs)
def optimizer_step(self, current_epoch, batch_nb, optimizer, optimizer_i):
# update generator opt every 2 steps
if optimizer_i == 0:
if batch_nb % 2 == 0 :
optimizer.step()
optimizer.zero_grad()
<span class="c1"># Alternating schedule for optimizer steps (ie: GANs) </span>
<span class="k">def</span> <span class="nf">optimizer_step</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">current_epoch</span><span class="p">,</span> <span class="n">batch_nb</span><span class="p">,</span> <span class="n">optimizer</span><span class="p">,</span> <span class="n">optimizer_i</span><span class="p">):</span>
<span class="c1"># update generator opt every 2 steps</span>
<span class="k">if</span> <span class="n">optimizer_i</span> <span class="o">==</span> <span class="mi">0</span><span class="p">:</span>
<span class="k">if</span> <span class="n">batch_nb</span> <span class="o">%</span> <span class="mi">2</span> <span class="o">==</span> <span class="mi">0</span> <span class="p">:</span>
<span class="n">optimizer</span><span class="o">.</span><span class="n">step</span><span class="p">()</span>
<span class="n">optimizer</span><span class="o">.</span><span class="n">zero_grad</span><span class="p">()</span>
# update discriminator opt every 4 steps
if optimizer_i == 1:
if batch_nb % 4 == 0 :
optimizer.step()
optimizer.zero_grad()
<span class="c1"># update discriminator opt every 4 steps</span>
<span class="k">if</span> <span class="n">optimizer_i</span> <span class="o">==</span> <span class="mi">1</span><span class="p">:</span>
<span class="k">if</span> <span class="n">batch_nb</span> <span class="o">%</span> <span class="mi">4</span> <span class="o">==</span> <span class="mi">0</span> <span class="p">:</span>
<span class="n">optimizer</span><span class="o">.</span><span class="n">step</span><span class="p">()</span>
<span class="n">optimizer</span><span class="o">.</span><span class="n">zero_grad</span><span class="p">()</span>
# ...
# add as many optimizers as you want
</code></pre>
<span class="c1"># ...</span>
<span class="c1"># add as many optimizers as you want </span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="on_before_zero_grad">on_before_zero_grad</h4>
<p>Called in the training loop after taking an optimizer step and before zeroing grads.
Good place to inspect weight information with weights updated.</p>
<p>Called once per optimizer</p>
<pre><code class="python">def on_before_zero_grad(self, optimizer):
# do something with the optimizer or inspect it.
</code></pre>
<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="k">def</span> <span class="nf">on_before_zero_grad</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">optimizer</span><span class="p">):</span>
<span class="c1"># do something with the optimizer or inspect it. </span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="on_after_backward">on_after_backward</h4>
<p>Called in the training loop after model.backward()
This is the ideal place to inspect or log gradient information </p>
<pre><code class="python">def on_after_backward(self):
# example to inspect gradient information in tensorboard
if self.trainer.global_step % 25 == 0: # don't make the tf file huge
params = self.state_dict()
for k, v in params.items():
grads = v
name = k
self.experiment.add_histogram(tag=name, values=grads, global_step=self.trainer.global_step)
</code></pre>
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
2
3
4
5
6
7
8</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="k">def</span> <span class="nf">on_after_backward</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
<span class="c1"># example to inspect gradient information in tensorboard</span>
<span class="k">if</span> <span class="bp">self</span><span class="o">.</span><span class="n">trainer</span><span class="o">.</span><span class="n">global_step</span> <span class="o">%</span> <span class="mi">25</span> <span class="o">==</span> <span class="mi">0</span><span class="p">:</span> <span class="c1"># don&#39;t make the tf file huge</span>
<span class="n">params</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">state_dict</span><span class="p">()</span>
<span class="k">for</span> <span class="n">k</span><span class="p">,</span> <span class="n">v</span> <span class="ow">in</span> <span class="n">params</span><span class="o">.</span><span class="n">items</span><span class="p">():</span>
<span class="n">grads</span> <span class="o">=</span> <span class="n">v</span>
<span class="n">name</span> <span class="o">=</span> <span class="n">k</span>
<span class="bp">self</span><span class="o">.</span><span class="n">experiment</span><span class="o">.</span><span class="n">add_histogram</span><span class="p">(</span><span class="n">tag</span><span class="o">=</span><span class="n">name</span><span class="p">,</span> <span class="n">values</span><span class="o">=</span><span class="n">grads</span><span class="p">,</span> <span class="n">global_step</span><span class="o">=</span><span class="bp">self</span><span class="o">.</span><span class="n">trainer</span><span class="o">.</span><span class="n">global_step</span><span class="p">)</span>
</pre></div>
</td></tr></table>
+11 -5
View File
@@ -495,13 +495,19 @@
<p>[<a href="https://github.com/williamFalcon/pytorch-lightning/blob/master/pytorch_lightning/models/trainer.py">Github Code</a>]</p>
<p>The lightning trainer abstracts best practices for running a training, val, test routine. It calls parts of your model when it wants to hand over full control and otherwise makes training assumptions which are now standard practice in AI research.</p>
<p>This is the basic use of the trainer:</p>
<pre><code class="python">from pytorch_lightning import Trainer
<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="kn">from</span> <span class="nn">pytorch_lightning</span> <span class="kn">import</span> <span class="n">Trainer</span>
model = LightningTemplate()
<span class="n">model</span> <span class="o">=</span> <span class="n">LightningTemplate</span><span class="p">()</span>
trainer = Trainer()
trainer.fit(model)
</code></pre>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</span><span class="p">()</span>
<span class="n">trainer</span><span class="o">.</span><span class="n">fit</span><span class="p">(</span><span class="n">model</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<p>But of course the fun is in all the advanced things it can do:</p>
<p><strong>Checkpointing</strong> </p>
+182 -65
View File
@@ -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">&#39;__main__&#39;</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">&#39;__main__&#39;</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">&#39;random_search&#39;</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):
&quot;&quot;&quot;
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=&#39;test_tube_exp&#39;,
debug=True,
save_dir=log_dir,
version=0,
autosave=False,
description='test demo'
description=&#39;test demo&#39;
)
# 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 = &#39;{}/{}/{}&#39;.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">&#39;add_email_here&#39;</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">&#39;48:00:00&#39;</span>
<span class="n">cluster</span><span class="o">.</span><span class="n">gpu_type</span> <span class="o">=</span> <span class="s1">&#39;1080ti&#39;</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">&#39;source activate pytorch_lightning&#39;</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">&#39;_&#39;</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">&#39;submitting jobs...&#39;</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>
File diff suppressed because one or more lines are too long
BIN
View File
Binary file not shown.