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

如何用ImageDataGenerator限制每个类别仅使用N张图像训练迁移学习模型?

如何用ImageDataGenerator限制每个类别的图像数量?

问题描述

我现在有10个类别(每个对应一个独立目录),每个目录下包含800张图像,打算用迁移学习训练模型。目前我用ImageDataGenerator加载数据的代码如下:

train_datagen = ImageDataGenerator(rescale=1./255, shear_range=0.2, zoom_range=0.2, horizontal_flip=True, validation_split=0.2) # set validation split
train_generator = train_datagen.flow_from_directory(
 train_data_dir,
 target_size=(img_height, img_width),
 batch_size=batch_size,
 class_mode='binary',
 subset='training') # set as training data
validation_generator = train_datagen.flow_from_directory(
 train_data_dir, # same directory as training data
 target_size=(img_height, img_width),
 batch_size=batch_size,
 class_mode='binary',
 subset='validation') # set as validation data
model.fit_generator(
 train_generator,
 steps_per_epoch = train_generator.samples // batch_size,
 validation_data = validation_generator,
 validation_steps = validation_generator.samples // batch_size,
 epochs = nb_epochs)

想请教下:能不能通过ImageDataGenerator限制每个目录只使用100张(或指定N张)图像,而不是全部800张?


解决方案

当然可以实现!ImageDataGenerator本身没有直接提供限制单类别样本数的参数,但我们有两种实用的方法来达成需求:

方法一:手动筛选并复制数据(最直观易上手)

这是最简单的方式,不需要修改太多代码:

  • 新建一个临时的数据集目录,结构和原目录完全一致(10个类别子目录)
  • 对每个类别目录,从原目录中复制你需要的N张图像(比如100张)到临时目录对应的子目录里
  • 之后直接让ImageDataGenerator从这个临时目录加载数据就行

这种方法的优势是逻辑清晰,数据完全可控,不容易出bug;唯一的小缺点是需要额外的磁盘空间来存储筛选后的数据集。

方法二:自定义生成器(更灵活,无需额外存储)

如果你不想额外占用磁盘空间,可以基于原数据集构建一个自定义生成器,只取每个类别前N张图像。这里给你一个可直接参考的实现:

import os
import glob
import numpy as np
from tensorflow.keras.utils import load_img, img_to_array
from sklearn.model_selection import train_test_split

# 配置参数
train_data_dir = "你的训练数据根目录"
img_height, img_width = 224, 224  # 可根据你的模型调整
batch_size = 32
N = 100  # 每个类别要保留的图像数量
nb_epochs = 10

# 第一步:收集每个类别中前N张图像的路径
class_image_paths = {}
for class_name in os.listdir(train_data_dir):
    class_dir = os.path.join(train_data_dir, class_name)
    if os.path.isdir(class_dir):
        # 匹配该类别下的所有图像文件(根据你的图像格式调整后缀,比如png)
        all_imgs = glob.glob(os.path.join(class_dir, "*.jpg"))
        class_image_paths[class_name] = all_imgs[:N]  # 只保留前N张

# 第二步:整理所有选中的图像和对应的标签
all_images = []
all_labels = []
class_to_index = {name: idx for idx, name in enumerate(class_image_paths.keys())}

for class_name, paths in class_image_paths.items():
    for img_path in paths:
        all_images.append(img_path)
        all_labels.append(class_to_index[class_name])

# 第三步:拆分训练集和验证集(保持类别分布一致)
train_imgs, val_imgs, train_labels, val_labels = train_test_split(
    all_images, all_labels, test_size=0.2, stratify=all_labels, random_state=42
)

# 第四步:定义自定义生成器,支持数据增强
def custom_data_generator(image_paths, labels, data_augmenter, batch_size):
    while True:
        # 每次迭代前打乱数据顺序
        shuffled_indices = np.random.permutation(len(image_paths))
        for start_idx in range(0, len(image_paths), batch_size):
            batch_indices = shuffled_indices[start_idx:start_idx+batch_size]
            batch_imgs = []
            batch_lbls = []
            
            for idx in batch_indices:
                # 加载并预处理图像
                img = load_img(image_paths[idx], target_size=(img_height, img_width))
                img_array = img_to_array(img)
                # 应用数据增强
                img_array = data_augmenter.random_transform(img_array)
                img_array = data_augmenter.standardize(img_array)
                
                batch_imgs.append(img_array)
                batch_lbls.append(labels[idx])
            
            yield np.array(batch_imgs), np.array(batch_lbls)

# 初始化数据增强器(这里去掉validation_split,因为我们自己拆分了数据集)
train_datagen = ImageDataGenerator(
    rescale=1./255,
    shear_range=0.2,
    zoom_range=0.2,
    horizontal_flip=True
)

# 创建训练和验证生成器
train_generator = custom_data_generator(train_imgs, train_labels, train_datagen, batch_size)
val_generator = custom_data_generator(val_imgs, val_labels, train_datagen, batch_size)

# 训练模型(注意:新版本Keras已弃用fit_generator,直接用fit即可)
model.fit(
    train_generator,
    steps_per_epoch=len(train_imgs) // batch_size,
    validation_data=val_generator,
    validation_steps=len(val_imgs) // batch_size,
    epochs=nb_epochs
)

小提醒

  • 如果你使用的是较新版本的Keras/TensorFlow,fit_generator已经被标记为弃用,直接使用model.fit()就可以支持生成器输入
  • 使用自定义生成器时,记得用stratify参数拆分数据集,这样能保证训练集和验证集中每个类别的样本比例一致,模型评估结果会更可靠

内容的提问来源于stack exchange,提问作者Jedi Nerd

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.06 20:09:10