TF Data API:如何为物体集合识别生成TensorFlow输入
没问题!TensorFlow的tf.data API完全能实现你需要的动态样本抽取逻辑,我给你梳理一套从图像目录开始的完整流程,涵盖随机选类别、动态调整每批次每类样本数的核心逻辑:
完整实现流程
1. 加载原始图像数据集
首先用tf.keras.utils.image_dataset_from_directory加载目录结构的数据集(比如ImageNet就是每个类别对应一个子文件夹),这个函数会自动处理目录结构,返回包含图像和对应标签的数据集:
import tensorflow as tf import numpy as np # 你的数据集根目录,结构如下: # dataset_root/ # class_0/ # img0.jpg # img1.jpg # ... # class_1/ # ... # ... dataset_root = "/path/to/your/dataset" # 加载基础数据集,先不设置固定批次,后续自定义批次逻辑 raw_dataset = tf.keras.utils.image_dataset_from_directory( dataset_root, image_size=(224, 224), # 根据你的模型输入尺寸调整 batch_size=None, # 单个元素为(单张图像, 对应标签) label_mode='int' )
2. 按类别分组数据集
我们需要把同一类别的图像归为一组,方便后续随机抽取。这里用tf.data.experimental.group_by_key实现分组,先把数据格式转换成(标签, 图像)的键值对形式:
# 转换为(key=标签, value=图像)的格式,适配group_by_key key_value_dataset = raw_dataset.map(lambda img, lbl: (lbl, img)) # 按标签分组,每个组包含对应类别的所有图像 grouped_dataset = tf.data.experimental.group_by_key( key_value_dataset, reduce_func=lambda key, dataset: dataset.batch(-1) # 把同类别所有图像合并成一个tensor ) # 将分组后的数据集转为列表(如果数据集极大,建议用迭代器而非列表,避免内存占用过高) class_groups = list(grouped_dataset.as_numpy_iterator()) # class_groups的每个元素格式:(类别标签, 该类别所有图像的数组)
3. 自定义动态批次生成逻辑
接下来实现核心需求:每个批次随机选择若干类别,每个类别随机抽取指定数量的样本(不同批次的样本数可以不同)。我们写一个生成器函数,再转换成TensorFlow数据集:
def dynamic_batch_generator(class_groups, min_classes=2, max_classes=5, min_imgs_per_cls=2, max_imgs_per_cls=4): while True: # 随机确定当前批次的类别数量 num_classes = np.random.randint(min_classes, max_classes + 1) # 随机确定当前批次每类的样本数量 num_imgs_per_cls = np.random.randint(min_imgs_per_cls, max_imgs_per_cls + 1) # 随机挑选num_classes个不同的类别 selected_classes = np.random.choice(len(class_groups), size=num_classes, replace=False) batch_imgs = [] batch_labels = [] for cls_idx in selected_classes: cls_label, cls_imgs = class_groups[cls_idx] # 从当前类别中随机抽取num_imgs_per_cls个样本 selected_img_indices = np.random.choice(cls_imgs.shape[0], size=num_imgs_per_cls, replace=False) selected_imgs = cls_imgs[selected_img_indices] batch_imgs.append(selected_imgs) # 生成对应数量的标签 batch_labels.append(np.full(num_imgs_per_cls, cls_label)) # 合并成一个完整批次 batch_imgs = np.concatenate(batch_imgs, axis=0) batch_labels = np.concatenate(batch_labels, axis=0) # 可选:打乱批次内的样本顺序,避免同类别样本集中 shuffle_indices = np.random.permutation(len(batch_imgs)) batch_imgs = batch_imgs[shuffle_indices] batch_labels = batch_labels[shuffle_indices] yield batch_imgs, batch_labels # 将生成器转换为TensorFlow数据集,定义输出格式 dynamic_batch_dataset = tf.data.Dataset.from_generator( lambda: dynamic_batch_generator(class_groups), output_signature=( tf.TensorSpec(shape=(None, 224, 224, 3), dtype=tf.float32), tf.TensorSpec(shape=(None,), dtype=tf.int32) ) ) # 开启预取优化,提升数据加载效率 dynamic_batch_dataset = dynamic_batch_dataset.prefetch(tf.data.AUTOTUNE)
4. 对接模型训练/推理
现在这个dynamic_batch_dataset可以直接作为模型的输入,比如用于训练你的物体集合识别模型:
# 示例模型(替换成你的自定义模型) model = tf.keras.applications.ResNet50(weights=None, input_shape=(224,224,3), num_classes=len(class_groups)) model.compile(optimizer='adam', loss='sparse_categorical_crossentropy', metrics=['accuracy']) # 训练模型,steps_per_epoch根据你的需求设定 model.fit(dynamic_batch_dataset, steps_per_epoch=100, epochs=10)
额外注意事项
- 大数据集适配:如果是像完整ImageNet这样的超大数据集,不要把
class_groups转换成列表,而是用grouped_dataset.shuffle(buffer_size=len(class_groups))打乱分组,每次迭代取N个组再抽取样本,避免内存溢出。 - 图像预处理:可以在
raw_dataset后添加map操作实现归一化、数据增强,比如:def preprocess_img(img, lbl): img = tf.keras.applications.resnet50.preprocess_input(img) img = tf.image.random_flip_left_right(img) # 随机翻转增强 return img, lbl raw_dataset = raw_dataset.map(preprocess_img, num_parallel_calls=tf.data.AUTOTUNE) - 参数灵活调整:你可以修改
dynamic_batch_generator里的min_classes、max_classes等参数,或者改成根据每个类别的样本数动态生成num_imgs_per_cls(比如不超过该类别总样本数)。
内容的提问来源于stack exchange,提问作者ArthurSeat
相关产品推荐
相关产品推荐

