Keras使用ImageDataGenerator做数据增强时无法复现运行结果怎么办
Keras ImageDataGenerator 增强场景下结果复现缺失配置说明
你遗漏的核心配置和修复方案如下:
- 强制数据加载使用单线程单进程模式,避免多线程下随机数生成顺序不可控。
ImageDataGenerator默认支持多并行worker处理数据增强,并行场景下随机数的调用顺序无法和全局种子的生成序列对齐,就算设置了全局种子也无法保证每次运行的随机变换顺序一致。请在model.fit()的参数中新增workers=1, use_multiprocessing=False,修改后fit代码如下:
my_model.fit( generator.flow(X_train, y_train, batch_size=batch_size, shuffle=False, seed=seed_num), validation_data=(X_val, y_val), callbacks=callbacks, epochs=epochs, shuffle=False, workers=1, use_multiprocessing=False )
- 新增CPU线程数限制,避免多线程运算带来的浮点误差累积和执行顺序差异。CPU多线程并行计算时,运算执行顺序不固定会带来微小的浮点误差,多轮迭代后会累积成可观测的结果差异。在导入TensorFlow之后添加以下配置:
tf.config.threading.set_inter_op_parallelism_threads(1) tf.config.threading.set_intra_op_parallelism_threads(1)
如果使用兼容v1的会话模式,补充以下配置固定会话参数:
session_conf = tf.compat.v1.ConfigProto(intra_op_parallelism_threads=1, inter_op_parallelism_threads=1) sess = tf.compat.v1.Session(graph=tf.compat.v1.get_default_graph(), config=session_conf) tf.compat.v1.keras.backend.set_session(sess)
- 补充Keras后端全局种子设置,覆盖Keras内部随机数生成器的种子配置,在导入backend后添加:
tf.keras.utils.set_random_seed(seed_num)
以上配置全部添加后,开启数据增强的训练结果即可复现。如果你不愿修改配置,也可以升级到TensorFlow 2.6及以上版本,该版本修复了2.4中ImageDataGenerator随机种子不生效的已知缺陷。
内容的提问来源于stack exchange,提问作者Kate
相关产品推荐
相关产品推荐

