Deployed f51b459 with MkDocs version: 1.0.4

This commit is contained in:
William Falcon
2019-08-31 02:03:39 -05:00
parent 55692eb3e3
commit 65a4a54205
18 changed files with 345 additions and 98 deletions
+2 -2
View File
@@ -152,7 +152,7 @@
</div>
<div class="md-source__repository">
GitHub
williamFalcon/pytorch-lightning
</div>
</a>
</div>
@@ -199,7 +199,7 @@
</div>
<div class="md-source__repository">
GitHub
williamFalcon/pytorch-lightning
</div>
</a>
</div>
@@ -156,7 +156,7 @@
</div>
<div class="md-source__repository">
GitHub
williamFalcon/pytorch-lightning
</div>
</a>
</div>
@@ -203,7 +203,7 @@
</div>
<div class="md-source__repository">
GitHub
williamFalcon/pytorch-lightning
</div>
</a>
</div>
@@ -369,6 +369,20 @@
validation_end
</a>
</li>
<li class="md-nav__item">
<a href="#test_step" title="test_step" class="md-nav__link">
test_step
</a>
</li>
<li class="md-nav__item">
<a href="#test_end" title="test_end" class="md-nav__link">
test_end
</a>
</li>
<li class="md-nav__item">
@@ -818,6 +832,20 @@
validation_end
</a>
</li>
<li class="md-nav__item">
<a href="#test_step" title="test_step" class="md-nav__link">
test_step
</a>
</li>
<li class="md-nav__item">
<a href="#test_end" title="test_end" class="md-nav__link">
test_end
</a>
</li>
<li class="md-nav__item">
@@ -978,7 +1006,9 @@
<p><strong>Optional</strong>: </p>
<ul>
<li><a href="./#validation_step">validation_step</a> </li>
<li><a href="./#validation_end">validation_end</a> </li>
<li><a href="./#validation_end">validation_end</a> </li>
<li><a href="./#test_step">test_step</a> </li>
<li><a href="./#test_end">test_end</a> </li>
<li><a href="./#val_dataloader">val_dataloader</a> </li>
<li><a href="./#test_dataloader">test_dataloader</a> </li>
<li><a href="./#on_save_checkpoint">on_save_checkpoint</a> </li>
@@ -1041,7 +1071,19 @@
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>
54
55
56
57
58
59
60
61
62
63
64
65
66</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>
@@ -1077,6 +1119,17 @@
<span class="n">avg_loss</span> <span class="o">=</span> <span class="n">torch</span><span class="o">.</span><span class="n">stack</span><span class="p">([</span><span class="n">x</span><span class="p">[</span><span class="s1">&#39;val_loss&#39;</span><span class="p">]</span> <span class="k">for</span> <span class="n">x</span> <span class="ow">in</span> <span class="n">outputs</span><span class="p">])</span><span class="o">.</span><span class="n">mean</span><span class="p">()</span>
<span class="k">return</span> <span class="p">{</span><span class="s1">&#39;avg_val_loss&#39;</span><span class="p">:</span> <span class="n">avg_loss</span><span class="p">}</span>
<span class="k">def</span> <span class="nf">test_step</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">batch</span><span class="p">,</span> <span class="n">batch_nb</span><span class="p">):</span>
<span class="c1"># OPTIONAL</span>
<span class="n">x</span><span class="p">,</span> <span class="n">y</span> <span class="o">=</span> <span class="n">batch</span>
<span class="n">y_hat</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">forward</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
<span class="k">return</span> <span class="p">{</span><span class="s1">&#39;test_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="k">def</span> <span class="nf">test_end</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">outputs</span><span class="p">):</span>
<span class="c1"># OPTIONAL</span>
<span class="n">avg_loss</span> <span class="o">=</span> <span class="n">torch</span><span class="o">.</span><span class="n">stack</span><span class="p">([</span><span class="n">x</span><span class="p">[</span><span class="s1">&#39;test_loss&#39;</span><span class="p">]</span> <span class="k">for</span> <span class="n">x</span> <span class="ow">in</span> <span class="n">outputs</span><span class="p">])</span><span class="o">.</span><span class="n">mean</span><span class="p">()</span>
<span class="k">return</span> <span class="p">{</span><span class="s1">&#39;avg_test_loss&#39;</span><span class="p">:</span> <span class="n">avg_loss</span><span class="p">}</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="c1"># REQUIRED</span>
<span class="k">return</span> <span class="p">[</span><span class="n">torch</span><span class="o">.</span><span class="n">optim</span><span class="o">.</span><span class="n">Adam</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">parameters</span><span class="p">(),</span> <span class="n">lr</span><span class="o">=</span><span class="mf">0.02</span><span class="p">)]</span>
@@ -1094,6 +1147,7 @@
<span class="nd">@pl.data_loader</span>
<span class="k">def</span> <span class="nf">test_dataloader</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
<span class="c1"># OPTIONAL</span>
<span class="c1"># can also return a list of test dataloaders</span>
<span class="k">return</span> <span class="n">DataLoader</span><span class="p">(</span><span class="n">MNIST</span><span class="p">(</span><span class="n">os</span><span class="o">.</span><span class="n">getcwd</span><span class="p">(),</span> <span class="n">train</span><span class="o">=</span><span class="bp">False</span><span class="p">,</span> <span class="n">download</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span> <span class="n">transform</span><span class="o">=</span><span class="n">transforms</span><span class="o">.</span><span class="n">ToTensor</span><span class="p">()),</span> <span class="n">batch_size</span><span class="o">=</span><span class="mi">32</span><span class="p">)</span>
</pre></div>
</td></tr></table>
@@ -1294,9 +1348,11 @@ the <a href="https://williamfalcon.github.io/pytorch-lightning/Trainer/hooks/#op
<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>
4
5</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># if you have one val dataloader:</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="c1"># if have multiple val dataloaders: </span>
<span class="c1"># if you have multiple val dataloaders: </span>
<span class="k">def</span> <span class="nf">validation_step</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">data_batch</span><span class="p">,</span> <span class="n">batch_nb</span><span class="p">,</span> <span class="n">dataloader_idx</span><span class="p">)</span>
</pre></div>
</td></tr></table>
@@ -1304,7 +1360,7 @@ the <a href="https://williamfalcon.github.io/pytorch-lightning/Trainer/hooks/#op
<p><strong>OPTIONAL</strong> <br />
If you don't need to validate you don't need to implement this method. </p>
<p>In this step you'd normally generate examples or calculate anything of interest such as accuracy. </p>
<p>The dict you return here will be available in the validation_end method. </p>
<p>The dict you return here will be available in the <code>validation_end</code> method. </p>
<p><strong>Params</strong> </p>
<table>
<thead>
@@ -1374,11 +1430,11 @@ If you don't need to validate you don't need to implement this method. </p>
26
27</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>
<span class="n">x</span><span class="p">,</span> <span class="n">y</span> <span class="o">=</span> <span class="n">data_batch</span>
<span class="c1"># implement your own</span>
<span class="n">out</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">forward</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
<span class="n">loss</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">loss</span><span class="p">(</span><span class="n">out</span><span class="p">,</span> <span class="n">x</span><span class="p">)</span>
<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">y</span><span class="p">)</span>
<span class="c1"># log 6 example images</span>
<span class="c1"># or generated text... or whatever</span>
@@ -1488,6 +1544,195 @@ If you don't need to validate you don't need to implement this method. </p>
</pre></div>
</td></tr></table>
<h3 id="test_step">test_step</h3>
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
2
3
4
5</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># if you have one test dataloader:</span>
<span class="k">def</span> <span class="nf">test_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="c1"># if you have multiple test dataloaders: </span>
<span class="k">def</span> <span class="nf">test_step</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">data_batch</span><span class="p">,</span> <span class="n">batch_nb</span><span class="p">,</span> <span class="n">dataloader_idx</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<p><strong>OPTIONAL</strong> <br />
If you don't need to test you don't need to implement this method. </p>
<p>In this step you'd normally generate examples or calculate anything of interest such as accuracy. </p>
<p>The dict you return here will be available in the <code>test_end</code> method. </p>
<p>This function is used when you execute <code>trainer.test()</code>.</p>
<p><strong>Params</strong> </p>
<table>
<thead>
<tr>
<th>Param</th>
<th>description</th>
</tr>
</thead>
<tbody>
<tr>
<td>data_batch</td>
<td>The output of your dataloader. A tensor, tuple or list</td>
</tr>
<tr>
<td>batch_nb</td>
<td>Integer displaying which batch this is</td>
</tr>
<tr>
<td>dataloader_i</td>
<td>Integer displaying which dataloader this is (only if multiple test datasets used)</td>
</tr>
</tbody>
</table>
<p><strong>Return</strong> </p>
<table>
<thead>
<tr>
<th>Return</th>
<th>description</th>
<th>optional</th>
</tr>
</thead>
<tbody>
<tr>
<td>dict</td>
<td>Dict or OrderedDict with metrics to display in progress bar. All keys must be tensors.</td>
<td>Y</td>
</tr>
</tbody>
</table>
<p><strong>Example</strong></p>
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre> 1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># CASE 1: A single test dataset</span>
<span class="k">def</span> <span class="nf">test_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="o">=</span> <span class="n">data_batch</span>
<span class="c1"># implement your own</span>
<span class="n">out</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">forward</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
<span class="n">loss</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">loss</span><span class="p">(</span><span class="n">out</span><span class="p">,</span> <span class="n">y</span><span class="p">)</span>
<span class="c1"># calculate acc</span>
<span class="n">labels_hat</span> <span class="o">=</span> <span class="n">torch</span><span class="o">.</span><span class="n">argmax</span><span class="p">(</span><span class="n">out</span><span class="p">,</span> <span class="n">dim</span><span class="o">=</span><span class="mi">1</span><span class="p">)</span>
<span class="n">test_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>
<span class="c1"># all optional...</span>
<span class="c1"># return whatever you need for the collation function test_end</span>
<span class="n">output</span> <span class="o">=</span> <span class="n">OrderedDict</span><span class="p">({</span>
<span class="s1">&#39;test_loss&#39;</span><span class="p">:</span> <span class="n">loss_test</span><span class="p">,</span>
<span class="s1">&#39;test_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">test_acc</span><span class="p">),</span> <span class="c1"># everything must be a tensor</span>
<span class="p">})</span>
<span class="c1"># return an optional dict</span>
<span class="k">return</span> <span class="n">output</span>
</pre></div>
</td></tr></table>
<p>If you pass in multiple test datasets, test_step will have an additional argument.</p>
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
2
3</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># CASE 2: multiple test datasets</span>
<span class="k">def</span> <span class="nf">test_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>test_dataloader</code>. </p>
<hr />
<h3 id="test_end">test_end</h3>
<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">test_end</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">outputs</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<p>If you didn't define a test_step, this won't be called. </p>
<p>Called at the end of the test step with the output of each test_step. Called once per test dataset. </p>
<p>The outputs here are strictly for the progress bar. If you don't need to display anything, don't return anything. </p>
<p><strong>Params</strong> </p>
<table>
<thead>
<tr>
<th>Param</th>
<th>description</th>
</tr>
</thead>
<tbody>
<tr>
<td>outputs</td>
<td>List of outputs you defined test_step</td>
</tr>
</tbody>
</table>
<p><strong>Return</strong> </p>
<table>
<thead>
<tr>
<th>Return</th>
<th>description</th>
<th>optional</th>
</tr>
</thead>
<tbody>
<tr>
<td>dict</td>
<td>Dict of OrderedDict with metrics to display in progress bar</td>
<td>Y</td>
</tr>
</tbody>
</table>
<p><strong>Example</strong></p>
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre> 1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="k">def</span> <span class="nf">test_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 test to aggregate outputs</span>
<span class="sd"> :param outputs: list of individual outputs of each test step</span>
<span class="sd"> :return:</span>
<span class="sd"> &quot;&quot;&quot;</span>
<span class="n">test_loss_mean</span> <span class="o">=</span> <span class="mi">0</span>
<span class="n">test_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">test_loss_mean</span> <span class="o">+=</span> <span class="n">output</span><span class="p">[</span><span class="s1">&#39;test_loss&#39;</span><span class="p">]</span>
<span class="n">test_acc_mean</span> <span class="o">+=</span> <span class="n">output</span><span class="p">[</span><span class="s1">&#39;test_acc&#39;</span><span class="p">]</span>
<span class="n">test_loss_mean</span> <span class="o">/=</span> <span class="nb">len</span><span class="p">(</span><span class="n">outputs</span><span class="p">)</span>
<span class="n">test_acc_mean</span> <span class="o">/=</span> <span class="nb">len</span><span class="p">(</span><span class="n">outputs</span><span class="p">)</span>
<span class="n">tqdm_dic</span> <span class="o">=</span> <span class="p">{</span><span class="s1">&#39;test_loss&#39;</span><span class="p">:</span> <span class="n">test_loss_mean</span><span class="o">.</span><span class="n">item</span><span class="p">(),</span> <span class="s1">&#39;test_acc&#39;</span><span class="p">:</span> <span class="n">test_acc_mean</span><span class="o">.</span><span class="n">item</span><span class="p">()}</span>
<span class="k">return</span> <span class="n">tqdm_dic</span>
</pre></div>
</td></tr></table>
<hr />
<h3 id="on_save_checkpoint">on_save_checkpoint</h3>
<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>
+2 -2
View File
@@ -156,7 +156,7 @@
</div>
<div class="md-source__repository">
GitHub
williamFalcon/pytorch-lightning
</div>
</a>
</div>
@@ -203,7 +203,7 @@
</div>
<div class="md-source__repository">
GitHub
williamFalcon/pytorch-lightning
</div>
</a>
</div>
+2 -2
View File
@@ -156,7 +156,7 @@
</div>
<div class="md-source__repository">
GitHub
williamFalcon/pytorch-lightning
</div>
</a>
</div>
@@ -203,7 +203,7 @@
</div>
<div class="md-source__repository">
GitHub
williamFalcon/pytorch-lightning
</div>
</a>
</div>
+2 -2
View File
@@ -156,7 +156,7 @@
</div>
<div class="md-source__repository">
GitHub
williamFalcon/pytorch-lightning
</div>
</a>
</div>
@@ -203,7 +203,7 @@
</div>
<div class="md-source__repository">
GitHub
williamFalcon/pytorch-lightning
</div>
</a>
</div>
+2 -2
View File
@@ -156,7 +156,7 @@
</div>
<div class="md-source__repository">
GitHub
williamFalcon/pytorch-lightning
</div>
</a>
</div>
@@ -203,7 +203,7 @@
</div>
<div class="md-source__repository">
GitHub
williamFalcon/pytorch-lightning
</div>
</a>
</div>
+2 -2
View File
@@ -156,7 +156,7 @@
</div>
<div class="md-source__repository">
GitHub
williamFalcon/pytorch-lightning
</div>
</a>
</div>
@@ -203,7 +203,7 @@
</div>
<div class="md-source__repository">
GitHub
williamFalcon/pytorch-lightning
</div>
</a>
</div>
+2 -2
View File
@@ -156,7 +156,7 @@
</div>
<div class="md-source__repository">
GitHub
williamFalcon/pytorch-lightning
</div>
</a>
</div>
@@ -203,7 +203,7 @@
</div>
<div class="md-source__repository">
GitHub
williamFalcon/pytorch-lightning
</div>
</a>
</div>
+2 -2
View File
@@ -156,7 +156,7 @@
</div>
<div class="md-source__repository">
GitHub
williamFalcon/pytorch-lightning
</div>
</a>
</div>
@@ -203,7 +203,7 @@
</div>
<div class="md-source__repository">
GitHub
williamFalcon/pytorch-lightning
</div>
</a>
</div>
+4 -2
View File
@@ -156,7 +156,7 @@
</div>
<div class="md-source__repository">
GitHub
williamFalcon/pytorch-lightning
</div>
</a>
</div>
@@ -203,7 +203,7 @@
</div>
<div class="md-source__repository">
GitHub
williamFalcon/pytorch-lightning
</div>
</a>
</div>
@@ -656,6 +656,8 @@ Lightning will run 5 steps of validation in the beginning of training as a sanit
<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>
<p>You can use <code>Trainer(nb_sanity_val_steps=0)</code> to skip the sanity check.</p>
+2 -2
View File
@@ -156,7 +156,7 @@
</div>
<div class="md-source__repository">
GitHub
williamFalcon/pytorch-lightning
</div>
</a>
</div>
@@ -203,7 +203,7 @@
</div>
<div class="md-source__repository">
GitHub
williamFalcon/pytorch-lightning
</div>
</a>
</div>
+2 -2
View File
@@ -156,7 +156,7 @@
</div>
<div class="md-source__repository">
GitHub
williamFalcon/pytorch-lightning
</div>
</a>
</div>
@@ -203,7 +203,7 @@
</div>
<div class="md-source__repository">
GitHub
williamFalcon/pytorch-lightning
</div>
</a>
</div>
+2 -2
View File
@@ -156,7 +156,7 @@
</div>
<div class="md-source__repository">
GitHub
williamFalcon/pytorch-lightning
</div>
</a>
</div>
@@ -203,7 +203,7 @@
</div>
<div class="md-source__repository">
GitHub
williamFalcon/pytorch-lightning
</div>
</a>
</div>
+48 -48
View File
@@ -156,7 +156,7 @@
</div>
<div class="md-source__repository">
GitHub
williamFalcon/pytorch-lightning
</div>
</a>
</div>
@@ -203,7 +203,7 @@
</div>
<div class="md-source__repository">
GitHub
williamFalcon/pytorch-lightning
</div>
</a>
</div>
@@ -703,58 +703,58 @@ The main function should have 3 arguments: <br />
49
50
51
52</pre></div></td><td class="code"><div class="codehilite"><pre><span></span>def main(hparams, cluster, results_dict):
&quot;&quot;&quot;
Main training routine specific for this project
:param hparams:
:return:
&quot;&quot;&quot;
# init experiment
log_dir = os.path.dirname(os.path.realpath(__file__))
exp = Experiment(
name=&#39;test_tube_exp&#39;,
debug=True,
save_dir=log_dir,
version=0,
autosave=False,
description=&#39;test demo&#39;
)
52</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="k">def</span> <span class="nf">main</span><span class="p">(</span><span class="n">hparams</span><span class="p">,</span> <span class="n">cluster</span><span class="p">,</span> <span class="n">results_dict</span><span class="p">):</span>
<span class="sd">&quot;&quot;&quot;</span>
<span class="sd"> Main training routine specific for this project</span>
<span class="sd"> :param hparams:</span>
<span class="sd"> :return:</span>
<span class="sd"> &quot;&quot;&quot;</span>
<span class="c1"># init experiment</span>
<span class="n">log_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">dirname</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">realpath</span><span class="p">(</span><span class="vm">__file__</span><span class="p">))</span>
<span class="n">exp</span> <span class="o">=</span> <span class="n">Experiment</span><span class="p">(</span>
<span class="n">name</span><span class="o">=</span><span class="s1">&#39;test_tube_exp&#39;</span><span class="p">,</span>
<span class="n">debug</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span>
<span class="n">save_dir</span><span class="o">=</span><span class="n">log_dir</span><span class="p">,</span>
<span class="n">version</span><span class="o">=</span><span class="mi">0</span><span class="p">,</span>
<span class="n">autosave</span><span class="o">=</span><span class="bp">False</span><span class="p">,</span>
<span class="n">description</span><span class="o">=</span><span class="s1">&#39;test demo&#39;</span>
<span class="p">)</span>
# set the hparams for the experiment
exp.argparse(hparams)
exp.save()
<span class="c1"># set the hparams for the experiment</span>
<span class="n">exp</span><span class="o">.</span><span class="n">argparse</span><span class="p">(</span><span class="n">hparams</span><span class="p">)</span>
<span class="n">exp</span><span class="o">.</span><span class="n">save</span><span class="p">()</span>
# build model
model = MyLightningModule(hparams)
<span class="c1"># build model</span>
<span class="n">model</span> <span class="o">=</span> <span class="n">MyLightningModule</span><span class="p">(</span><span class="n">hparams</span><span class="p">)</span>
# callbacks
early_stop = EarlyStopping(
monitor=hparams.early_stop_metric,
patience=hparams.early_stop_patience,
verbose=True,
mode=hparams.early_stop_mode
)
<span class="c1"># callbacks</span>
<span class="n">early_stop</span> <span class="o">=</span> <span class="n">EarlyStopping</span><span class="p">(</span>
<span class="n">monitor</span><span class="o">=</span><span class="n">hparams</span><span class="o">.</span><span class="n">early_stop_metric</span><span class="p">,</span>
<span class="n">patience</span><span class="o">=</span><span class="n">hparams</span><span class="o">.</span><span class="n">early_stop_patience</span><span class="p">,</span>
<span class="n">verbose</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span>
<span class="n">mode</span><span class="o">=</span><span class="n">hparams</span><span class="o">.</span><span class="n">early_stop_mode</span>
<span class="p">)</span>
model_save_path = &#39;{}/{}/{}&#39;.format(hparams.model_save_path, exp.name, exp.version)
checkpoint = ModelCheckpoint(
filepath=model_save_path,
save_function=None,
save_best_only=True,
verbose=True,
monitor=hparams.model_save_monitor_value,
mode=hparams.model_save_monitor_mode
)
<span class="n">model_save_path</span> <span class="o">=</span> <span class="s1">&#39;{}/{}/{}&#39;</span><span class="o">.</span><span class="n">format</span><span class="p">(</span><span class="n">hparams</span><span class="o">.</span><span class="n">model_save_path</span><span class="p">,</span> <span class="n">exp</span><span class="o">.</span><span class="n">name</span><span class="p">,</span> <span class="n">exp</span><span class="o">.</span><span class="n">version</span><span class="p">)</span>
<span class="n">checkpoint</span> <span class="o">=</span> <span class="n">ModelCheckpoint</span><span class="p">(</span>
<span class="n">filepath</span><span class="o">=</span><span class="n">model_save_path</span><span class="p">,</span>
<span class="n">save_function</span><span class="o">=</span><span class="bp">None</span><span class="p">,</span>
<span class="n">save_best_only</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span>
<span class="n">verbose</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span>
<span class="n">monitor</span><span class="o">=</span><span class="n">hparams</span><span class="o">.</span><span class="n">model_save_monitor_value</span><span class="p">,</span>
<span class="n">mode</span><span class="o">=</span><span class="n">hparams</span><span class="o">.</span><span class="n">model_save_monitor_mode</span>
<span class="p">)</span>
# configure trainer
trainer = Trainer(
experiment=exp,
cluster=cluster,
checkpoint_callback=checkpoint,
early_stop_callback=early_stop,
)
<span class="c1"># configure trainer</span>
<span class="n">trainer</span> <span class="o">=</span> <span class="n">Trainer</span><span class="p">(</span>
<span class="n">experiment</span><span class="o">=</span><span class="n">exp</span><span class="p">,</span>
<span class="n">cluster</span><span class="o">=</span><span class="n">cluster</span><span class="p">,</span>
<span class="n">checkpoint_callback</span><span class="o">=</span><span class="n">checkpoint</span><span class="p">,</span>
<span class="n">early_stop_callback</span><span class="o">=</span><span class="n">early_stop</span><span class="p">,</span>
<span class="p">)</span>
# train model
trainer.fit(model)
<span class="c1"># train model</span>
<span class="n">trainer</span><span class="o">.</span><span class="n">fit</span><span class="p">(</span><span class="n">model</span><span class="p">)</span>
</pre></div>
</td></tr></table>
+2 -2
View File
@@ -156,7 +156,7 @@
</div>
<div class="md-source__repository">
GitHub
williamFalcon/pytorch-lightning
</div>
</a>
</div>
@@ -203,7 +203,7 @@
</div>
<div class="md-source__repository">
GitHub
williamFalcon/pytorch-lightning
</div>
</a>
</div>
File diff suppressed because one or more lines are too long
+14 -14
View File
@@ -2,72 +2,72 @@
<urlset xmlns="http://www.sitemaps.org/schemas/sitemap/0.9">
<url>
<loc>None</loc>
<lastmod>2019-08-24</lastmod>
<lastmod>2019-08-31</lastmod>
<changefreq>daily</changefreq>
</url>
<url>
<loc>None</loc>
<lastmod>2019-08-24</lastmod>
<lastmod>2019-08-31</lastmod>
<changefreq>daily</changefreq>
</url>
<url>
<loc>None</loc>
<lastmod>2019-08-24</lastmod>
<lastmod>2019-08-31</lastmod>
<changefreq>daily</changefreq>
</url>
<url>
<loc>None</loc>
<lastmod>2019-08-24</lastmod>
<lastmod>2019-08-31</lastmod>
<changefreq>daily</changefreq>
</url>
<url>
<loc>None</loc>
<lastmod>2019-08-24</lastmod>
<lastmod>2019-08-31</lastmod>
<changefreq>daily</changefreq>
</url>
<url>
<loc>None</loc>
<lastmod>2019-08-24</lastmod>
<lastmod>2019-08-31</lastmod>
<changefreq>daily</changefreq>
</url>
<url>
<loc>None</loc>
<lastmod>2019-08-24</lastmod>
<lastmod>2019-08-31</lastmod>
<changefreq>daily</changefreq>
</url>
<url>
<loc>None</loc>
<lastmod>2019-08-24</lastmod>
<lastmod>2019-08-31</lastmod>
<changefreq>daily</changefreq>
</url>
<url>
<loc>None</loc>
<lastmod>2019-08-24</lastmod>
<lastmod>2019-08-31</lastmod>
<changefreq>daily</changefreq>
</url>
<url>
<loc>None</loc>
<lastmod>2019-08-24</lastmod>
<lastmod>2019-08-31</lastmod>
<changefreq>daily</changefreq>
</url>
<url>
<loc>None</loc>
<lastmod>2019-08-24</lastmod>
<lastmod>2019-08-31</lastmod>
<changefreq>daily</changefreq>
</url>
<url>
<loc>None</loc>
<lastmod>2019-08-24</lastmod>
<lastmod>2019-08-31</lastmod>
<changefreq>daily</changefreq>
</url>
<url>
<loc>None</loc>
<lastmod>2019-08-24</lastmod>
<lastmod>2019-08-31</lastmod>
<changefreq>daily</changefreq>
</url>
<url>
<loc>None</loc>
<lastmod>2019-08-24</lastmod>
<lastmod>2019-08-31</lastmod>
<changefreq>daily</changefreq>
</url>
</urlset>
BIN
View File
Binary file not shown.