You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.23 20:13:10