開始使用免費開始

掌握 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)
編輯並執行程式碼