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> </ul>
<hr /> <hr />
<h3 id="minimal-example">Minimal example</h3> <h3 id="minimal-example">Minimal example</h3>
<pre><code class="python">import os <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre> 1
import torch 2
from torch.nn import functional as F 3
from torch.utils.data import DataLoader 4
from torchvision.datasets import MNIST 5
import torchvision.transforms as transforms 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): <span class="k">def</span> <span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
super(CoolModel, self).__init__() <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>
# not the best model... <span class="c1"># not the best model...</span>
self.l1 = torch.nn.Linear(28 * 28, 10) <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): <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>
return torch.relu(self.l1(x.view(x.size(0), -1))) <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): <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>
# REQUIRED <span class="c1"># REQUIRED</span>
x, y = batch <span class="n">x</span><span class="p">,</span> <span class="n">y</span> <span class="o">=</span> <span class="n">batch</span>
y_hat = self.forward(x) <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>
return {'loss': F.cross_entropy(y_hat, y)(y_hat, y)} <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): <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>
# OPTIONAL <span class="c1"># OPTIONAL</span>
x, y = batch <span class="n">x</span><span class="p">,</span> <span class="n">y</span> <span class="o">=</span> <span class="n">batch</span>
y_hat = self.forward(x) <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>
return {'val_loss': F.cross_entropy(y_hat, y)(y_hat, y)} <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): <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>
# OPTIONAL <span class="c1"># OPTIONAL</span>
avg_loss = torch.stack([x['val_loss'] for x in outputs]).mean() <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>
return {'avg_val_loss': avg_loss} <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): <span class="k">def</span> <span class="nf">configure_optimizers</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
# REQUIRED <span class="c1"># REQUIRED</span>
return [torch.optim.Adam(self.parameters(), lr=0.02)] <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 <span class="nd">@pl.data_loader</span>
def tng_dataloader(self): <span class="k">def</span> <span class="nf">tng_dataloader</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
return DataLoader(MNIST(os.getcwd(), train=True, download=True, transform=transforms.ToTensor()), batch_size=32) <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 <span class="nd">@pl.data_loader</span>
def val_dataloader(self): <span class="k">def</span> <span class="nf">val_dataloader</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
# OPTIONAL <span class="c1"># OPTIONAL</span>
# can also return a list of val dataloaders <span class="c1"># can also return a list of val dataloaders</span>
return DataLoader(MNIST(os.getcwd(), train=True, download=True, transform=transforms.ToTensor()), batch_size=32) <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 <span class="nd">@pl.data_loader</span>
def test_dataloader(self): <span class="k">def</span> <span class="nf">test_dataloader</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
# OPTIONAL <span class="c1"># OPTIONAL</span>
return DataLoader(MNIST(os.getcwd(), train=True, download=True, transform=transforms.ToTensor()), batch_size=32) <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>
</code></pre> </pre></div>
</td></tr></table>
<hr /> <hr />
<h3 id="how-do-these-methods-fit-into-the-broader-training">How do these methods fit into the broader training?</h3> <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> <h2 id="required-methods">Required Methods</h2>
<h3 id="training_step">training_step</h3> <h3 id="training_step">training_step</h3>
<pre><code class="python">def training_step(self, data_batch, batch_nb) <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>
</code></pre> </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>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> <p><strong>Params</strong> </p>
@@ -1102,57 +1157,90 @@ class CoolModel(pl.LightningModule):
</tbody> </tbody>
</table> </table>
<p><strong>Example</strong></p> <p><strong>Example</strong></p>
<pre><code class="python">def training_step(self, data_batch, batch_nb): <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre> 1
x, y, z = data_batch 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 <span class="c1"># implement your own</span>
out = self.forward(x) <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>
loss = self.loss(out, x) <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 = { <span class="n">output</span> <span class="o">=</span> <span class="p">{</span>
'loss': loss, # required <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>
'prog': {'tng_loss': loss, 'batch_nb': batch_nb} # optional <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 <span class="c1"># return a dict</span>
return output <span class="k">return</span> <span class="n">output</span>
</code></pre> </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> <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) <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
def training_step(self, data_batch, batch_nb, optimizer_idx): 2
if optimizer_idx == 0: 3
# do training_step with encoder 4
if optimizer_idx == 1: 5
# do training_step with decoder 6</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># Multiple optimizers (ie: GANs) </span>
</code></pre> <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 /> <hr />
<h3 id="tng_dataloader">tng_dataloader</h3> <h3 id="tng_dataloader">tng_dataloader</h3>
<pre><code class="python">@pl.data_loader <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
def tng_dataloader(self) 2</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="nd">@pl.data_loader</span>
</code></pre> <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> <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> <h5 id="return">Return</h5>
<p>PyTorch DataLoader</p> <p>PyTorch DataLoader</p>
<p><strong>Example</strong></p> <p><strong>Example</strong></p>
<pre><code class="python">@pl.data_loader <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre> 1
def tng_dataloader(self): 2
transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,), (1.0,))]) 3
dataset = MNIST(root='/path/to/mnist/', train=True, transform=transform, download=True) 4
loader = torch.utils.data.DataLoader( 5
dataset=dataset, 6
batch_size=self.hparams.batch_size, 7
shuffle=True 8
) 9
return loader 10</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="nd">@pl.data_loader</span>
</code></pre> <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 /> <hr />
<h3 id="configure_optimizers">configure_optimizers</h3> <h3 id="configure_optimizers">configure_optimizers</h3>
<pre><code class="python">def configure_optimizers(self) <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>
</code></pre> </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. <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> 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> <h5 id="return_1">Return</h5>
<p>List or Tuple - List of optimizers with an optional second list of learning-rate schedulers</p> <p>List or Tuple - List of optimizers with an optional second list of learning-rate schedulers</p>
<p><strong>Example</strong></p> <p><strong>Example</strong></p>
<pre><code class="python"># most cases <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre> 1
def configure_optimizers(self): 2
opt = Adam(self.parameters(), lr=0.01) 3
return [opt] 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 <span class="c1"># gan example, with scheduler for discriminator</span>
def configure_optimizers(self): <span class="k">def</span> <span class="nf">configure_optimizers</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
generator_opt = Adam(self.model_gen.parameters(), lr=0.01) <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>
disriminator_opt = Adam(self.model_disc.parameters(), lr=0.02) <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>
discriminator_sched = CosineAnnealing(discriminator_opt, T_max=10) <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>
return [generator_opt, disriminator_opt], [discriminator_sched] <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>
</code></pre> </pre></div>
</td></tr></table>
<p>If you need to control how often those optimizers step or override the default .step() schedule, override <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> 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> <h2 id="optional-methods">Optional Methods</h2>
<h3 id="validation_step">validation_step</h3> <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: <span class="c1"># if have multiple val dataloaders: </span>
def validation_step(self, data_batch, batch_nb, dataloader_idx) <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>
</code></pre> </pre></div>
</td></tr></table>
<p><strong>OPTIONAL</strong> <br /> <p><strong>OPTIONAL</strong> <br />
If you don't need to validate you don't need to implement this method. </p> 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> </tbody>
</table> </table>
<p><strong>Example</strong></p> <p><strong>Example</strong></p>
<pre><code class="python"># CASE 1: A single validation dataset <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre> 1
def validation_step(self, data_batch, batch_nb): 2
x, y, z = data_batch 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 <span class="c1"># implement your own</span>
out = self.forward(x) <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>
loss = self.loss(out, x) <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 <span class="c1"># calculate acc</span>
labels_hat = torch.argmax(out, dim=1) <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>
val_acc = torch.sum(y == labels_hat).item() / (len(y) * 1.0) <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... <span class="c1"># all optional...</span>
# return whatever you need for the collation function validation_end <span class="c1"># return whatever you need for the collation function validation_end</span>
output = OrderedDict({ <span class="n">output</span> <span class="o">=</span> <span class="n">OrderedDict</span><span class="p">({</span>
'val_loss': loss_val, <span class="s1">&#39;val_loss&#39;</span><span class="p">:</span> <span class="n">loss_val</span><span class="p">,</span>
'val_acc': torch.tensor(val_acc), # everything must be a tensor <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 <span class="c1"># return an optional dict</span>
return output <span class="k">return</span> <span class="n">output</span>
</code></pre> </pre></div>
</td></tr></table>
<p>If you pass in multiple validation datasets, validation_step will have an additional argument.</p> <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 <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
def validation_step(self, data_batch, batch_nb, dataset_idx): 2
# dataset_idx tells you which dataset this is. 3</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># CASE 2: multiple validation datasets</span>
</code></pre> <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> <p>The <code>dataset_idx</code> corresponds to the order of datasets returned in <code>val_dataloader</code>. </p>
<hr /> <hr />
<h3 id="validation_end">validation_end</h3> <h3 id="validation_end">validation_end</h3>
<pre><code class="python">def validation_end(self, outputs) <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>
</code></pre> </pre></div>
</td></tr></table>
<p>If you didn't define a validation_step, this won't be called. </p> <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> <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> </tbody>
</table> </table>
<p><strong>Example</strong></p> <p><strong>Example</strong></p>
<pre><code class="python">def validation_end(self, outputs): <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre> 1
&quot;&quot;&quot; 2
Called at the end of validation to aggregate outputs 3
:param outputs: list of individual outputs of each validation step 4
:return: 5
&quot;&quot;&quot; 6
val_loss_mean = 0 7
val_acc_mean = 0 8
for output in outputs: 9
val_loss_mean += output['val_loss'] 10
val_acc_mean += output['val_acc'] 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) <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>
val_acc_mean /= len(outputs) <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>
tqdm_dic = {'val_loss': val_loss_mean.item(), 'val_acc': val_acc_mean.item()} <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>
return tqdm_dic <span class="k">return</span> <span class="n">tqdm_dic</span>
</code></pre> </pre></div>
</td></tr></table>
<hr /> <hr />
<h3 id="on_save_checkpoint">on_save_checkpoint</h3> <h3 id="on_save_checkpoint">on_save_checkpoint</h3>
<pre><code class="python">def on_save_checkpoint(self, checkpoint) <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>
</code></pre> </pre></div>
</td></tr></table>
<p>Called by lightning to checkpoint your model. Lightning saves the training state (current epoch, global_step, etc) <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 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> <h5 id="return_2">Return</h5>
<p>Nothing</p> <p>Nothing</p>
<p><strong>Example</strong></p> <p><strong>Example</strong></p>
<pre><code class="python">def on_save_checkpoint(self, checkpoint): <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
# 99% of use cases you don't need to implement this method 2
checkpoint['something_cool_i_want_to_save'] = my_cool_pickable_object 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>
</code></pre> <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 /> <hr />
<h3 id="on_load_checkpoint">on_load_checkpoint</h3> <h3 id="on_load_checkpoint">on_load_checkpoint</h3>
<pre><code class="python">def on_load_checkpoint(self, checkpoint) <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>
</code></pre> </pre></div>
</td></tr></table>
<p>Called by lightning to restore your model. Lighting auto-restores global step, epoch, etc... <p>Called by lightning to restore your model. Lighting auto-restores global step, epoch, etc...
It also restores the model state_dict. 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> <h5 id="return_3">Return</h5>
<p>Nothing </p> <p>Nothing </p>
<p><strong>Example</strong></p> <p><strong>Example</strong></p>
<pre><code class="python">def on_load_checkpoint(self, checkpoint): <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
# 99% of the time you don't need to implement this method 2
self.something_cool_i_want_to_save = checkpoint['something_cool_i_want_to_save'] 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>
</code></pre> <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 /> <hr />
<h3 id="val_dataloader">val_dataloader</h3> <h3 id="val_dataloader">val_dataloader</h3>
<pre><code class="python">@pl.data_loader <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
def tng_dataloader(self) 2</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="nd">@pl.data_loader</span>
</code></pre> <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 /> <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> 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> <h5 id="return_4">Return</h5>
<p>PyTorch DataLoader or list of PyTorch Dataloaders. </p> <p>PyTorch DataLoader or list of PyTorch Dataloaders. </p>
<p><strong>Example</strong></p> <p><strong>Example</strong></p>
<pre><code class="python">@pl.data_loader <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre> 1
def val_dataloader(self): 2
transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,), (1.0,))]) 3
dataset = MNIST(root='/path/to/mnist/', train=False, transform=transform, download=True) 4
loader = torch.utils.data.DataLoader( 5
dataset=dataset, 6
batch_size=self.hparams.batch_size, 7
shuffle=True 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 <span class="c1"># can also return multiple dataloaders </span>
@pl.data_loader <span class="nd">@pl.data_loader</span>
def val_dataloader(self): <span class="k">def</span> <span class="nf">val_dataloader</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
return [loader_a, loader_b, ..., loader_n] <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>
</code></pre> </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> <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> which matches the order here. </p>
<hr /> <hr />
<h3 id="test_dataloader">test_dataloader</h3> <h3 id="test_dataloader">test_dataloader</h3>
<pre><code class="python">@pl.data_loader <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
def test_dataloader(self) 2</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="nd">@pl.data_loader</span>
</code></pre> <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 /> <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> 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> <h5 id="return_5">Return</h5>
<p>PyTorch DataLoader</p> <p>PyTorch DataLoader</p>
<p><strong>Example</strong></p> <p><strong>Example</strong></p>
<pre><code class="python">@pl.data_loader <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre> 1
def test_dataloader(self): 2
transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,), (1.0,))]) 3
dataset = MNIST(root='/path/to/mnist/', train=False, transform=transform, download=True) 4
loader = torch.utils.data.DataLoader( 5
dataset=dataset, 6
batch_size=self.hparams.batch_size, 7
shuffle=True 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 <span class="k">return</span> <span class="n">loader</span>
</code></pre> </pre></div>
</td></tr></table>
<hr /> <hr />
<h3 id="update_tng_log_metrics">update_tng_log_metrics</h3> <h3 id="update_tng_log_metrics">update_tng_log_metrics</h3>
<pre><code class="python">def update_tng_log_metrics(self, logs) <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>
</code></pre> </pre></div>
</td></tr></table>
<p>Called by lightning right before it logs metrics for this batch. <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> This is a chance to ammend or add to the metrics about to be logged.</p>
<h5 id="return_6">Return</h5> <h5 id="return_6">Return</h5>
<p>Dict </p> <p>Dict </p>
<p><strong>Example</strong></p> <p><strong>Example</strong></p>
<pre><code class="python">def update_tng_log_metrics(self, logs): <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
# modify or add to logs 2
return logs 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>
</code></pre> <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 /> <hr />
<h3 id="add_model_specific_args">add_model_specific_args</h3> <h3 id="add_model_specific_args">add_model_specific_args</h3>
<pre><code class="python">@staticmethod <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
def add_model_specific_args(parent_parser, root_dir) 2</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="nd">@staticmethod</span>
</code></pre> <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. <p>Lightning has a list of default argparse commands.
This method is your chance to add or modify commands specific to your model. 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> <h5 id="return_7">Return</h5>
<p>An argument parser</p> <p>An argument parser</p>
<p><strong>Example</strong></p> <p><strong>Example</strong></p>
<pre><code class="python">@staticmethod <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre> 1
def add_model_specific_args(parent_parser, root_dir): 2
parser = HyperOptArgumentParser(strategy=parent_parser.strategy, parents=[parent_parser]) 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 <span class="c1"># param overwrites</span>
# parser.set_defaults(gradient_clip=5.0) <span class="c1"># parser.set_defaults(gradient_clip=5.0)</span>
# network params <span class="c1"># network params</span>
parser.opt_list('--drop_prob', default=0.2, options=[0.2, 0.5], type=float, tunable=False) <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>
parser.add_argument('--in_features', default=28*28) <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>
parser.add_argument('--out_features', default=10) <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>
parser.add_argument('--hidden_dim', default=50000) # use 500 for CPU, 50000 for GPU to see speed difference <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 <span class="c1"># data</span>
parser.add_argument('--data_root', default=os.path.join(root_dir, 'mnist'), type=str) <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) <span class="c1"># training params (opt)</span>
parser.opt_list('--learning_rate', default=0.001, type=float, options=[0.0001, 0.0005, 0.001, 0.005], <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>
tunable=False) <span class="n">tunable</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>
parser.opt_list('--batch_size', default=256, type=int, options=[32, 64, 128, 256], tunable=False) <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>
parser.opt_list('--optimizer_name', default='adam', type=str, options=['adam'], tunable=False) <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>
return parser <span class="k">return</span> <span class="n">parser</span>
</code></pre> </pre></div>
</td></tr></table>
+32 -17
View File
@@ -564,26 +564,39 @@
<hr /> <hr />
<h3 id="freeze">freeze</h3> <h3 id="freeze">freeze</h3>
<p>Freeze all params for inference</p> <p>Freeze all params for inference</p>
<pre><code class="python">model = MyLightningModule(...) <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
model.freeze() 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>
</code></pre> <span class="n">model</span><span class="o">.</span><span class="n">freeze</span><span class="p">()</span>
</pre></div>
</td></tr></table>
<hr /> <hr />
<h3 id="load_from_metrics">load_from_metrics</h3> <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. <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> 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( <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre> 1
weights_path='/path/to/pytorch_checkpoint.ckpt', 2
tags_csv='/path/to/test_tube/experiment/version/meta_tags.csv', 3
on_gpu=True, 4
map_location=None 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 <span class="c1"># predict</span>
pretrained_model.eval() <span class="n">pretrained_model</span><span class="o">.</span><span class="n">eval</span><span class="p">()</span>
pretrained_model.freeze() <span class="n">pretrained_model</span><span class="o">.</span><span class="n">freeze</span><span class="p">()</span>
y_hat = pretrained_model(x) <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>
</code></pre> </pre></div>
</td></tr></table>
<p><strong>Params</strong> </p> <p><strong>Params</strong> </p>
<table> <table>
@@ -617,9 +630,11 @@ y_hat = pretrained_model(x)
<hr /> <hr />
<h3 id="unfreeze">unfreeze</h3> <h3 id="unfreeze">unfreeze</h3>
<p>Unfreeze all params for inference</p> <p>Unfreeze all params for inference</p>
<pre><code class="python">model = MyLightningModule(...) <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
model.unfreeze() 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>
</code></pre> <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 /> <hr />
<h4 id="experiment">experiment</h4> <h4 id="experiment">experiment</h4>
<p>An instance of test-tube Experiment which you can use to log anything for tensorboarX. </p> <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(...) <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
self.experiment.log({'val_loss': 0.9}) 2
self.experiment.add_scalars(...) 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>
</code></pre> <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 /> <hr />
<h4 id="global_step">global_step</h4> <h4 id="global_step">global_step</h4>
@@ -683,10 +686,13 @@ self.experiment.add_scalars(...)
<hr /> <hr />
<h4 id="trainer">trainer</h4> <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> <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 <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
self.trainer.current_epoch 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>
</code></pre> <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> <h2 id="debugging">Debugging</h2>
<p>The LightningModule also offers these tricks to help debug. </p> <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> <h4 id="example_input_array">example_input_array</h4>
<p>In the LightningModule init, you can set a dummy tensor for this property <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> to get a print out of sizes coming into and out of every layer. </p>
<pre><code class="python">def __init__(self): <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
# put the dimensions of the first input to your system 2
self.example_input_array = torch.rand(5, 28 * 28) 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>
</code></pre> <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 /> <hr />
<h3 id="model-saving">Model saving</h3> <h3 id="model-saving">Model saving</h3>
<p>To enable checkpointing, define the checkpoint callback and give it to the trainer.</p> <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( <span class="n">checkpoint_callback</span> <span class="o">=</span> <span class="n">ModelCheckpoint</span><span class="p">(</span>
filepath='/path/to/store/weights.ckpt', <span class="n">filepath</span><span class="o">=</span><span class="s1">&#39;/path/to/store/weights.ckpt&#39;</span><span class="p">,</span>
save_best_only=True, <span class="n">save_best_only</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span>
verbose=True, <span class="n">verbose</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span>
monitor='val_loss', <span class="n">monitor</span><span class="o">=</span><span class="s1">&#39;val_loss&#39;</span><span class="p">,</span>
mode='min' <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) <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>
</code></pre> </pre></div>
</td></tr></table>
<hr /> <hr />
<h3 id="restoring-training-session">Restoring training session</h3> <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 /> 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> 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> <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) <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>
trainer = Trainer(experiment=exp) <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 <span class="c1"># this fit call loads model weights and trainer state</span>
# the trainer continues seamlessly from where you left off <span class="c1"># the trainer continues seamlessly from where you left off</span>
# without having to do anything else. <span class="c1"># without having to do anything else.</span>
trainer.fit(model) <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>
</code></pre> </pre></div>
</td></tr></table>
<p>The trainer restores:<br /> <p>The trainer restores:<br />
- global_step <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 <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> 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> <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"> <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre> 1
self.global_step = checkpoint['global_step'] 2
self.current_epoch = checkpoint['epoch'] 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 <span class="c1"># restore the optimizers</span>
optimizer_states = checkpoint['optimizer_states'] <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>
for optimizer, opt_state in zip(self.optimizers, optimizer_states): <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>
optimizer.load_state_dict(opt_state) <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 <span class="c1"># restore the lr schedulers</span>
lr_schedulers = checkpoint['lr_schedulers'] <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>
for scheduler, lrs_state in zip(self.lr_schedulers, lr_schedulers): <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>
scheduler.load_state_dict(lrs_state) <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 <span class="c1"># uses the model you passed into trainer </span>
model.load_state_dict(checkpoint['state_dict']) <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>
</code></pre> </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. <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> For multi-node training you must use DistributedDataParallel. </p>
<p>You can toggle between each mode by setting this flag.</p> <p>You can toggle between each mode by setting this flag.</p>
<pre><code class="python"># DEFAULT uses DataParallel <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
trainer = Trainer(distributed_backend='dp') 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 <span class="c1"># change to distributed data parallel</span>
trainer = Trainer(distributed_backend='ddp') <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>
</code></pre> </pre></div>
</td></tr></table>
<p>If you request multiple nodes, the back-end will auto-switch to ddp. <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> 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> <h4 id="cuda-flags">CUDA flags</h4>
<p>CUDA flags make certain GPUs visible to your script. <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> 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 <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
# os.environ[&quot;CUDA_DEVICE_ORDER&quot;] = &quot;PCI_BUS_ID&quot; 2
# os.environ[&quot;CUDA_VISIBLE_DEVICES&quot;] = &quot;0&quot; 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>
</code></pre> <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 /> <hr />
<h4 id="16-bit-mixed-precision">16-bit mixed precision</h4> <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 /> <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> 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 <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
$ cd apex 2
$ pip install -v --no-cache-dir --global-option=&quot;--cpp_ext&quot; --global-option=&quot;--cuda_ext&quot; ./ 3</pre></div></td><td class="code"><div class="codehilite"><pre><span></span>$ git clone https://github.com/NVIDIA/apex
</code></pre> $ <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> <p>then set this use_amp to True.</p>
<pre><code class="python"># DEFAULT <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
trainer = Trainer(amp_level='O2', use_amp=False) 2</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># DEFAULT</span>
</code></pre> <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 /> <hr />
<h4 id="single-gpu">Single-gpu</h4> <h4 id="single-gpu">Single-gpu</h4>
<p>Make sure you're on a GPU machine. </p> <p>Make sure you're on a GPU machine. </p>
<pre><code class="python"># DEFAULT <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
trainer = Trainer(gpus=[0]) 2</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># DEFAULT</span>
</code></pre> <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 /> <hr />
<h4 id="multi-gpu">multi-gpu</h4> <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. <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> 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) <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
trainer = Trainer(gpus=[0,1,2,3,4,5,6,7], distributed_backend='dp') 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 <span class="c1"># RECOMMENDED use DistributedDataParallel</span>
trainer = Trainer(gpus=[0,1,2,3,4,5,6,7], distributed_backend='ddp') <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>
</code></pre> </pre></div>
</td></tr></table>
<hr /> <hr />
<h4 id="multi-node">Multi-node</h4> <h4 id="multi-node">Multi-node</h4>
<p>Multi-node training is easily done by specifying these flags.</p> <p>Multi-node training is easily done by specifying these flags.</p>
<pre><code class="python"># train on 12*8 GPUs <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
trainer = Trainer(gpus=[0,1,2,3,4,5,6,7], nb_gpu_nodes=12) 2</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># train on 12*8 GPUs</span>
</code></pre> <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> <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( <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre> 1
hyperparam_optimizer=test_tube.HyperOptArgumentParser(), 2
log_path='/some/path/to/save', 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 <span class="c1"># OPTIONAL FLAGS WHICH MAY BE CLUSTER DEPENDENT</span>
# which interface your nodes use for communication <span class="c1"># which interface your nodes use for communication</span>
cluster.add_command('export NCCL_SOCKET_IFNAME=^docker0,lo') <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 <span class="c1"># see output of the NCCL connection process</span>
# NCCL is how the nodes talk to each other <span class="c1"># NCCL is how the nodes talk to each other</span>
cluster.add_command('export NCCL_DEBUG=INFO') <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. <span class="c1"># setting a master port here is a good idea.</span>
cluster.add_command('export MASTER_PORT=%r' % PORT) <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 <span class="c1"># good to load the latest NCCL version</span>
cluster.load_modules(['NCCL/2.4.7-1-cuda.10.0']) <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 <span class="c1"># configure cluster</span>
cluster.per_experiment_nb_nodes = 12 <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>
cluster.per_experiment_nb_gpus = 8 <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') <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>
</code></pre> </pre></div>
</td></tr></table>
<p>Finally, make sure to add a distributed sampler to your dataset. The distributed sampler copies a <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> portion of your dataset onto each GPU. (World_size = gpus_per_node * nb_nodes). </p>
<pre><code class="python"># ie: this: <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
dataset = myDataset() 2
dataloader = Dataloader(dataset) 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: <span class="c1"># becomes:</span>
dataset = myDataset() <span class="n">dataset</span> <span class="o">=</span> <span class="n">myDataset</span><span class="p">()</span>
dist_sampler = torch.utils.data.distributed.DistributedSampler(dataset) <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>
dataloader = Dataloader(dataset, sampler=dist_sampler) <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>
</code></pre> </pre></div>
</td></tr></table>
<hr /> <hr />
<h4 id="self-balancing-architecture">Self-balancing architecture</h4> <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> <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 /> <hr />
<h4 id="display-metrics-in-progress-bar">Display metrics in progress bar</h4> <h4 id="display-metrics-in-progress-bar">Display metrics in progress bar</h4>
<pre><code class="python"># DEFAULT <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
trainer = Trainer(progress_bar=True) 2</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># DEFAULT</span>
</code></pre> <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 /> <hr />
<h4 id="log-metric-row-every-k-batches">Log metric row every k batches</h4> <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> <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) <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
trainer = Trainer(add_log_row_interval=10) 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>
</code></pre> <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 /> <hr />
<h4 id="process-position">Process position</h4> <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. <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> Lightning will stack progress bars according to this value. </p>
<pre><code class="python"># DEFAULT <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
trainer = Trainer(process_position=0) 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 <span class="c1"># if this is the second model on the node, show the second progress bar below</span>
trainer = Trainer(process_position=1) <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>
</code></pre> </pre></div>
</td></tr></table>
<hr /> <hr />
<h4 id="save-a-snapshot-of-all-hyperparameters">Save a snapshot of all hyperparameters</h4> <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. <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> 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(...) <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>
Trainer(experiment=exp) <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>
</code></pre> </pre></div>
</td></tr></table>
<hr /> <hr />
<h4 id="snapshot-code-for-a-training-run">Snapshot code for a training run</h4> <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. <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> 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) <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>
Trainer(experiment=exp) <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>
</code></pre> </pre></div>
</td></tr></table>
<hr /> <hr />
<h3 id="tensorboard-support">Tensorboard support</h3> <h3 id="tensorboard-support">Tensorboard support</h3>
<p>In the LightningModule you can access the experiment logger by doing:</p> <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 <span class="c1"># add image</span>
# Look at PyTorch SummaryWriter docs for what you can do. <span class="c1"># Look at PyTorch SummaryWriter docs for what you can do. </span>
self.experiment.add_image(...) <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>
</code></pre> </pre></div>
</td></tr></table>
<p>The experiment object is a strict subclass of PyTorch SummaryWriter. However, this class <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), also snapshots every detail about the experiment (data folder paths, code, hyperparams),
and allows you to visualize it using tensorboard.</p> 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 <span class="c1"># exp hyperparams</span>
args = HyperOptArgumentParser() <span class="n">args</span> <span class="o">=</span> <span class="n">HyperOptArgumentParser</span><span class="p">()</span>
hparams = args.parse_args() <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 <span class="c1"># this is a summaryWriter with nicer logging structure</span>
exp = Experiment(save_dir='/some/path', create_git_tag=True) <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). <span class="c1"># track experiment details (must be ArgumentParser or HyperOptArgumentParser).</span>
# each option in the parser is tracked <span class="c1"># each option in the parser is tracked</span>
exp.argparse(hparams) <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>
exp.tag({'description': 'running demo'}) <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 <span class="c1"># trainer uses the exp object to log exp data</span>
trainer = Trainer(experiment=exp) <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>
trainer.fit(model) <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: <span class="c1"># view logs at:</span>
# tensorboard --logdir /some/path <span class="c1"># tensorboard --logdir /some/path </span>
</code></pre> </pre></div>
</td></tr></table>
<hr /> <hr />
<h4 id="write-logs-file-to-csv-every-k-batches">Write logs file to csv every k batches</h4> <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> <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) <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
trainer = Trainer(log_save_interval=100) 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>
</code></pre> <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> <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>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> <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 <span class="c1"># subclass of argparse</span>
parser = HyperOptArgumentParser(strategy='random_search') <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>
parser.add_argument('--learning_rate', default=0.002, type=float, help='the learning rate') <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 <span class="c1"># let&#39;s enable optimizing over the number of layers in the network</span>
parser.opt_list('--nb_layers', default=2, type=int, tunable=True, options=[2, 4, 8]) <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() <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>
</code></pre> </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> <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 <span class="c1"># hyperparameters is a test-tube hyper params object</span>
# see https://williamfalcon.github.io/test-tube/hyperparameter_optimization/HyperOptArgumentParser/ <span class="c1"># see https://williamfalcon.github.io/test-tube/hyperparameter_optimization/HyperOptArgumentParser/</span>
hyperparams = args.parse() <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 <span class="c1"># init cluster</span>
cluster = SlurmCluster( <span class="n">cluster</span> <span class="o">=</span> <span class="n">SlurmCluster</span><span class="p">(</span>
hyperparam_optimizer=hyperparams, <span class="n">hyperparam_optimizer</span><span class="o">=</span><span class="n">hyperparams</span><span class="p">,</span>
log_path='/path/to/log/results/to', <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>
python_cmd='python3' <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...) <span class="c1"># let the cluster know where to email for a change in job status (ie: complete, fail, etc...)</span>
cluster.notify_job_status(email='some@email.com', on_done=True, on_fail=True) <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 <span class="c1"># set the job options. In this instance, we&#39;ll run 20 different models</span>
# each with its own set of hyperparameters giving each one 1 GPU (ie: taking up 20 GPUs) <span class="c1"># each with its own set of hyperparameters giving each one 1 GPU (ie: taking up 20 GPUs)</span>
cluster.per_experiment_nb_gpus = 8 <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.per_experiment_nb_nodes = 5 <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 <span class="c1"># we&#39;ll request 10GB of memory per node</span>
cluster.memory_mb_per_node = 10000 <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 <span class="c1"># set a walltime of 10 minues</span>
cluster.job_time = '10:00' <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>
</code></pre> </pre></div>
</td></tr></table>
<p>(3). Give trainer the cluster_manager in your main function: </p> <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, _): <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>
# hparams has a specific set of hyperparams <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 <span class="c1"># give the trainer the cluster object</span>
trainer = Trainer(cluster=cluster_manager) <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>
trainer.fit(my_model) <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>
</code></pre> </td></tr></table>
<p>(4). Start the grid search </p> <p>(4). Start the grid search </p>
<pre><code class="python"># run the models on the cluster <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
cluster.optimize_parallel_cluster_gpu( 2
train_fx, 3
nb_trials=20, 4
job_name='my_grid_search_exp_name', 5
job_display_name='my_exp') 6</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># run the models on the cluster</span>
</code></pre> <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> <p>That's it! The SlurmCluster object will automatically checkpoint the lightning model and resubmit if it runs into the walltime!</p>
<hr /> <hr />
<h4 id="walltime-auto-resubmit">Walltime auto-resubmit</h4> <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 <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> a slurm cluster object.</p>
<pre><code class="python">def my_main_fx(hparams, slurm_manager, _): <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
trainer = Trainer(cluster=slurm_manager) 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>
</code></pre> <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). <p>(See the grid search example above for cluster configuration).
With this feature lightning will: </p> With this feature lightning will: </p>
+45 -24
View File
@@ -607,54 +607,75 @@
<hr /> <hr />
<h4 id="accumulated-gradients">Accumulated gradients</h4> <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> <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) <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
trainer = Trainer(accumulate_grad_batches=1) 2</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># DEFAULT (ie: no accumulated grads)</span>
</code></pre> <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 /> <hr />
<h4 id="force-training-for-min-or-max-epochs">Force training for min or max epochs</h4> <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> <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 <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
trainer = Trainer(min_nb_epochs=1, max_nb_epochs=1000) 2</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># DEFAULT</span>
</code></pre> <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 /> <hr />
<h4 id="force-disable-early-stop">Force disable early stop</h4> <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> <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 <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
trainer = Trainer(enable_early_stop=True) 2</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># DEFAULT</span>
</code></pre> <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 /> <hr />
<h4 id="gradient-clipping">Gradient Clipping</h4> <h4 id="gradient-clipping">Gradient Clipping</h4>
<p>Gradient clipping may be enabled to avoid exploding gradients. <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> 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) <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
trainer = Trainer(gradient_clip=0) 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 <span class="c1"># clip gradients with norm above 0.5</span>
trainer = Trainer(gradient_clip=0.5) <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>
</code></pre> </pre></div>
</td></tr></table>
<hr /> <hr />
<h4 id="inspect-gradient-norms">Inspect gradient norms</h4> <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> <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) <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
trainer = Trainer(track_grad_norm=-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) <span class="c1"># track the LP norm (P=2 here)</span>
trainer = Trainer(track_grad_norm=2) <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>
</code></pre> </pre></div>
</td></tr></table>
<hr /> <hr />
<h4 id="set-how-much-of-the-training-set-to-check">Set how much of the training set to check</h4> <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> <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 <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
trainer = Trainer(train_percent_check=1.0) 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 <span class="c1"># check 10% only</span>
trainer = Trainer(train_percent_check=0.1) <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>
</code></pre> </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 /> <hr />
<h4 id="check-validation-every-n-epochs">Check validation every n epochs</h4> <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> <p>If you have a small dataset you might want to check validation every n epochs</p>
<pre><code class="python"># DEFAULT <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
trainer = Trainer(check_val_every_n_epoch=1) 2</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># DEFAULT</span>
</code></pre> <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 /> <hr />
<h4 id="set-how-much-of-the-validation-set-to-check">Set how much of the validation set to check</h4> <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> <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 <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
trainer = Trainer(val_percent_check=1.0) 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 <span class="c1"># check 10% only</span>
trainer = Trainer(val_percent_check=0.1) <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>
</code></pre> </pre></div>
</td></tr></table>
<hr /> <hr />
<h4 id="set-how-much-of-the-test-set-to-check">Set how much of the test set to check</h4> <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> <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 <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
trainer = Trainer(test_percent_check=1.0) 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 <span class="c1"># check 10% only</span>
trainer = Trainer(test_percent_check=0.1) <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>
</code></pre> </pre></div>
</td></tr></table>
<hr /> <hr />
<h4 id="set-validation-check-frequency-within-1-training-epoch">Set validation check frequency within 1 training epoch</h4> <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> <p>For large datasets it's often desirable to check validation multiple times within a training loop</p>
<pre><code class="python"># DEFAULT <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
trainer = Trainer(val_check_interval=0.95) 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 <span class="c1"># check every .25 of an epoch </span>
trainer = Trainer(val_check_interval=0.25) <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>
</code></pre> </pre></div>
</td></tr></table>
<hr /> <hr />
<h4 id="set-the-number-of-validation-sanity-steps">Set the number of validation sanity steps</h4> <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> <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 <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
trainer = Trainer(nb_sanity_val_steps=5) 2</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># DEFAULT</span>
</code></pre> <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> <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. <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> Use this to debug a full run of your program quickly</p>
<pre><code class="python"># DEFAULT <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
trainer = Trainer(fast_dev_run=False) 2</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># DEFAULT</span>
</code></pre> <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 /> <hr />
<h4 id="inspect-gradient-norms">Inspect gradient norms</h4> <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> <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) <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
trainer = Trainer(track_grad_norm=-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) <span class="c1"># track the LP norm (P=2 here)</span>
trainer = Trainer(track_grad_norm=2) <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>
</code></pre> </pre></div>
</td></tr></table>
<hr /> <hr />
<h4 id="make-model-overfit-on-subset-of-data">Make model overfit on subset of data</h4> <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> <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) <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
trainer = Trainer(overfit_pct=0.0) 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 <span class="c1"># overfit on 1% of data </span>
trainer = Trainer(overfit_pct=0.01) <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>
</code></pre> </pre></div>
</td></tr></table>
<hr /> <hr />
<h4 id="print-the-parameter-count-by-layer">Print the parameter count by layer</h4> <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 /> <hr />
<h4 id="print-which-gradients-are-nan">Print which gradients are nan</h4> <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> <p>This option prints a list of tensors with nan gradients.</p>
<pre><code class="python"># DEFAULT <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
trainer = Trainer(print_nan_grads=False) 2</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># DEFAULT</span>
</code></pre> <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 /> <hr />
<h4 id="log-gpu-usage">Log GPU usage</h4> <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 /> <hr />
<h4 id="on_epoch_start">on_epoch_start</h4> <h4 id="on_epoch_start">on_epoch_start</h4>
<p>Called in the training loop at the very beginning of the epoch. </p> <p>Called in the training loop at the very beginning of the epoch. </p>
<pre><code class="python">def on_epoch_start(self): <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
# do something when the epoch starts 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>
</code></pre> <span class="c1"># do something when the epoch starts</span>
</pre></div>
</td></tr></table>
<hr /> <hr />
<h4 id="on_epoch_end">on_epoch_end</h4> <h4 id="on_epoch_end">on_epoch_end</h4>
<p>Called in the training loop at the very end of the epoch. </p> <p>Called in the training loop at the very end of the epoch. </p>
<pre><code class="python">def on_epoch_end(self): <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
# do something when the epoch ends 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>
</code></pre> <span class="c1"># do something when the epoch ends </span>
</pre></div>
</td></tr></table>
<hr /> <hr />
<h4 id="on_batch_start">on_batch_start</h4> <h4 id="on_batch_start">on_batch_start</h4>
<p>Called in the training loop before anything happens for that batch. </p> <p>Called in the training loop before anything happens for that batch. </p>
<pre><code class="python">def on_batch_start(self): <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
# do something when the batch starts 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>
</code></pre> <span class="c1"># do something when the batch starts</span>
</pre></div>
</td></tr></table>
<hr /> <hr />
<h4 id="on_batch_end">on_batch_end</h4> <h4 id="on_batch_end">on_batch_end</h4>
<p>Called in the training loop after the batch. </p> <p>Called in the training loop after the batch. </p>
<pre><code class="python">def on_batch_end(self): <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
# do something when the batch ends 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>
</code></pre> <span class="c1"># do something when the batch ends </span>
</pre></div>
</td></tr></table>
<hr /> <hr />
<h4 id="on_pre_performance_check">on_pre_performance_check</h4> <h4 id="on_pre_performance_check">on_pre_performance_check</h4>
<p>Called at the very beginning of the validation loop. </p> <p>Called at the very beginning of the validation loop. </p>
<pre><code class="python">def on_pre_performance_check(self): <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
# do something before validation starts 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>
</code></pre> <span class="c1"># do something before validation starts </span>
</pre></div>
</td></tr></table>
<hr /> <hr />
<h4 id="on_post_performance_check">on_post_performance_check</h4> <h4 id="on_post_performance_check">on_post_performance_check</h4>
<p>Called at the very end of the validation loop. </p> <p>Called at the very end of the validation loop. </p>
<pre><code class="python">def on_post_performance_check(self): <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
# do something before validation end 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>
</code></pre> <span class="c1"># do something before validation end</span>
</pre></div>
</td></tr></table>
<hr /> <hr />
<h4 id="on_tng_metrics">on_tng_metrics</h4> <h4 id="on_tng_metrics">on_tng_metrics</h4>
<p>Called in the training loop, right before metrics are logged. <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 Although you can log at any time by using self.experiment, you can use
this callback to modify what will be logged.</p> this callback to modify what will be logged.</p>
<pre><code class="python">def on_tng_metrics(self, metrics): <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
# do something before validation end 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>
</code></pre> <span class="c1"># do something before validation end</span>
</pre></div>
</td></tr></table>
<hr /> <hr />
<h4 id="optimizer_step">optimizer_step</h4> <h4 id="optimizer_step">optimizer_step</h4>
<p>Calls .step() and .zero_grad for each optimizer.<br /> <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> You can override this method to adjust how you do the optimizer step for each optimizer</p>
<p>Called once per optimizer</p> <p>Called once per optimizer</p>
<pre><code class="python"># DEFAULT <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre> 1
def optimizer_step(self, current_epoch, batch_nb, optimizer, optimizer_i): 2
optimizer.step() 3
optimizer.zero_grad() 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) <span class="c1"># Alternating schedule for optimizer steps (ie: GANs) </span>
def optimizer_step(self, current_epoch, batch_nb, optimizer, optimizer_i): <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>
# update generator opt every 2 steps <span class="c1"># update generator opt every 2 steps</span>
if optimizer_i == 0: <span class="k">if</span> <span class="n">optimizer_i</span> <span class="o">==</span> <span class="mi">0</span><span class="p">:</span>
if batch_nb % 2 == 0 : <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>
optimizer.step() <span class="n">optimizer</span><span class="o">.</span><span class="n">step</span><span class="p">()</span>
optimizer.zero_grad() <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 <span class="c1"># update discriminator opt every 4 steps</span>
if optimizer_i == 1: <span class="k">if</span> <span class="n">optimizer_i</span> <span class="o">==</span> <span class="mi">1</span><span class="p">:</span>
if batch_nb % 4 == 0 : <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>
optimizer.step() <span class="n">optimizer</span><span class="o">.</span><span class="n">step</span><span class="p">()</span>
optimizer.zero_grad() <span class="n">optimizer</span><span class="o">.</span><span class="n">zero_grad</span><span class="p">()</span>
# ... <span class="c1"># ...</span>
# add as many optimizers as you want <span class="c1"># add as many optimizers as you want </span>
</code></pre> </pre></div>
</td></tr></table>
<hr /> <hr />
<h4 id="on_before_zero_grad">on_before_zero_grad</h4> <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. <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> Good place to inspect weight information with weights updated.</p>
<p>Called once per optimizer</p> <p>Called once per optimizer</p>
<pre><code class="python">def on_before_zero_grad(self, optimizer): <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
# do something with the optimizer or inspect it. 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>
</code></pre> <span class="c1"># do something with the optimizer or inspect it. </span>
</pre></div>
</td></tr></table>
<hr /> <hr />
<h4 id="on_after_backward">on_after_backward</h4> <h4 id="on_after_backward">on_after_backward</h4>
<p>Called in the training loop after model.backward() <p>Called in the training loop after model.backward()
This is the ideal place to inspect or log gradient information </p> This is the ideal place to inspect or log gradient information </p>
<pre><code class="python">def on_after_backward(self): <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
# example to inspect gradient information in tensorboard 2
if self.trainer.global_step % 25 == 0: # don't make the tf file huge 3
params = self.state_dict() 4
for k, v in params.items(): 5
grads = v 6
name = k 7
self.experiment.add_histogram(tag=name, values=grads, global_step=self.trainer.global_step) 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>
</code></pre> <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>[<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>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> <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() <span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</span><span class="p">()</span>
trainer.fit(model) <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>
</code></pre> </pre></div>
</td></tr></table>
<p>But of course the fun is in all the advanced things it can do:</p> <p>But of course the fun is in all the advanced things it can do:</p>
<p><strong>Checkpointing</strong> </p> <p><strong>Checkpointing</strong> </p>
+182 -65
View File
@@ -602,9 +602,11 @@
<h3 id="template-model-definition">Template model definition</h3> <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> <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 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 /> <hr />
<h3 id="trainer-example">Trainer Example</h3> <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. <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 Inside the main we parse training arguments with whatever hyperparameters we want. Your LightningModule will have a
chance to add hyperparameters. </p> 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 <span class="c1"># use default args given by lightning</span>
root_dir = os.path.split(os.path.dirname(sys.modules['__main__'].__file__))[0] <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>
parent_parser = HyperOptArgumentParser(strategy='random_search', add_help=False) <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>
add_default_args(parent_parser, root_dir) <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 <span class="c1"># allow model to overwrite or extend args</span>
parser = ExampleModel.add_model_specific_args(parent_parser) <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>
hyperparams = parser.parse_args() <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 <span class="c1"># train model</span>
main(hyperparams) <span class="n">main</span><span class="p">(</span><span class="n">hyperparams</span><span class="p">)</span>
</code></pre> </pre></div>
</td></tr></table>
<p><strong>Main Function</strong> </p> <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. <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 /> - hparams: a configuration of hyperparameters. <br />
- slurm_manager: Slurm cluster manager object (can be None) - 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> - 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; &quot;&quot;&quot;
Main training routine specific for this project Main training routine specific for this project
:param hparams: :param hparams:
@@ -644,12 +712,12 @@ The main function should have 3 arguments: <br />
# init experiment # init experiment
log_dir = os.path.dirname(os.path.realpath(__file__)) log_dir = os.path.dirname(os.path.realpath(__file__))
exp = Experiment( exp = Experiment(
name='test_tube_exp', name=&#39;test_tube_exp&#39;,
debug=True, debug=True,
save_dir=log_dir, save_dir=log_dir,
version=0, version=0,
autosave=False, autosave=False,
description='test demo' description=&#39;test demo&#39;
) )
# set the hparams for the experiment # set the hparams for the experiment
@@ -667,7 +735,7 @@ The main function should have 3 arguments: <br />
mode=hparams.early_stop_mode 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( checkpoint = ModelCheckpoint(
filepath=model_save_path, filepath=model_save_path,
save_function=None, save_function=None,
@@ -687,73 +755,122 @@ The main function should have 3 arguments: <br />
# train model # train model
trainer.fit(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 <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 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> 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> <p>So, calling main(hyperparams) runs the model with the default argparse arguments. </p>
<pre><code class="python">main(hyperparams) <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>
</code></pre> </pre></div>
</td></tr></table>
<hr /> <hr />
<h4 id="cpu-hyperparameter-search">CPU hyperparameter search</h4> <h4 id="cpu-hyperparameter-search">CPU hyperparameter search</h4>
<pre><code class="python"># run a grid search over 20 hyperparameter combinations. <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
hyperparams.optimize_parallel_cpu( 2
main_local, 3
nb_trials=20, 4
nb_workers=1 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>
</code></pre> <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 /> <hr />
<h4 id="hyperparameter-search-on-a-single-or-multiple-gpus">Hyperparameter search on a single or multiple GPUs</h4> <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. <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
hyperparams.optimize_parallel_gpu( 2
main_local, 3
nb_trials=20, 4
nb_workers=1, 5
gpus=[0,1,2,3] 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>
</code></pre> <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 /> <hr />
<h4 id="hyperparameter-search-on-a-slurm-hpc-cluster">Hyperparameter search on a SLURM HPC cluster</h4> <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): <table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre> 1
# enable cluster training 2
cluster = SlurmCluster( 3
hyperparam_optimizer=hyperparams, 4
log_path=hyperparams.tt_save_path, 5
test_tube_exp_name=hyperparams.tt_name 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 <span class="c1"># email for cluster coms</span>
cluster.notify_job_status(email='add_email_here', on_done=True, on_fail=True) <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 <span class="c1"># configure cluster</span>
cluster.per_experiment_nb_gpus = hyperparams.per_experiment_nb_gpus <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>
cluster.job_time = '48:00:00' <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>
cluster.gpu_type = '1080ti' <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>
cluster.memory_mb_per_node = 48000 <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 <span class="c1"># any modules for code to run in env</span>
cluster.add_command('source activate pytorch_lightning') <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 <span class="c1"># name of exp</span>
job_display_name = hyperparams.tt_name.split('_')[0] <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>
job_display_name = job_display_name[0:3] <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 <span class="c1"># run hopt</span>
print('submitting jobs...') <span class="k">print</span><span class="p">(</span><span class="s1">&#39;submitting jobs...&#39;</span><span class="p">)</span>
cluster.optimize_parallel_cluster_gpu( <span class="n">cluster</span><span class="o">.</span><span class="n">optimize_parallel_cluster_gpu</span><span class="p">(</span>
main, <span class="n">main</span><span class="p">,</span>
nb_trials=hyperparams.nb_hopt_trials, <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>
job_name=job_display_name <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 <span class="c1"># run cluster hyperparameter search </span>
optimize_on_cluster(hyperparams) <span class="n">optimize_on_cluster</span><span class="p">(</span><span class="n">hyperparams</span><span class="p">)</span>
</code></pre> </pre></div>
</td></tr></table>
File diff suppressed because one or more lines are too long
BIN
View File
Binary file not shown.