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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.29 07:08:05