未打乱测试数据时Torchmetrics Accuracy计算异常的原因
出现这种现象的核心原因大概率是以下两种情况之一:
测试阶段模型未进入评估模式,随机层持续生效
如果你在测试代码中没有正确让模型进入eval()模式(比如PyTorch Lightning默认会自动处理,但如果手动调用了train()或者修改了层的状态),模型中的dropout、随机深度等随机层会在测试时继续随机失活神经元。当测试集按类别连续输入时,连续的同类样本会被随机层干扰,导致预测稳定性极差,错误率飙升;而打乱测试集后,不同类别的样本交替输入,随机失活的影响被平均,准确率自然回到正常的80-85%。手动评估时你应该是手动开启了eval()模式,所以结果正确。torchmetrics.Accuracy的参数设置与输入不匹配,且测试集连续放大了错误
如果你直接将模型输出的logits传入Accuracy,但没有设置prediction_type="logits"参数,torchmetrics会默认把输入的logits当成类别索引(只能是0或1)来计算准确率。由于logits是连续的实数(比如猫的logits是[10, -10]),和标签的0/1几乎不匹配,所以未打乱时连续输入同类样本会导致准确率接近随机水平(40%左右);而打乱后,样本混合使得logits与标签的匹配概率被修正,准确率恢复正常。
另外还有一种小概率情况:测试集的标签编码和你手动评估时的类别对应关系完全反向,且测试集按类别连续排列。比如代码中标签0代表狗,但你手动评估时把模型输出的0当成猫,导致手动评估认为正确,但代码中计算时这部分预测会被判定为错误。不过这种情况和你手动评估的结果矛盾,概率较低。
内容的提问来源于stack exchange,提问作者cplusplus

