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

如何在PyTorch中手动计算整个数据集的误差?

在PyTorch 0.3.1中计算CIFAR10数据集的整体误差

看起来你已经搭好了基础框架,不过error_criterion函数还没写完,而且要计算整个数据集的误差,还需要在遍历数据时累计错误数,最后再推导整体误差率。我帮你把代码补全并整理好,同时拆解关键步骤:

完整可复现代码

import torch
from torch.autograd import Variable
import torch.optim as optim
import torchvision
import torchvision.transforms as transforms
from math import inf
from pdb import set_trace as st

def error_criterion(outputs, labels):
    # 取出模型输出中概率最高的类别索引
    max_vals, max_indices = torch.max(outputs, 1)
    # 统计当前批次中预测错误的样本数量
    incorrect = (max_indices != labels).sum().item()  # PyTorch 0.3.1用item()获取标量值
    return incorrect

# 数据预处理流程
transform = transforms.Compose(
    [transforms.ToTensor(),
     transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))])

# 加载CIFAR10训练集
trainset = torchvision.datasets.CIFAR10(root='./data', train=True,
                                        download=True, transform=transform)
trainloader = torch.utils.data.DataLoader(trainset, batch_size=4,
                                          shuffle=True, num_workers=2)

# 定义一个基础CNN模型(你可以替换成自己的模型)
class Net(torch.nn.Module):
    def __init__(self):
        super(Net, self).__init__()
        self.conv1 = torch.nn.Conv2d(3, 6, 5)
        self.pool = torch.nn.MaxPool2d(2, 2)
        self.conv2 = torch.nn.Conv2d(6, 16, 5)
        self.fc1 = torch.nn.Linear(16 * 5 * 5, 120)
        self.fc2 = torch.nn.Linear(120, 84)
        self.fc3 = torch.nn.Linear(84, 10)

    def forward(self, x):
        x = self.pool(torch.nn.functional.relu(self.conv1(x)))
        x = self.pool(torch.nn.functional.relu(self.conv2(x)))
        x = x.view(-1, 16 * 5 * 5)
        x = torch.nn.functional.relu(self.fc1(x))
        x = torch.nn.functional.relu(self.fc2(x))
        x = self.fc3(x)
        return x

net = Net()
criterion = torch.nn.CrossEntropyLoss()
optimizer = optim.SGD(net.parameters(), lr=0.001, momentum=0.9)

# 计算整个数据集误差的核心函数
def calculate_dataset_error(loader, model):
    model.eval()  # 切换到评估模式,关闭训练特有的层(如Dropout、BatchNorm)
    total_incorrect = 0
    total_samples = 0
    with torch.no_grad():  # PyTorch 0.3.1也可以用Variable(..., volatile=True),效果一致
        for data in loader:
            images, labels = data
            inputs = Variable(images, volatile=True)  # 评估阶段无需计算梯度,节省资源
            outputs = model(inputs)
            # 累计当前批次的错误数
            total_incorrect += error_criterion(outputs, labels)
            total_samples += labels.size(0)
    # 计算整体误差率
    error_rate = total_incorrect / total_samples
    model.train()  # 切回训练模式,不影响后续训练流程
    return error_rate

# 示例:训练2个epoch后计算数据集误差
for epoch in range(2):
    running_loss = 0.0
    for i, data in enumerate(trainloader, 0):
        inputs, labels = data
        inputs, labels = Variable(inputs), Variable(labels)
        
        optimizer.zero_grad()
        
        outputs = net(inputs)
        loss = criterion(outputs, labels)
        loss.backward()
        optimizer.step()
        
        running_loss += loss.item()
        if i % 2000 == 1999:
            print('[%d, %5d] loss: %.3f' %
                  (epoch + 1, i + 1, running_loss / 2000))
            running_loss = 0.0

print('Finished Training')

# 计算并打印整个训练集的误差率
train_error_rate = calculate_dataset_error(trainloader, net)
print(f'整个训练集的误差率: {train_error_rate:.4f}')

关键细节说明

  • error_criterion函数:专注于统计单批次错误样本数,而非直接计算误差率,这样更便于累计整个数据集的错误总量。
  • calculate_dataset_error函数:
    • 切换eval()模式是为了避免训练层干扰评估结果;
    • volatile=True是PyTorch 0.3.1中推荐的评估模式,禁止梯度计算,大幅降低内存占用;
    • 遍历完所有数据后,用总错误数除以总样本数得到最终误差率。
  • PyTorch 0.3.1特性适配:
    • 必须手动将Tensor转为Variable才能进行自动求导;
    • 获取标量值用.item(),也可以用.data[0]替代;
    • 评估时的梯度禁用逻辑要适配旧版本语法。

内容的提问来源于stack exchange,提问作者Charlie Parker

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.25 03:24:46