如何使用Keras对无法全量载入内存的大规模数据集训练自编码器
问题根因
当前实现存在4个核心错误,直接导致模型无法遍历全量数据集完成训练:
- 错误调用
next(iter(X_train)):该操作仅会从数据集中提取第一个批次的样本,后续迭代不会继续读取剩余数据,模型跑完这个单批次就会终止训练。 - 输入输出数据集未绑定:为原始图像、标注图像分别创建独立的
tf.data.Dataset实例后,直接分别传入fit的x、y参数时,两个独立迭代器无法保证样本一一对应,还会出现迭代步数不匹配的问题。 - 参数冲突:在
image_dataset_from_directory中已经设置batch_size=128完成数据集分块,此时再给fit传入batch_size=10完全无效——Keras接收已分batch的tf.data数据集时,会自动忽略fit方法内的batch_size参数。 - 数据流水线缺失优化:未配置预加载、缓存、打乱逻辑,既会拖慢训练速度,也会影响模型收敛效果。
大规模数据集正确训练方案
对于无法全量载入内存的图像数据集,使用tf.data流水线做流式读取是Keras官方推荐方案,不需要手动做map转存,实现步骤如下:
- 分别加载输入、标注数据集时先关闭shuffle,保证两个数据集的样本顺序一致
- 用
tf.data.Dataset.zip将两个数据集绑定为(输入样本, 标注)的标准结构 - 对绑定后的数据集做shuffle、缓存、预取优化,提升IO读取效率
- 训练时直接传入处理好的数据集对象,不需要额外指定batch_size
参考实现代码:
import tensorflow as tf from tensorflow import keras # 统一配置参数,不要重复定义批次大小 image_size_ = 256 # 替换为实际使用的图像尺寸 batch_size = 128 epochs_number = 1 # 加载训练集输入,先关闭shuffle保证和标注顺序对齐 train_X = keras.utils.image_dataset_from_directory( directory=train_image_folder, labels=None, label_mode=None, batch_size=batch_size, image_size=(image_size_, image_size_), shuffle=False ) # 加载训练集标注 train_Y = keras.utils.image_dataset_from_directory( directory=train_annot_folder, labels=None, label_mode=None, batch_size=batch_size, image_size=(image_size_, image_size_), shuffle=False ) # 加载验证集输入 val_X = keras.utils.image_dataset_from_directory( directory=val_image_folder, labels=None, label_mode=None, batch_size=batch_size, image_size=(image_size_, image_size_), shuffle=False ) # 加载验证集标注 val_Y = keras.utils.image_dataset_from_directory( directory=val_annot_folder, labels=None, label_mode=None, batch_size=batch_size, image_size=(image_size_, image_size_), shuffle=False ) # 绑定输入和标注为标准数据集结构 train_ds = tf.data.Dataset.zip((train_X, train_Y)) val_ds = tf.data.Dataset.zip((val_X, val_Y)) # 流水线优化 # shuffle缓冲区大小可根据显存/内存调整,建议不小于单batch大小的5倍 train_ds = train_ds.shuffle(buffer_size=1000) # 内存足够时直接用.cache()缓存预处理后的数据,内存不足则传入磁盘路径做持久化缓存,例:.cache("./train_cache.tmp") train_ds = train_ds.cache().prefetch(buffer_size=tf.data.AUTOTUNE) val_ds = val_ds.cache().prefetch(buffer_size=tf.data.AUTOTUNE) # 加载模型 autoencoder = keras.models.load_model("Model") # 启动训练,无需传入batch_size参数 autoencoder.fit( x=train_ds, epochs=epochs_number, validation_data=val_ds )
注意事项
- 必须保证输入文件夹和标注文件夹内的文件命名、排序逻辑完全一致,否则zip后会出现样本和标注错配的问题,训练完全无效。
- 如果训练时出现显存不足,直接调小全局
batch_size参数即可,不要在fit方法内重复设置batch_size。 - 若内存不足以缓存全量预处理数据,给
cache()传入本地磁盘路径即可,首次训练会生成缓存文件,后续训练启动速度会大幅提升。 - 不要完全关闭shuffle,否则样本顺序固定会严重影响模型收敛效果。
内容的提问来源于stack exchange,提问作者Nuno Pessanha Santos
相关产品推荐
相关产品推荐

