Doskonalenie metody forward
Po zdefiniowaniu warstw w metodzie __init__ to metoda forward decyduje o tym, jak dane przez nie przepływają. W PyTorch Lightning taki podział sprawia, że kod jest przejrzysty i łatwy w utrzymaniu. Wiesz już, jak zbudować konstruktor – czas skupić się na metodzie forward i zadbać o to, by logika klasyfikacji była czytelna i gotowa do trenowania. Warstwy w __init__ są już dla ciebie zdefiniowane, więc możesz skoncentrować się wyłącznie na przepływie danych.
Biblioteki lightning.pytorch i torch.nn zostały już zaimportowane odpowiednio jako pl i nn.
To ćwiczenie jest częścią kursu
Skalowalne modele AI z PyTorch Lightning
Instrukcje do ćwiczenia
- Zaimplementuj metodę
forwardwewnątrz klasyClassifierModel. - Zastosuj aktywację ReLU po warstwie ukrytej.
Interaktywne ćwiczenie praktyczne
Spróbuj tego ćwiczenia, uzupełniając ten przykładowy kod.
class ClassifierModel(pl.LightningModule):
def __init__(self, input_dim, hidden_dim, output_dim):
super().__init__()
self.hidden = nn.Linear(input_dim, hidden_dim)
self.output = nn.Linear(hidden_dim, output_dim)
# Define forward method
def ____(self, ____):
# Complete the forward pass
x = self.hidden(x)
x = ____(x)
x = self.output(x)
return x