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

如何在Keras中拼接原始训练图像数据集与增强后的训练数据集?

解决方案

要实现原始训练数据与增强后数据的拼接,你可以通过以下两种实用方式完成:

方法一:双生成器合并输出(适合大数据集)

这种方式不需要提前加载所有数据到内存,通过tf.data.Dataset合并两个生成器的输出,兼顾效率与内存友好性:

from tensorflow.keras.preprocessing.image import ImageDataGenerator
import tensorflow as tf

# 定义仅做归一化的生成器(加载原始训练数据)
original_train_datagen = ImageDataGenerator(
    rescale=1./255,
    validation_split=0.2
)

# 定义带数据增强的生成器
augmented_train_datagen = ImageDataGenerator(
    rescale=1./255, 
    rotation_range=5,
    zoom_range = 0.1,
    width_shift_range=0.1,
    height_shift_range=0.1,
    validation_split=0.2
)

test_datagen = ImageDataGenerator(rescale=1./255)

train_dir = 'train_separated'
test_dir = 'test_separated'
batch_size = 128
img_height = 100
img_width = 100
num_classes = 10

# 加载原始训练数据(关闭 shuffle 保证后续合并顺序对应)
original_train_data = original_train_datagen.flow_from_directory(
    train_dir,
    target_size=(img_height, img_width),
    batch_size=batch_size,
    class_mode='categorical', 
    subset='training',
    shuffle=False
)

# 加载增强后的训练数据
augmented_train_data = augmented_train_datagen.flow_from_directory(
    train_dir,
    target_size=(img_height, img_width),
    batch_size=batch_size,
    class_mode='categorical', 
    subset='training',
    shuffle=False
)

# 将两个生成器转为 tf.data.Dataset 格式
original_ds = tf.data.Dataset.from_generator(
    lambda: original_train_data,
    output_types=(tf.float32, tf.float32),
    output_shapes=((None, img_height, img_width, 3), (None, num_classes))
)

augmented_ds = tf.data.Dataset.from_generator(
    lambda: augmented_train_data,
    output_types=(tf.float32, tf.float32),
    output_shapes=((None, img_height, img_width, 3), (None, num_classes))
)

# 合并数据集并打乱,重新设置批次大小
combined_train_ds = original_ds.concatenate(augmented_ds).shuffle(original_train_data.samples * 2)
combined_train_ds = combined_train_ds.unbatch().batch(batch_size)

# 验证集与测试集加载逻辑保持不变
val_data = original_train_datagen.flow_from_directory(
    train_dir,
    target_size=(img_height, img_width),
    batch_size=batch_size,
    class_mode='categorical', 
    subset='validation')

test_data = test_datagen.flow_from_directory(
    test_dir,
    target_size=(img_height, img_width),
    batch_size=batch_size,
    class_mode='categorical')

方法二:内存拼接(适合小数据集)

如果数据集规模较小,可以直接将原始数据与增强数据全部加载到内存后拼接,逻辑更直观:

from tensorflow.keras.preprocessing.image import ImageDataGenerator
import numpy as np

# 定义生成器
original_train_datagen = ImageDataGenerator(rescale=1./255, validation_split=0.2)
augmented_train_datagen = ImageDataGenerator(
    rescale=1./255, 
    rotation_range=5,
    zoom_range = 0.1,
    width_shift_range=0.1,
    height_shift_range=0.1,
    validation_split=0.2
)

train_dir = 'train_separated'
batch_size = 128
img_height = 100
img_width = 100

# 加载所有原始训练数据
original_train_data = original_train_datagen.flow_from_directory(
    train_dir,
    target_size=(img_height, img_width),
    batch_size=batch_size,
    class_mode='categorical', 
    subset='training',
    shuffle=False
)
x_original, y_original = next(original_train_data)
for i in range(len(original_train_data)-1):
    img_batch, label_batch = next(original_train_data)
    x_original = np.concatenate((x_original, img_batch))
    y_original = np.concatenate((y_original, label_batch))

# 生成并加载增强训练数据(与原始数据数量一致)
augmented_train_data = augmented_train_datagen.flow_from_directory(
    train_dir,
    target_size=(img_height, img_width),
    batch_size=batch_size,
    class_mode='categorical', 
    subset='training',
    shuffle=False
)
x_augmented, y_augmented = next(augmented_train_data)
for i in range(len(augmented_train_data)-1):
    img_batch, label_batch = next(augmented_train_data)
    x_augmented = np.concatenate((x_augmented, img_batch))
    y_augmented = np.concatenate((y_augmented, label_batch))

# 拼接原始与增强数据并打乱
x_combined = np.concatenate((x_original, x_augmented))
y_combined = np.concatenate((y_original, y_augmented))

shuffle_idx = np.random.permutation(len(x_combined))
x_combined = x_combined[shuffle_idx]
y_combined = y_combined[shuffle_idx]

# 后续直接用拼接后的数组训练模型即可,例如:model.fit(x_combined, y_combined, epochs=10, validation_data=val_data)

关键注意事项

  • 验证集禁止添加数据增强,否则会干扰模型评估的准确性
  • 若需要更多增强样本,可多次迭代增强生成器,生成多份数据后再与原始数据拼接
  • 方法一优先用于大数据集,避免内存溢出;方法二更适合快速验证小数据集的增强效果

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.26 01:22:01