掌握 init 方法
在 PyTorch Lightning 中,__init__ 方法是設定 LightningModule 的核心。你會在這裡定義模型層、參數,以及訓練前的初始設定。透過把這個設定步驟清楚分離,PyTorch Lightning 能讓專案更容易維護與擴充。在本練習中,你將專注於為分類模型初始化類別。
本練習屬於課程
Scalable AI Models with PyTorch Lightning
練習說明
- 建立一個名為
ClassifierModel的類別,並從pl.LightningModule繼承。 - 初始化父類別,以使用 LightningModule 的功能。
動手互動練習
試著完成這個範例程式碼,體驗一下這個練習。
import lightning.pytorch as pl
import torch.nn as nn
# Create the class
class ClassifierModel(____):
# Create init method
def __init__(self, input_dim, output_dim):
____
self.classifier = nn.Linear(input_dim, output_dim)