mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-10 12:21:57 +08:00
Deployed 53ec3bc with MkDocs version: 1.0.4
This commit is contained in:
@@ -1013,7 +1013,7 @@ class CoolModel(pl.LightningModule):
|
||||
y_hat = self.forward(x)
|
||||
return {'loss': F.cross_entropy(y_hat, y)(y_hat, y)}
|
||||
|
||||
def validation_step(self, batch, batch_nb, dataloader_i):
|
||||
def validation_step(self, batch, batch_nb):
|
||||
# OPTIONAL
|
||||
x, y = batch
|
||||
y_hat = self.forward(x)
|
||||
@@ -1191,7 +1191,7 @@ If you don't need to validate you don't need to implement this method. </p>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>dataloader_i</td>
|
||||
<td>Integer displaying which dataloader this is</td>
|
||||
<td>Integer displaying which dataloader this is (only if multiple val datasets used)</td>
|
||||
</tr>
|
||||
</tbody>
|
||||
</table>
|
||||
@@ -1213,7 +1213,8 @@ If you don't need to validate you don't need to implement this method. </p>
|
||||
</tbody>
|
||||
</table>
|
||||
<p><strong>Example</strong></p>
|
||||
<pre><code class="python">def validation_step(self, data_batch, batch_nb):
|
||||
<pre><code class="python"># CASE 1: A single validation dataset
|
||||
def validation_step(self, data_batch, batch_nb):
|
||||
x, y, z = data_batch
|
||||
|
||||
# implement your own
|
||||
@@ -1235,6 +1236,13 @@ If you don't need to validate you don't need to implement this method. </p>
|
||||
return output
|
||||
</code></pre>
|
||||
|
||||
<p>If you pass in multiple validation datasets, validation_step will have an additional argument.</p>
|
||||
<pre><code class="python"># CASE 2: multiple validation datasets
|
||||
def validation_step(self, data_batch, batch_nb, dataset_idx):
|
||||
# dataset_idx tells you which dataset this is.
|
||||
</code></pre>
|
||||
|
||||
<p>The <code>dataset_idx</code> corresponds to the order of datasets returned in <code>val_dataloader</code>. </p>
|
||||
<hr />
|
||||
<h3 id="validation_end">validation_end</h3>
|
||||
<pre><code class="python">def validation_end(self, outputs)
|
||||
@@ -1356,6 +1364,8 @@ def val_dataloader(self):
|
||||
return [loader_a, loader_b, ..., loader_n]
|
||||
</code></pre>
|
||||
|
||||
<p>In the case where you return multiple val_dataloaders, the validation_step will have an arguement <code>dataset_idx</code>
|
||||
which matches the order here. </p>
|
||||
<hr />
|
||||
<h3 id="test_dataloader">test_dataloader</h3>
|
||||
<pre><code class="python">@pl.data_loader
|
||||
|
||||
@@ -478,6 +478,13 @@
|
||||
on_tng_metrics
|
||||
</a>
|
||||
|
||||
</li>
|
||||
|
||||
<li class="md-nav__item">
|
||||
<a href="#optimizer_step" title="optimizer_step" class="md-nav__link">
|
||||
optimizer_step
|
||||
</a>
|
||||
|
||||
</li>
|
||||
|
||||
<li class="md-nav__item">
|
||||
@@ -613,6 +620,13 @@
|
||||
on_tng_metrics
|
||||
</a>
|
||||
|
||||
</li>
|
||||
|
||||
<li class="md-nav__item">
|
||||
<a href="#optimizer_step" title="optimizer_step" class="md-nav__link">
|
||||
optimizer_step
|
||||
</a>
|
||||
|
||||
</li>
|
||||
|
||||
<li class="md-nav__item">
|
||||
@@ -707,6 +721,31 @@ this callback to modify what will be logged.</p>
|
||||
# do something before validation end
|
||||
</code></pre>
|
||||
|
||||
<hr />
|
||||
<h4 id="optimizer_step">optimizer_step</h4>
|
||||
<p>Calls .step() and .zero_grad for each optimizer.<br />
|
||||
You can override this method to adjust how you do the optimizer step for each optimizer</p>
|
||||
<p>Called once per optimizer</p>
|
||||
<pre><code class="python"># DEFAULT
|
||||
def optimizer_step(self, current_epoch, batch_nb, optimizer, optimizer_i):
|
||||
optimizer.step()
|
||||
optimizer.zero_grad()
|
||||
|
||||
# Alternating schedule for optimizer steps (ie: GANs)
|
||||
def optimizer_step(self, current_epoch, batch_nb, optimizer, optimizer_i):
|
||||
# update generator opt every 2 steps
|
||||
if optimizer_i == 0:
|
||||
if batch_nb % 2 == 0 :
|
||||
optimizer.step()
|
||||
optimizer.zero_grad()
|
||||
|
||||
# update discriminator opt every 4 steps
|
||||
if optimizer_i == 1:
|
||||
if batch_nb % 4 == 0 :
|
||||
optimizer.step()
|
||||
optimizer.zero_grad()
|
||||
</code></pre>
|
||||
|
||||
<hr />
|
||||
<h4 id="on_before_zero_grad">on_before_zero_grad</h4>
|
||||
<p>Called in the training loop after taking an optimizer step and before zeroing grads.
|
||||
|
||||
File diff suppressed because one or more lines are too long
Binary file not shown.
Reference in New Issue
Block a user