PyTorch+CIFAR100早停机制实现:代码整合位置咨询
如何在PyTorch CIFAR100训练代码中整合早停机制
需求背景
我是PyTorch新手,基于CIFAR100数据集开发模型,希望在现有训练代码中加入早停机制。目前没有单独划分验证集,计划用测试集临时作为验证集使用。现有训练函数、早停类代码如下,想了解如何将早停调用逻辑整合进训练流程。
现有代码
训练函数
def train(net,trainloader,epochs,use_gpu = True): ... net.train() for epoch in range(epochs): print ("Epoch {}/{}".format(epoch+1, epochs)) running_loss = 0.0 running_corrects = 0 for i, data in enumerate(trainloader, 0): images, labels = data[0].to(device), data[1].to(device) optimizer.zero_grad() outputs = net(images) _, preds = torch.max(outputs, 1) loss = criterion(outputs, labels) loss.backward() optimizer.step() running_loss += loss.item() epoch_loss = running_loss/len(trainloader.dataset) print('Loss: {}'.format(epoch_loss))
早停类
class EarlyStopper: def __init__(self, patience=1, min_delta=0): self.patience = patience self.min_delta = min_delta self.counter = 0 self.min_validation_loss = float('inf') def early_stop(self, validation_loss): if validation_loss < self.min_validation_loss: self.min_validation_loss = validation_loss self.counter = 0 elif validation_loss > (self.min_validation_loss + self.min_delta): self.counter += 1 if self.counter >= self.patience: return True return False
整合步骤与完整代码
要实现早停,需要先添加验证函数,再重构训练逻辑,让每轮训练后都能评估验证损失并触发早停判断。
1. 编写验证函数
这个函数负责在测试集(当作验证集)上计算损失,注意要关闭梯度计算,避免影响训练:
def validate(net, valloader, criterion, device): net.eval() # 切换到评估模式 val_loss = 0.0 with torch.no_grad(): # 关闭梯度计算 for images, labels in valloader: images, labels = images.to(device), labels.to(device) outputs = net(images) loss = criterion(outputs, labels) val_loss += loss.item() # 计算平均验证损失 avg_val_loss = val_loss / len(valloader.dataset) net.train() # 切回训练模式 return avg_val_loss
2. 重构训练逻辑,整合早停
把原来的训练循环拆出来,每轮训练后调用验证函数,再用早停类判断是否停止训练:
import torch import numpy as np from torchvision.datasets import CIFAR100 from torch.utils.data import DataLoader import torchvision.transforms as transforms # 初始化CIFAR100数据集(训练集+测试集) transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,0.5,0.5), (0.5,0.5,0.5))]) trainset = CIFAR100(root='./data', train=True, download=True, transform=transform) trainloader = DataLoader(trainset, batch_size=64, shuffle=True, num_workers=2) testset = CIFAR100(root='./data', train=False, download=True, transform=transform) valloader = DataLoader(testset, batch_size=64, shuffle=False, num_workers=2) # 用测试集当验证集 # 假设你已经定义了模型、optimizer、criterion、device # 示例:device = torch.device("cuda" if torch.cuda.is_available() else "cpu") # 初始化早停器 early_stopper = EarlyStopper(patience=3, min_delta=0.01) # 注意min_delta别设太大,原代码的10不合理,这里改小 # 主训练循环 n_epochs = 50 # 预设最大训练轮数 for epoch in range(n_epochs): print(f"Epoch {epoch+1}/{n_epochs}") running_loss = 0.0 # 单轮训练 net.train() for images, labels in trainloader: images, labels = images.to(device), labels.to(device) optimizer.zero_grad() outputs = net(images) _, preds = torch.max(outputs, 1) loss = criterion(outputs, labels) loss.backward() optimizer.step() running_loss += loss.item() # 计算训练损失 train_loss = running_loss / len(trainloader.dataset) print(f'Train Loss: {train_loss:.4f}') # 计算验证损失 val_loss = validate(net, valloader, criterion, device) print(f'Validation Loss: {val_loss:.4f}') # 检查是否早停 if early_stopper.early_stop(val_loss): print("早停触发,停止训练") break
关键注意事项
- 验证模式切换:验证时必须调用
net.eval(),评估完要切回net.train(),否则模型的BatchNorm、Dropout等层会工作异常。 - 梯度关闭:用
torch.no_grad()包裹验证过程,大幅减少显存占用,提升速度。 - 参数调整:原代码中
min_delta=10不合理,因为CIFAR100的损失通常在0-10之间,建议设为0.01或0.001,避免误触发早停。 - 测试集使用:用测试集当验证集只是临时方案,后续最好从训练集中划分出单独的验证集,避免最终模型评估时数据泄露。
内容的提问来源于stack exchange,提问作者seni
相关产品推荐
相关产品推荐

