使用 LightningModule 定义模型

使用 PyTorch Lightning 构建可扩展 AI 模型

Sergiy Tkachuk

Director, GenAI Productivity

聚焦 LightningModule

  1. 封装模型架构
  2. 将训练逻辑组织为单一、可管理的单元
  3. 为深度学习项目提供清晰有序的蓝图

PyTorch LightningModule diagram

使用 PyTorch Lightning 构建可扩展 AI 模型

定义 init 方法

关键任务:

  • 模型初始化
  • super()
    • 训练循环自动处理
    • 日志记录
    • 检查点保存
  • 初始化后定义模型层
  • 模块化,易维护
import lightning.pytorch as pl
import torch.nn as nn

class ClassificationModel(pl.LightningModule):
    def __init__(self, input_dim,
                 hidden_dim, num_class):
          # Initialize parent class
        super().__init__()

# First layer self.layer1 = nn.Linear(input_dim, hidden_dim) # Activation function self.relu = nn.ReLU() # Output layer self.layer2 = nn.Linear(hidden_dim, num_class)
使用 PyTorch Lightning 构建可扩展 AI 模型

实现 forward 方法

关键步骤:

  • 定义网络中的数据流
  • 顺序处理输入
    • 线性变换
    • 激活
    • 最后一层与输出
import lightning.pytorch as pl
import torch.nn as nn

class ClassificationModel(pl.LightningModule):
    def __init__(self, input_dim,
                 hidden_dim, num_class):
          ...

def forward(self, x):
x = self.layer1(x) # Pass input
x = nn.ReLU(x) # Apply activation
x = self.layer2(x) # Compute output
return x # Return result
使用 PyTorch Lightning 构建可扩展 AI 模型

示例:手写数字分类

import lightning.pytorch as pl
from torch.utils.data import DataLoader
from torchvision.datasets import MNIST
from torchvision import transforms

transform = transforms.ToTensor() train_ds = MNIST(root='.', train=True, download=True, transform=transform) test_ds = MNIST(root='.', train=False, download=True, transform=transform) train_loader = DataLoader(train_ds, batch_size=64, shuffle=True) test_loader = DataLoader(test_ds, batch_size=64)
model = ClassificationModel(input_dim=28*28, hidden_dim=128, num_class=10)
trainer = pl.Trainer(max_epochs=3, accelerator='auto') trainer.fit(model, train_loader, test_loader)
使用 PyTorch Lightning 构建可扩展 AI 模型

将模型用于分类任务

$$

  • 聚焦分类用例
  • 全流程在 LightningModule 内完成
  • 输出供 softmax 使用的原始值
  • 与 Lightning Trainer 集成
class ClassificationModel(pl.LightningModule):
  def __init__(self, input_dim, 
               hidden_dim, output_dim):
    super().__init__()

self.hid = nn.Linear(input_dim, hidden_dim) self.out = nn.Linear(hidden_dim, output_dim)
def forward(self, x):
x = self.hidden(x) x = nn.ReLU(x) x = self.output(x)
return x
使用 PyTorch Lightning 构建可扩展 AI 模型

Passons à la pratique !

使用 PyTorch Lightning 构建可扩展 AI 模型

Preparing Video For Download...