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

使用flow_from_dataframe时,如何初始化y_true和y_pred生成混淆矩阵与分类报告?

解决Keras flow_from_dataframe无法获取classes属性的问题

问题原因

当使用flow_from_dataframe并设置class_mode="raw"时,返回的DataFrameIterator对象不会生成classes和class_indices属性。这个模式的设计是直接返回原始标签值,而非像flow_from_directory默认分类模式那样生成编码后的标签集合,因此无法通过val_set.classes获取真实标签。

解决方案

方法1:直接从原始验证集DataFrame提取真实标签

既然你已经持有验证集的DataFrame(val),直接从对应列提取真实标签即可,同时注意关闭验证集迭代器的shuffle保证顺序一致:

# 加载数据时关闭验证集的shuffle,确保预测结果与真实标签顺序匹配
val_set = val_datagen.flow_from_dataframe(
    val,
    path,
    x_col="image_name",
    y_col="level",
    class_mode="raw",
    color_mode="rgb",
    batch_size=32,
    target_size=(64, 64),
    shuffle=False)  # 关键:必须关闭shuffle

# 执行预测
Y_pred = model.predict(val_set)
y_pred = np.argmax(Y_pred, axis=1)

# 从原始DataFrame获取真实标签
y_true = val["level"].values
# 获取排序后的类别标签(与模型输出类别顺序对齐)
class_labels = sorted(val["level"].unique())

# 生成混淆矩阵和分类报告
print('Confusion Matrix')
print(confusion_matrix(y_true, y_pred))
print('Classification Report')
# 若标签是数值类型,转为字符串作为目标名称更易读
print(classification_report(y_true, y_pred, target_names=[str(lbl) for lbl in class_labels]))

方法2:修改class_mode为分类模式(可选)

如果希望保持和flow_from_directory类似的属性,可以将class_mode改为"sparse"(适用于整数标签的多分类)或"categorical"(适用于one-hot编码标签),此时DataFrameIterator会自动生成classes和class_indices属性:

# 修改class_mode为sparse(假设标签是整数类型)
val_set = val_datagen.flow_from_dataframe(
    val,
    path,
    x_col="image_name",
    y_col="level",
    class_mode="sparse",
    color_mode="rgb",
    batch_size=32,
    target_size=(64, 64),
    shuffle=False)

# 此时可直接使用val_set.classes
y_true = val_set.classes
class_labels = list(val_set.class_indices.keys())

# 后续预测和评估代码不变
Y_pred = model.predict(val_set)
y_pred = np.argmax(Y_pred, axis=1)
print(confusion_matrix(y_true, y_pred))
print(classification_report(y_true, y_pred, target_names=class_labels))

注意:使用class_mode="sparse"或"categorical"时,需确保模型的输出层与该模式匹配(例如sparse_categorical_crossentropy损失对应sparse模式)。

内容的提问来源于stack exchange,提问作者Sheetal Nagar

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.09 19:15:24