如何在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
相关产品推荐
相关产品推荐

