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

基于image_dataset_from_directory的FER2013混淆矩阵生成方法正确性疑问

问题分析与解决方案

两种方法的核心逻辑都是正确的,出现显著差异的原因大概率是测试集的随机打乱导致样本顺序不一致,或是多余的标签转换步骤引入了潜在问题。以下是具体分析和修正方案:

核心问题排查

  1. 测试集打乱的随机状态干扰
    tf.keras.preprocessing.image_dataset_from_directory默认shuffle=True,即使设置了seed=42,TensorFlow的全局随机状态仍可能在两次运行(方法1/方法2)之间被其他操作篡改,导致两次迭代测试集的样本顺序完全不同,最终混淆矩阵结果自然差异巨大。

  2. 多余的标签转换步骤
    两种方法中都对原始整数标签执行了to_categorical再np.argmax的冗余操作,虽然理论上不会出错,但多余的转换步骤增加了潜在风险(比如num_classes设置错误),且完全没有必要。

修正后的正确代码

第一步:固定测试集顺序

创建测试集时显式设置shuffle=False,保证样本顺序完全固定:

test_data = tf.keras.preprocessing.image_dataset_from_directory(
        test_directory,
        image_size=(48, 48),
        batch_size=64,
        seed=42,
        color_mode="grayscale",
        shuffle=False  # 关键:测试集无需打乱,固定顺序避免随机干扰
)

方法1:合并完整数据集后生成混淆矩阵(简化版)

去掉冗余的标签转换,直接使用原始整数标签计算:

x_test_complete = np.concatenate([x for x, y in test_data], axis=0)
y_test_complete = np.concatenate([y for x, y in test_data], axis=0)
y_test_pred = model.predict(x_test_complete, verbose=0)

cm = confusion_matrix(
        y_test_complete, 
        np.argmax(y_test_pred, axis=1), 
        labels=range(7)
)

方法2:逐批累加混淆矩阵(简化版)

同样去掉冗余的to_categorical步骤:

conf_matrix = np.zeros((7, 7))

for x_test, y_test in test_data.as_numpy_iterator():
        y_test_pred = model.predict(x_test, verbose=0)
        batch_cm = confusion_matrix(
                y_test,
                np.argmax(y_test_pred, axis=1),
                labels=range(7)
        )
        conf_matrix += batch_cm

验证一致性

完成上述修正后,两种方法生成的混淆矩阵应完全一致。如果仍有差异,可进一步排查:

  • 对比两种方法中第一个批次的标签和预测结果,确认输入样本和模型输出完全一致
  • 检查模型是否存在推理模式下的异常(比如未关闭的Dropout层,可通过model.evaluate验证模型稳定性)

内容的提问来源于stack exchange,提问作者Murillo Carvalho

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.17 14:34:55