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

使用TFLite ModelMaker训练后,测试集混淆矩阵与准确率不符的问题

问题分析与修正方案

你遇到的核心问题是手动提取标签和预测结果的代码存在多处错误,导致混淆矩阵与模型评估的准确率不符。以下是具体问题和修正后的代码:

具体错误点

  1. 数据集迭代后耗尽:同一个ds被先用于提取标签,再用于预测,但TensorFlow Dataset是一次性迭代对象,遍历后会空,导致预测结果错误。
  2. 变量名笔误:混淆矩阵中使用test_labs和test_preds,但实际定义的变量是test_labels和test_pred,属于拼写错误。
  3. 预测索引转换不可靠:通过classes.index(pred[0][0])用类别名称反查索引,若类别名称与model.index_to_label的顺序不匹配(或存在重复名称),会导致索引错误。
  4. 标签提取逻辑可能错误:若test_data.gen_dataset()返回的是批量数据,label[0].numpy()的提取方式仅适用于batch_size=1的情况,通用性差。

修正后的代码

步骤1:正确提取测试标签与图像

先将测试集转换为列表缓存,避免重复生成数据集和迭代耗尽的问题:

import numpy as np
import tensorflow as tf
from tflite_model_maker import image_classifier

# 已有的训练代码保持不变
data = image_classifier.DataLoader.from_folder(data_root)
train_data, rest_data = data.split(0.7)
validation_data, test_data = rest_data.split(0.5)
model = image_classifier.create(train_data, validation_data=validation_data, epochs=20)

# 缓存测试集样本
test_samples = list(test_data.gen_dataset(batch_size=1).unbatch())
# 提取测试标签(已映射为类别索引)
test_labels = [label.numpy() for _, label in test_samples]
# 提取测试图像(用于预测)
test_images = np.array([image.numpy() for image, _ in test_samples])

步骤2:正确获取预测结果

直接通过预测概率取最大索引,避免类别名称转换的风险:

# 获取预测概率
predictions = model.predict(test_images)
# 对每个预测结果取argmax得到类别索引
test_pred = [np.argmax(pred) for pred in predictions]

步骤3:生成混淆矩阵并验证

# 计算混淆矩阵
confusion_mat = tf.math.confusion_matrix(test_labels, test_pred, num_classes=4)
print("混淆矩阵:")
print(confusion_mat.numpy())

# 再次验证模型准确率
loss, accuracy = model.evaluate(test_data)
print(f"模型准确率:{accuracy:.2%}")

为什么之前的代码会出错?

  • 模型evaluate方法直接调用test_data内部的数据集处理逻辑,是可靠的,所以准确率结果正确;而你手动处理数据集时,因为迭代耗尽、索引转换错误等问题,导致混淆矩阵完全失真。
  • 使用model.predict_top_k返回的是类别名称,再通过index反查索引的方式,不如直接从预测概率取索引稳定,尤其是当类别名称存在特殊字符或排序变化时,容易出错。

内容的提问来源于stack exchange,提问作者André

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.14 15:01:03