如何用PyTorch打印各类别预测准确率?代码问题排查
问题排查与修正方案
我帮你梳理了代码里的几个关键错误,这些正是导致预测准确率不符合预期的核心原因:
1. 硬编码的Batch Size循环
你写了for i in range(4)来遍历batch里的样本,但实际数据集的batch size不一定固定为4(比如最后一个batch的样本数通常会更少)。这会引发两种问题:要么索引越界报错,要么漏算/多算样本,最终让统计结果完全失真。
2. 输入变量名错误
你用了outputs = model(inputs),但从dataloaders['val']加载的变量是images, labels = data,这里应该用images作为模型输入,而非未定义的inputs——这会导致模型用了错误的输入数据,预测结果自然不可能正确。
3. 统计变量未重置
class_correct和class_total应该在每个epoch的验证阶段开始前重新初始化,否则会累加多个epoch的统计结果,导致准确率计算彻底偏离预期。
4. 模型模式未切换
验证阶段必须把模型切换到eval()模式,关闭dropout、batch norm等训练专属的层行为,否则会严重影响预测结果的稳定性和准确性。
修正后的完整代码片段
把这段代码整合到你的迁移学习流程中,替换原有的统计部分:
num_epochs = 1 num_classes = 3 # 你的数据集类别数 for epoch in range(num_epochs): print(f'Epoch {epoch+1}/{num_epochs}') print('-' * 10) # 每个epoch包含训练和验证两个阶段 for phase in ['train', 'val']: if phase == 'train': model.train() # 切换到训练模式 else: model.eval() # 切换到评估模式 # 初始化当前epoch验证阶段的类别统计变量 class_correct = list(0. for _ in range(num_classes)) class_total = list(0. for _ in range(num_classes)) running_loss = 0.0 running_corrects = 0 # 遍历数据加载器 for data in dataloaders[phase]: images, labels = data # 把数据移到指定设备(GPU/CPU) images = images.to(device) labels = labels.to(device) # 训练阶段清空梯度 optimizer.zero_grad() # 前向传播:训练阶段开启梯度,验证阶段关闭 with torch.set_grad_enabled(phase == 'train'): outputs = model(images) # 使用正确的输入变量images _, preds = torch.max(outputs, 1) loss = criterion(outputs, labels) # 训练阶段执行反向传播和优化 if phase == 'train': loss.backward() optimizer.step() # 累计全局的loss和正确数 running_loss += loss.item() * images.size(0) running_corrects += torch.sum(preds == labels.data) # 验证阶段统计每个类别的正确数和总数 if phase == 'val': correct_mask = (preds == labels.data).squeeze() # 遍历当前batch的所有样本,用实际batch大小替代硬编码的4 for i in range(labels.size(0)): label = labels.data[i] class_correct[label] += correct_mask[i].item() class_total[label] += 1 # 计算当前阶段的整体loss和准确率 epoch_loss = running_loss / dataset_sizes[phase] epoch_acc = running_corrects.double() / dataset_sizes[phase] print(f'{phase} Loss: {epoch_loss:.4f} Acc: {epoch_acc:.4f}') # 验证阶段结束后打印每个类别的准确率 if phase == 'val': print('\nPer class accuracy:') for i in range(num_classes): if class_total[i] > 0: acc = 100 * class_correct[i] / class_total[i] print(f'Accuracy of class {i}: {class_correct[i]:.0f} / {class_total[i]:.0f} = {acc:.4f} %') else: print(f'Accuracy of class {i}: No samples in validation set') print()
关键修正点说明
- 用
labels.size(0)替代硬编码的4,适配任意batch size,避免索引错误和统计遗漏; - 把
model(inputs)改为model(images),使用从数据加载器中获取的正确输入; - 在每个epoch的验证阶段开始时重新初始化
class_correct和class_total,确保统计只针对当前epoch的验证数据; - 通过
model.eval()和torch.set_grad_enabled(phase == 'train'),确保验证阶段模型处于正确的评估状态; - 统计时用
.item()把张量转为普通数值,避免张量累加导致的计算异常。
这样修改后,每个类别的准确率统计就会和全局的running_corrects结果完全匹配,比如你提到的running_corrects = 2 + 2的情况就能得到正确的统计结果。
内容的提问来源于stack exchange,提问作者Boooooooooms
相关产品推荐
相关产品推荐

