TensorFlow1.15.x环境下Keras训练进度条异常、epoch重复问题求助
TensorFlow 1.15.x 下fit训练epoch重复打印问题修复方案
该异常是TensorFlow 1.15内置Keras的进度条组件已知bug,仅为日志输出异常,不会影响模型训练、验证的计算结果,也没有真的重复运行epoch,只是动态进度条的计数逻辑出错导致重复打印。可通过以下方案在不更换1.15.x版本的前提下修复:
- 方案1:修改steps参数为整数,关闭浮点数step逻辑
你当前代码中steps_per_epoch和validation_steps直接使用除法得到的是浮点数,1.15版本的进度条对浮点数step的处理存在计数异常,将两个参数改为整数即可修复,示例修改如下:history = model.fit( generator_train, epochs=3, # 改为整除得到整数,也可以用int()强制转换,需要遍历所有样本可改用math.ceil() validation_steps=generator_valid.samples // generator_valid.batch_size, steps_per_epoch=generator_train.samples // generator_train.batch_size, validation_data=generator_valid, ) - 方案2:调整日志输出级别,禁用动态进度条
直接在model.fit中添加verbose=2参数,每个epoch仅输出一行日志,跳过动态进度条的渲染逻辑,从根源避免进度条计数bug,示例如下:
该方案不需要修改核心逻辑,仅调整输出形式,不会影响训练效果。history = model.fit( generator_train, epochs=3, validation_steps=generator_valid.samples / generator_valid.batch_size, steps_per_epoch=generator_train.samples / generator_train.batch_size, validation_data=generator_valid, verbose=2 ) - 方案3:替换为独立版本Keras
安装兼容TF1.15的独立Keras 2.3.1版本:pip install keras==2.3.1
代码中将所有tf.keras的导入替换为keras即可,独立版Keras不存在该进度条bug,接口完全兼容无需额外修改核心逻辑。 - 方案4:运行时打补丁修复内置进度条bug
如果不想修改现有逻辑,可以在训练代码前添加如下猴子补丁,修复内置Progbar类的计数错误:import tensorflow as tf from tensorflow.python.keras.utils import generic_utils # 修复进度条计数bug def patched_progbar_update(self, current, values=None, force=False): if self._seen_so_far == 0: self._start = generic_utils.time.time() self._seen_so_far = current if values is not None: for k, v in values: if k not in self._values: self._values[k] = [v * (current - self._last_update), current - self._last_update] else: self._values[k][0] += v * (current - self._last_update) self._values[k][1] += current - self._last_update self._last_update = current if (current >= self.target and self.target is not None) or force: self._finalize_state() return True return False tf.keras.utils.Progbar.update = patched_progbar_update
内容的提问来源于stack exchange,提问作者GarvielLoken
相关产品推荐
相关产品推荐

