mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-10 12:21:57 +08:00
Docs (#1024)
* added checkpoint defaults * added checkpoint defaults * added checkpoint defaults * added checkpoint defaults * added checkpoint defaults * added checkpoint defaults * added checkpoint defaults * added checkpoint defaults * added checkpoint defaults * added checkpoint defaults * added checkpoint defaults * docs * docs * docs * docs * docs * docs * docs * docs * docs * docs * docs * docs * docs * docs * docs * added community examples * added community examples
This commit is contained in:
@@ -0,0 +1,62 @@
|
||||
Child Modules
|
||||
-------------
|
||||
Research projects tend to test different approaches to the same dataset.
|
||||
This is very easy to do in Lightning with inheritance.
|
||||
|
||||
For example, imaging we now want to train an Autoencoder to use as a feature extractor for MNIST images.
|
||||
Recall that `CoolMNIST` already defines all the dataloading etc... The only things
|
||||
that change in the `Autoencoder` model are the init, forward, training, validation and test step.
|
||||
|
||||
.. code-block::
|
||||
|
||||
class Encoder(torch.nn.Module):
|
||||
...
|
||||
|
||||
class AutoEncoder(CoolMNIST):
|
||||
def __init__(self):
|
||||
self.encoder = Encoder()
|
||||
self.decoder = Decoder()
|
||||
|
||||
def forward(self, x):
|
||||
generated = self.decoder(x)
|
||||
|
||||
def training_step(self, batch, batch_idx):
|
||||
x, _ = batch
|
||||
|
||||
representation = self.encoder(x)
|
||||
x_hat = self.forward(representation)
|
||||
|
||||
loss = MSE(x, x_hat)
|
||||
return loss
|
||||
|
||||
def validation_step(self, batch, batch_idx):
|
||||
return self._shared_eval(batch, batch_idx, 'val'):
|
||||
|
||||
def test_step(self, batch, batch_idx):
|
||||
return self._shared_eval(batch, batch_idx, 'test'):
|
||||
|
||||
def _shared_eval(self, batch, batch_idx, prefix):
|
||||
x, y = batch
|
||||
representation = self.encoder(x)
|
||||
x_hat = self.forward(representation)
|
||||
|
||||
loss = F.nll_loss(logits, y)
|
||||
return {f'{prefix}_loss': loss}
|
||||
|
||||
and we can train this using the same trainer
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
autoencoder = AutoEncoder()
|
||||
trainer = Trainer()
|
||||
trainer.fit(autoencoder)
|
||||
|
||||
And remember that the forward method is to define the practical use of a LightningModule.
|
||||
In this case, we want to use the `AutoEncoder` to extract image representations
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
some_images = torch.Tensor(32, 1, 28, 28)
|
||||
representations = autoencoder(some_images)
|
||||
|
||||
..
|
||||
Reference in New Issue
Block a user