如何用Keras ImageDataGenerator生成同维度同顺序的MNIST随机旋转数组
一次性获取MNIST随机旋转后的完整数据集
嘿,我完全懂你的困惑——ImageDataGenerator的教程大多围绕训练时的batch和epoch循环,但其实要一次性处理整个数据集、拿到和原数组维度完全一致的结果,根本不需要那么复杂!
下面给你两种简单直接的方法,都能满足你的需求:
方法1:用ImageDataGenerator一次性生成全部样本
ImageDataGenerator的flow方法返回的是一个生成器,只要把batch_size设为你的数据集总样本数,就能一次拿到所有旋转后的图片,不用循环epoch。不过要注意它默认期望的是通道在后((样本数, 高, 宽, 通道数))的格式,而你的数据是通道在前((5000,1,28,28)),所以需要先做维度转换:
import numpy as np from keras.preprocessing.image import ImageDataGenerator # 假设你的原始数据X是(5000, 1, 28, 28)的numpy数组 # 1. 转换为ImageDataGenerator需要的通道在后格式 X_channels_last = np.transpose(X, (0, 2, 3, 1)) # 变成(5000,28,28,1) # 初始化生成器并拟合数据 datagen = ImageDataGenerator(rotation_range=360) datagen.fit(X_channels_last) # 2. 一次性生成所有旋转样本,batch_size设为样本总数 # shuffle=False可以保持和原数据的顺序对应(如果不需要顺序可以去掉) gen = datagen.flow(X_channels_last, batch_size=len(X_channels_last), shuffle=False) X_rotated_channels_last = next(gen) # 3. 转换回你需要的通道在前格式 X_rotated = np.transpose(X_rotated_channels_last, (0, 3, 1, 2)) # 验证形状是否一致 print(X_rotated.shape) # 输出:(5000, 1, 28, 28)
方法2:用scikit-image手动遍历旋转(更直观)
如果你觉得Keras的生成器有点绕,也可以用scikit-image的rotate函数直接遍历每张图片处理,逻辑更清晰:
import numpy as np from skimage.transform import rotate X_rotated = np.copy(X) # 复制原数组避免修改原数据 for idx in range(X.shape[0]): # 生成0到360之间的随机旋转角度 random_angle = np.random.uniform(0, 360) # 取出单张图片(形状是(28,28)) single_img = X[idx, 0, :, :] # 旋转图片,preserve_range=True保持像素值范围和原数据一致 rotated_img = rotate(single_img, random_angle, preserve_range=True) # 把旋转后的图片放回结果数组 X_rotated[idx, 0, :, :] = rotated_img # 验证形状 print(X_rotated.shape) # 输出:(5000, 1, 28, 28)
两种方法都能达到你的需求,选哪种看你的习惯:
- 用
ImageDataGenerator的好处是后续如果需要加其他数据增强(比如平移、缩放),直接在初始化时加参数就行,而且和Keras训练流程兼容; - 用
scikit-image的方法更直白,适合快速调试小数据集,不用处理维度转换的细节。
内容的提问来源于stack exchange,提问作者Didier
相关产品推荐
相关产品推荐

