如何按类拆分MNIST数据集并添加指定增强?代码报错求助
问题修正:MNIST数据集采样与数据增强实现
需求说明
- 从含70000样本、10个类别的MNIST数据集中提取总计22000个样本,保持类别分布一致
- 训练集共20000样本:每类20张原始图像 + 1980张增强图像
- 测试集共2000样本:每类200张原始图像
- 仅使用剪切(shear)、旋转(rotation)、宽度平移(width-shift)、高度平移(height-shift)作为数据增强方式
报错原因
- 索引越界:原代码选择测试集样本时,误用了训练集的类别索引(
class_indices来自y_train_full)去访问测试集数据x_test_full,但测试集仅含10000个样本,训练集索引远超出这个范围,直接触发IndexError。 - 增强图像数量不足:原代码仅生成100张增强图,不符合每类1980张的要求。
- 输入维度不匹配:
ImageDataGenerator.flow要求输入为4D张量(样本数×高×宽×通道数),但MNIST原始数据是3D结构(样本数×28×28),会导致增强过程异常。
修正后的代码
import numpy as np from tensorflow import keras from tensorflow.keras.preprocessing.image import ImageDataGenerator # 加载MNIST数据集 (x_train_full, y_train_full), (x_test_full, y_test_full) = keras.datasets.mnist.load_data() # 归一化数据,并扩展维度为4D(适配ImageDataGenerator的输入要求) x_train_full = x_train_full / 255.0 x_train_full = np.expand_dims(x_train_full, axis=-1) x_test_full = x_test_full / 255.0 x_test_full = np.expand_dims(x_test_full, axis=-1) # 创建数据增强生成器 data_gen = ImageDataGenerator( shear_range=0.2, rotation_range=20, width_shift_range=0.2, height_shift_range=0.2 ) # 初始化空列表存储最终数据集 x_train, y_train, x_test, y_test = [], [], [], [] # 遍历每个类别 for class_n in range(10): # ---------------------- 处理训练集原始样本 ---------------------- train_class_indices = np.where(y_train_full == class_n)[0] # 随机选择20张原始训练样本 selected_train_indices = np.random.choice(train_class_indices, 20, replace=False) x_train.append(x_train_full[selected_train_indices]) y_train.append(y_train_full[selected_train_indices]) # ---------------------- 处理测试集样本 ---------------------- test_class_indices = np.where(y_test_full == class_n)[0] # 随机选择200张原始测试样本 selected_test_indices = np.random.choice(test_class_indices, 200, replace=False) x_test.append(x_test_full[selected_test_indices]) y_test.append(y_test_full[selected_test_indices]) # ---------------------- 生成增强训练样本 ---------------------- # 基于选中的20张原始图生成1980张增强图 aug_generator = data_gen.flow( x_train_full[selected_train_indices], y_train_full[selected_train_indices], batch_size=20 # 每次用20张原始图生成20张增强图 ) # 计算需要迭代的次数:1980 / 20 = 99次 total_aug = 1980 aug_images = [] aug_labels = [] for _ in range(total_aug // aug_generator.batch_size): batch_imgs, batch_labels = next(aug_generator) aug_images.append(batch_imgs) aug_labels.append(batch_labels) # 将增强样本加入训练集 x_train.append(np.concatenate(aug_images)) y_train.append(np.concatenate(aug_labels)) # 合并所有类别数据,得到最终数据集 x_train = np.concatenate(x_train) y_train = np.concatenate(y_train) x_test = np.concatenate(x_test) y_test = np.concatenate(y_test) # 打印数据集形状验证 print("训练集形状: ", x_train.shape) # 应为(20000, 28, 28, 1) print("训练标签形状: ", y_train.shape) # 应为(20000,) print("测试集形状: ", x_test.shape) # 应为(2000, 28, 28, 1) print("测试标签形状: ", y_test.shape) # 应为(2000,)
关键修正点
- 测试集索引分离:单独从测试集标签中提取对应类别的索引,彻底解决跨数据集访问的索引越界问题。
- 维度适配:将MNIST的3D数据扩展为4D,满足数据增强生成器的输入格式要求。
- 增强数量匹配:通过循环迭代生成器,确保每类生成恰好1980张增强图像。
- 数据合并优化:统一处理原始样本与增强样本的合并逻辑,保证最终数据集结构正确。
内容的提问来源于stack exchange,提问作者Reihan Maulana
相关产品推荐
相关产品推荐

