TensorFlow image_dataset_from_directory标签映射错误问题排查
问题原因及修正方案
核心错误:标签解析逻辑错误
你设置的label_mode="int",这意味着数据集返回的标签是单个整数(对应类别的索引),而非one-hot编码的数组。但你在显示标签时用了labels[i].numpy().argmax()——这个方法是用来从one-hot数组中提取类别索引的,对单个整数使用argmax()会始终返回0,所以你看到的永远是class_names列表的第一个元素,也就是第一个子目录的名称。删除第一个子目录后,class_names的第一个元素变成下一个子目录,自然就显示那个名称了。
次要错误:数据集加载代码的语法问题
你的加载代码里有两处缺少逗号的语法错误,会导致Python解析失败:
batch_size=batch_size行末尾没有逗号subset="training"行末尾没有逗号
修正后的代码
1. 修正数据集加载代码
training_dataset= tf.keras.utils.image_dataset_from_directory( labeled_data_dir, image_size=image_size, # 调整图像大小(按需修改) batch_size=batch_size, # 训练批次大小(按需修改) seed=9, validation_split=0.2, subset="training", labels="inferred", label_mode="int" )
2. 修正图像与标签显示代码
# 先获取数据集的类别名称列表 class_names = training_dataset.class_names for images, labels in training_dataset.take(1): # 取一个批次的数据 for i in range(9): # 显示批次中前9张图片 ax = plt.subplot(3, 3, i + 1) plt.imshow(images[i].numpy().astype("uint8")) # 直接用标签的整数值作为索引取类别名称,无需argmax() plt.title(class_names[labels[i].numpy()]) plt.axis("off") plt.show()
内容的提问来源于stack exchange,提问作者JDT
相关产品推荐
相关产品推荐

