多模型误分类场景下,如何从TensorFlow的image_dataset_from_directory中根据索引列表获取对应图片
解决方法:从
image_dataset_from_directory中按索引提取图片 好问题!因为image_dataset_from_directory返回的是分批迭代的tf.data.Dataset对象,没法直接像普通列表那样用索引取值,不过结合你已经设置shuffle=False的前提(这个很关键,保证了图片顺序和文件夹完全一致),我们可以用两种方式拿到目标图片:
方法1:把数据集转换成可直接索引的数组(适合小数据集)
如果你的数据集不大,可以先把整个数据集的图片和标签拼接成完整的numpy数组,之后就能直接用索引提取了:
步骤1:将数据集转为完整数组
在加载train_ds之后,执行以下代码:
import numpy as np # 遍历所有批次,收集图片和标签 all_images = [] all_labels = [] for batch_images, batch_labels in train_ds: all_images.append(batch_images.numpy()) all_labels.append(batch_labels.numpy()) # 拼接成完整的数组 all_images = np.concatenate(all_images, axis=0) # shape: (总图片数, 高度, 宽度, 通道数) all_labels = np.concatenate(all_labels, axis=0) # shape: (总图片数,)
步骤2:提取共同误分类的图片
假设你已经通过求交集得到了所有模型都误分类的索引集合(比如common_miss_indices),直接用索引数组提取即可:
# 提取目标图片和对应标签 missed_images = all_images[list(common_miss_indices)] missed_labels = all_labels[list(common_miss_indices)]
方法2:遍历批次时筛选目标索引(适合大数据集)
如果数据集很大,一次性加载所有图片会占用过多内存,可以在遍历批次的过程中,根据索引范围筛选目标图片:
步骤1:先得到共同误分类的索引集合(用set提高查询效率)
首先把每个模型的误分类索引求交集,得到所有模型都判断错误的索引:
# 假设missclassified_train_folders是你收集的每个模型的误分类索引列表 common_miss_set = set(missclassified_train_folders[0]) for indices in missclassified_train_folders[1:]: common_miss_set.intersection_update(set(indices))
步骤2:遍历批次筛选目标图片
missed_images = [] missed_labels = [] current_start_idx = 0 # 跟踪当前批次的起始索引 for batch_images, batch_labels in train_ds: batch_size = batch_images.shape[0] # 计算当前批次覆盖的索引范围 batch_idx_range = range(current_start_idx, current_start_idx + batch_size) # 找出当前批次中属于共同误分类的索引位置 target_positions = [i for i, idx in enumerate(batch_idx_range) if idx in common_miss_set] if target_positions: # 提取当前批次中的目标图片和标签 missed_images.extend(batch_images.numpy()[target_positions]) missed_labels.extend(batch_labels.numpy()[target_positions]) current_start_idx += batch_size
额外补充:保存/查看提取的图片
image_dataset_from_directory默认会把图片转成float32格式,像素值缩放到[0,1]区间,如果要保存或查看,需要先转成uint8格式:
from PIL import Image # 示例:保存第一张误分类图片 img_array = missed_images[0] img_array = (img_array * 255).astype(np.uint8) # 还原到0-255的像素范围 img = Image.fromarray(img_array) img.save("missclassified_example.jpg")
内容的提问来源于stack exchange,提问作者albertovpd
相关产品推荐
相关产品推荐

