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

如何同时对X_train与y_train应用数据增强并生成数据集

Hey there! Let's tackle your problem step by step. Since both your input X_train and output y_train are 28×28 images (and you're using zca_whitening), we need to make sure the ImageDataGenerator learns from both datasets during fitting, then generate and save the augmented data without feeding it directly into training.

Step 1: Prepare Your Data

First, let's get your MNIST data ready as you did, but we'll work with copies to avoid overwriting the original dataset:

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

# Load and preprocess MNIST data
(X_train, y_train), (X_test, y_test) = mnist.load_data()
X_train = X_train.reshape((X_train.shape[0], 28, 28, 1)).astype('float32')
y_train = X_train.copy()  # In your case y equals X; adjust this if your actual y is different

Step 2: Fit the DataGenerator on Combined X and y

To make the generator learn statistical features from both X_train and y_train, we'll concatenate them into a single dataset before fitting. This ensures the ZCA whitening transformation accounts for the distribution of both input and output images:

# Combine X and y along the sample axis (total samples become 2×original count)
combined_data = np.concatenate([X_train, y_train], axis=0)

# Initialize the generator with zca_whitening enabled
datagen = ImageDataGenerator(zca_whitening=True)

# Fit the generator on the combined dataset
datagen.fit(combined_data)

Step 3: Generate and Save Augmented Data

We have two scenarios depending on your exact needs:

Scenario 1: Only X gets augmented, y stays as original

If you just need to apply whitening to X_train and keep y_train unchanged (useful for some image-to-image tasks, even though y=X in your example):

batch_size = 32
# Create a generator that takes X and y, outputs augmented X + original y
generator = datagen.flow(X_train, y_train, batch_size=batch_size, shuffle=False)

# Calculate total batches needed to cover all samples
total_batches = int(np.ceil(X_train.shape[0] / batch_size))

# Collect augmented data
augmented_X = []
augmented_y = []

for _ in range(total_batches):
    x_batch, y_batch = next(generator)
    augmented_X.append(x_batch)
    augmented_y.append(y_batch)

# Convert lists to numpy arrays
augmented_X = np.concatenate(augmented_X, axis=0)
augmented_y = np.concatenate(augmented_y, axis=0)

# Save the datasets to disk
np.save('augmented_X_train.npy', augmented_X)
np.save('augmented_y_train.npy', augmented_y)

Scenario 2: Both X and y get the same augmentation

If you need y_train to undergo the exact same ZCA whitening transformation as X_train (which makes sense for paired image tasks like yours), we'll combine them along the channel dimension first, generate augmented samples, then split them back:

# Combine X and y along the channel axis (each sample becomes 28×28×2)
combined_samples = np.concatenate([X_train, y_train], axis=-1)

# Create generator for combined samples (no need to refit if you already did it on combined_data)
generator = datagen.flow(combined_samples, batch_size=batch_size, shuffle=False)

# Collect augmented combined samples
augmented_combined = []
for _ in range(total_batches):
    augmented_combined.append(next(generator))

augmented_combined = np.concatenate(augmented_combined, axis=0)

# Split back into augmented X and y
augmented_X = augmented_combined[..., 0:1]  # Extract first channel (original X)
augmented_y = augmented_combined[..., 1:2]  # Extract second channel (original y)

# Save the results
np.save('augmented_X_train.npy', augmented_X)
np.save('augmented_y_train.npy', augmented_y)

Quick Notes

  • Use shuffle=False if you need augmented samples to stay paired with their original counterparts (set to True if you don't care about order).
  • Adjust batch_size based on your available memory—smaller batches use less RAM.
  • You can verify the results by loading the saved .npy files and checking their shape/values with np.load() and print(augmented_X.shape).

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.06 10:07:32