为何Caffe Accuracy Layer与PyCaffe计算的准确率存在差异?
解决PyCaffe中Accuracy Layer结果与自定义计算不一致的问题
我之前也碰到过一模一样的坑,咱们一步步拆解问题原因,再给出修正方案:
最常见的原因:Caffe返回的是单个batch的准确率,而非全局平均
默认情况下,当你跑完solver.test_nets[0].forward()一次,solver.test_nets[0].blobs['accuracy'].data存储的是当前这个测试batch的准确率,而不是所有测试batch的平均准确率。
但你的自定义run_test函数是累加所有正确样本数后除以总样本数,得到的是全局平均准确率——这两者自然会有差异,尤其是当测试集样本数不是batch size整数倍时,最后一个batch的准确率会拉偏结果。
其他可能的原因
1. 最后一个batch的样本数未处理
如果你的测试集总样本数不能被batch size整除,最后一个batch的实际样本数会小于设定的batch_size_test。如果自定义代码里直接用固定的batch_size_test计算总数,就会和Caffe的计算逻辑(自动用实际样本数计算该batch的准确率)产生误差。
2. Accuracy Layer的top_k参数不匹配
如果你的Accuracy Layer配置了top_k: N(比如top5准确率),那它会统计“预测结果前N个类别包含正确标签”的样本数;但如果你的自定义代码只判断top1的预测是否正确,结果肯定不一致。
修正后的自定义测试函数
下面是调整后的代码,完美对齐Caffe的计算逻辑:
def run_test(solver, test_iter): ''' Tests the network on all test set and calculates the test accuracy ''' correct = 0 total_samples = 0 caffe_acc_sum = 0 # 用来累加Caffe每个batch的准确率,最后计算全局平均 # 可选:重置测试网络状态,避免之前的运行残留影响 solver.test_nets[0].reset() for test_it in range(test_iter): solver.test_nets[0].forward() # 获取当前batch的实际样本数(处理最后一个batch不满的情况) current_batch_size = solver.test_nets[0].blobs['data'].data.shape[0] # --- 自定义计算部分 --- # 替换成你的网络输出层名称(比如fc8、prob等) pred = solver.test_nets[0].blobs['prob'].data.argmax(axis=1) # 获取真实标签 label = solver.test_nets[0].blobs['label'].data # 计算当前batch的正确样本数 batch_correct = (pred == label).sum() # --- 累加统计 --- correct += batch_correct total_samples += current_batch_size # --- 对比Caffe的batch准确率 --- caffe_batch_acc = solver.test_nets[0].blobs['accuracy'].data caffe_acc_sum += caffe_batch_acc * current_batch_size print(f"Batch {test_it+1}: Caffe batch acc = {caffe_batch_acc:.4f}, My batch acc = {batch_correct/current_batch_size:.4f}") # 计算全局平均准确率 my_global_acc = correct / total_samples caffe_global_acc = caffe_acc_sum / total_samples print(f"\nTotal Test Accuracy:") print(f"My Calculation = {my_global_acc:.4f}") print(f"Caffe Average = {caffe_global_acc:.4f}") return my_global_acc
额外注意事项
- 确认你的输出层名称(比如
prob)和标签层名称(label)与网络配置一致; - 如果Accuracy Layer用了
top_k参数,自定义代码也要对应调整:比如判断正确标签是否在预测概率前N的类别里; - 每次测试前调用
reset()可以避免网络缓存的中间结果干扰,保证测试的独立性。
内容的提问来源于stack exchange,提问作者Hossein
相关产品推荐
相关产品推荐

