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

TensorFlow处理CIFAR10可视化预测图像报TypeError如何解决

问题根源

报错的核心原因是CIFAR10数据集加载后的train_labels、test_labels默认是二维数组(形状为(样本数, 1)),当你在plot_value_array函数中执行true_label = true_label[i]时,拿到的不是普通整数标量,而是单元素numpy数组(比如array([5], dtype=int32)),用数组作为索引访问matplotlib的bar对象列表,就会触发这个类型错误。

解决方案

方案一:提前压缩标签数组维度(推荐)

在数据集加载完成后添加如下代码,把标签数组转为一维:

train_labels = train_labels.squeeze()
test_labels = test_labels.squeeze()

squeeze()会自动删除数组中长度为1的维度,后续取索引时就能直接拿到整数标量,一劳永逸解决问题。

方案二:修改函数内的取值逻辑

如果不想修改原始标签数组,可以在两个绘图函数中手动把单元素数组转为标量:

# 修改plot_image函数内的取值代码
true_label, img = true_label[i].item(), img[i]

# 修改plot_value_array函数内的取值代码
true_label = true_label[i].item()

.item()方法可以直接把单元素numpy数组转为Python原生整数。

额外优化说明

你当前的预测逻辑存在冗余:CNN的输出层已经设置了activation='softmax',后续又叠加了一层Softmax(),等于对输出做了两次归一化,会导致预测概率结果不准确。可以直接删除probability_model的定义,预测代码改为:

predictions = cnn.predict(test_images)

内容的提问来源于stack exchange,提问作者Aditya Aryan

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.25 20:45:03