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

多模型误分类场景下,如何从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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.28 16:12:36