用 LightningDataModule 管理数据

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

Sergiy Tkachuk

Director, GenAI Productivity

模型训练的数据准备

  • 数据准备不当会导致训练问题
    • 训练变慢
    • 频繁中断
    • 无法收敛

训练数据准备.png

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

为何使用 LightningDataModule?

$$

  • 📂 集中管理数据集

$$

  • 📊 标准化数据准备流程

$$

  • 🚀 简化训练与评估阶段
使用 PyTorch Lightning 构建可扩展 AI 模型

用 LightningDataModule 管理数据

关键方法:

  • prepare_data:下载并设置数据
  • setup:将数据划分为训练、验证、测试集
class ImageDataModule(pl.LightningDataModule):
    def __init__(self, data_dir="./data", batch_size=32):
        super().__init__()
        ...

def prepare_data(self): datasets.MNIST(self.data_dir, train=True, download=True)
def setup(self, stage=None): dataset = datasets.MNIST(self.data_dir, train=True, transform=self.transform) self.train_data, self.val_data = random_split(dataset, [55000, 5000]) self.test_data = datasets.MNIST(self.data_dir, train=False, transform=self.transform)
使用 PyTorch Lightning 构建可扩展 AI 模型

创建训练 DataLoader

$$

  • 提供训练数据批次
  • 帮助优化 GPU 利用率
  • 高效遍历大型数据集
def train_dataloader(self):
    return DataLoader(self.train_data, batch_size=self.batch_size, shuffle=True)
使用 PyTorch Lightning 构建可扩展 AI 模型

创建验证 DataLoader

$$

  • 提供模型验证数据
  • 帮助监控泛化性能
  • 通过打乱确保评估一致性
def val_dataloader(self):
    return DataLoader(self.val_data, batch_size=self.batch_size)
使用 PyTorch Lightning 构建可扩展 AI 模型

创建测试 DataLoader

$$

  • 为训练结束后的最终评估提供数据
  • 模拟真实场景下的性能评估
  • 确保无偏的性能测量
def test_dataloader(self):
    return DataLoader(self.test_data, batch_size=self.batch_size)
使用 PyTorch Lightning 构建可扩展 AI 模型

将 DataModule 连接到 LightningModule

  • 模块化设计将数据与模型逻辑分离

PyTorch Lightning 示意图

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

将 DataModule 连接到 LightningModule

  • 模块化设计分离数据与模型逻辑
  • LightningDataModuleLightningModule 配合使用
  • 标准化流程提升可复现性

包含 DataModule 与 LightningModule 的 PyTorch Lightning 示意图

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

Vamos praticar!

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

Preparing Video For Download...