【PL 基础】如何训练一个模型
- 1. 如何训练一个模型
- 2. 验证和测试模型
1. 如何训练一个模型
import os
import torch
from torch import nn
import torch.nn.functional as F
from torchvision import transforms
from torchvision.datasets import MNIST
from torch.utils.data import DataLoader
import lightning as L# 定义 PyTorch nn.模块
class Encoder(nn.Module):def __init__(self):super().__init__()self.l1 = nn.Sequential(nn.Linear(28 * 28, 64), nn.ReLU(), nn.Linear(64, 3))def forward(self, x):return self.l1(x)class Decoder(nn.Module):def __init__(self):super().__init__()self.l1 = nn.Sequential(nn.Linear(3, 64), nn.ReLU(), nn.Linear(64, 28 * 28))def forward(self, x):return self.l1(x)# 定义 LightningModule
class LitAutoEncoder(L.LightningModule):def __init__(self, encoder, decoder):super().__init__()self.encoder = encoderself.decoder = decoderdef training_step(self, batch, batch_idx):# training_step defines the train loop.x, _ = batchx = x.view(x.size(0), -1)z = self.encoder(x)x_hat = self.decoder(z)loss = F.mse_loss(x_hat, x)return lossdef configure_optimizers(self):# 在 configure_optimizers 中,为模型定义优化器。optimizer = torch.optim.Adam(self.parameters(), lr=1e-3)return optimizer# 定义训练数据集
dataset = MNIST(os.getcwd(), download=True, transform=transforms.ToTensor())
train_loader = DataLoader(dataset)# model
autoencoder = LitAutoEncoder(Encoder(), Decoder())# train model
trainer = L.Trainer()
trainer.fit(model=autoencoder, train_dataloaders=train_loader)
在代码背后,Lightning Trainer 相当于运行以下训练循环:
autoencoder = LitAutoEncoder(Encoder(), Decoder())
optimizer = autoencoder.configure_optimizers()for batch_idx, batch in enumerate(train_loader):loss = autoencoder.training_step(batch, batch_idx)loss.backward()optimizer.step()optimizer.zero_grad()
2. 验证和测试模型
为了确保模型可以泛化到看不见的数据集(即:发布论文或在生产环境中),数据集通常分为两部分,训练拆分和测试拆分。
测试集在训练期间不使用,只有在训练模型后才使用,以查看模型在现实世界中的表现。
import torch.utils.data as data
from torchvision import datasets
import torchvision.transforms as transforms# Load data sets
transform = transforms.ToTensor()
train_set = datasets.MNIST(root="MNIST", download=True, train=True, transform=transform)
test_set = datasets.MNIST(root="MNIST", download=True, train=False, transform=transform)
- 要添加测试循环,请实现
LightningModule的test_step方法
class LitAutoEncoder(L.LightningModule):def training_step(self, batch, batch_idx):...def test_step(self, batch, batch_idx):# this is the test loopx, _ = batchx = x.view(x.size(0), -1)z = self.encoder(x)x_hat = self.decoder(z)test_loss = F.mse_loss(x_hat, x)self.log("test_loss", test_loss)
使用测试循环进行训练,模型完成训练后,调用 .test
from torch.utils.data import DataLoader# initialize the Trainer
trainer = Trainer()# test the model
trainer.test(model, dataloaders=DataLoader(test_set))
- 添加验证循环。
在训练期间,通常的做法是使用训练分割的一小部分来确定模型何时完成训练。
# use 20% of training data for validation
train_set_size = int(len(train_set) * 0.8)
valid_set_size = len(train_set) - train_set_size# split the train set into two
seed = torch.Generator().manual_seed(42)
train_set, valid_set = data.random_split(train_set, [train_set_size, valid_set_size], generator=seed)
要添加验证循环,请实现 LightningModule 的 validation_step 方法
class LitAutoEncoder(L.LightningModule):def training_step(self, batch, batch_idx):...def validation_step(self, batch, batch_idx):# this is the validation loopx, _ = batchx = x.view(x.size(0), -1)z = self.encoder(x)x_hat = self.decoder(z)val_loss = F.mse_loss(x_hat, x)self.log("val_loss", val_loss)
要运行验证循环,请将验证集传入 .fit
from torch.utils.data import DataLoadertrain_loader = DataLoader(train_set)
valid_loader = DataLoader(valid_set)
model = LitAutoEncoder(...)# train with both splits
trainer = L.Trainer()
trainer.fit(model, train_loader, valid_loader)