使用 Lightning 回调优化训练

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

Sergiy Tkachuk

Director, GenAI Productivity

什么是回调?

$$

  • 在训练关键阶段执行的函数
  • 添加自定义操作而不使代码混乱
  • 提升灵活性与可控性

回调概览

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

什么是回调?

$$

from lightning.pytorch.callbacks import Callback

class MyPrintingCallback(Callback):

def on_train_start(self, trainer, pl_module): print("Training is starting")
def on_train_end(self, trainer, pl_module): print("Training is ending")
  • 在训练各阶段添加自定义操作
使用 PyTorch Lightning 构建可扩展 AI 模型

Lightning 的 ModelCheckpoint 回调

$$

  • 按设定间隔自动保存模型

  • 自选监控指标

  • 仅保留最佳模型

from lightning.pytorch.callbacks
import ModelCheckpoint

checkpoint_callback = ModelCheckpoint(
monitor='val_loss',
dirpath='my/path/',
filename='{epoch}-{val_loss:.2f}',
save_top_k=1,
mode='min'
)
1 https://lightning.ai/docs/pytorch/stable/api/lightning.pytorch.callbacks.ModelCheckpoint.html
使用 PyTorch Lightning 构建可扩展 AI 模型

Lightning 的 EarlyStopping 回调

$$

  • 监控某个指标
  • 当指标不再提升时停止训练
from lightning.pytorch.callbacks
import EarlyStopping

early_stopping_callback = EarlyStopping(

monitor='val_loss',
patience=3,
mode='min'
)
1 https://lightning.ai/docs/pytorch/stable/api/lightning.pytorch.callbacks.EarlyStopping.html
使用 PyTorch Lightning 构建可扩展 AI 模型

自定义并使用 Lightning 回调

from lightning.pytorch import Trainer
from lightning.pytorch.callbacks import EarlyStopping, ModelCheckpoint

checkpoint = ModelCheckpoint( monitor='val_accuracy', save_top_k=2, mode='max')
early_stopping = EarlyStopping( monitor='val_accuracy', patience=5, mode='max')
trainer = Trainer(max_epochs=50, callbacks=[checkpoint, early_stopping])
1 https://lightning.ai/docs/pytorch/stable/common/trainer.html
使用 PyTorch Lightning 构建可扩展 AI 模型

Passons à la pratique !

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

Preparing Video For Download...