mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-11 12:31:23 +08:00
Deployed f51b459 with MkDocs version: 1.0.4
This commit is contained in:
@@ -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">'val_loss'</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">'avg_val_loss'</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">'test_loss'</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">'test_loss'</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">'avg_test_loss'</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">'test_loss'</span><span class="p">:</span> <span class="n">loss_test</span><span class="p">,</span>
|
||||
<span class="s1">'test_acc'</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">"""</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"> """</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">'test_loss'</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">'test_acc'</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">'test_loss'</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">'test_acc'</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>
|
||||
|
||||
@@ -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>
|
||||
|
||||
@@ -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>
|
||||
|
||||
@@ -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>
|
||||
|
||||
@@ -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>
|
||||
|
||||
@@ -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>
|
||||
|
||||
@@ -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>
|
||||
|
||||
@@ -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>
|
||||
|
||||
@@ -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>
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -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>
|
||||
|
||||
@@ -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
@@ -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>
|
||||
|
||||
@@ -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):
|
||||
"""
|
||||
Main training routine specific for this project
|
||||
:param hparams:
|
||||
:return:
|
||||
"""
|
||||
# init experiment
|
||||
log_dir = os.path.dirname(os.path.realpath(__file__))
|
||||
exp = Experiment(
|
||||
name='test_tube_exp',
|
||||
debug=True,
|
||||
save_dir=log_dir,
|
||||
version=0,
|
||||
autosave=False,
|
||||
description='test demo'
|
||||
)
|
||||
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">"""</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"> """</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">'test_tube_exp'</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">'test demo'</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 = '{}/{}/{}'.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">'{}/{}/{}'</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
@@ -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
@@ -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>
|
||||
Binary file not shown.
Reference in New Issue
Block a user