Perfecționarea metodei forward
După ce ai definit straturile în metoda __init__, metoda forward stabilește modul în care datele circulă prin ele. În PyTorch Lightning, această separare menține codul curat și ușor de întreținut. Ai văzut deja cum se structurează constructorul – acum e momentul să te concentrezi pe pasul forward, asigurându-te că logica de clasificare este clară și optimizată pentru antrenament. Straturile din __init__ sunt deja definite pentru tine, astfel încât te poți concentra exclusiv pe fluxul forward.
lightning.pytorch și torch.nn au fost deja importate ca pl, respectiv nn.
Acest exercițiu face parte din cursul
Modele AI scalabile cu PyTorch Lightning
Instrucțiuni pentru exercițiu
- Implementează metoda
forwardîn interiorul claseiClassifierModel. - Aplică o activare ReLU după stratul ascuns.
Exercițiu interactiv practic
Încearcă acest exercițiu completând acest cod de exemplu.
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