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
相关产品推荐
相关产品推荐

