TensorFlow模型训练时Keras速度远低于Estimator的原因是什么?
训练速度差异的核心原因
- 参数使用错误:你的
get_dataset()已经对数据做了batch(batch_size)操作,但在Keras的fit调用中又传入了batch_size=batch_size参数,会导致Keras对已经分批的数据集做二次分批,单步处理的样本量是预期的256倍,计算量陡增。 - 无限数据集适配问题:你定义的
data_generator是无限循环生成数据,没有终止条件。Estimator本身默认支持无限输入,会按照训练逻辑直接执行;但Keras的fit默认会尝试遍历数据集来获取总步数,探测不到终止条件就会持续产生冗余的预读取开销,这也是你看到Keras输出里步数显示Unknown的核心原因。 - 执行模式的默认差异:Estimator的训练逻辑默认全部包装为静态计算图执行,没有eager模式的开销;而Keras在TF2.x中虽然默认也会做图编译,但对自定义输入、自定义损失的适配过程会产生额外的追踪开销,尤其是和旧版
tf.compat.v1API混用的时候,编译优化的覆盖率更低。 - FeatureColumn的版本兼容问题:TF2.6对
tf.feature_column的底层实现做了调整,存在已知的性能退化问题,所以你会看到两个训练方式在2.6版本下速度都比2.5低很多。
Keras训练速度优化方案
- 修正fit调用参数:
- 移除
fit中的batch_size参数,避免二次分批 - 显式指定
steps_per_epoch参数,比如你要每轮跑100步就设置steps_per_epoch=100,避免Keras探测无限数据集的长度开销,修改后的调用示例:
model.fit(get_dataset(), epochs=1, steps_per_epoch=100) - 移除
- 替换性能较低的组件:将
DenseFeatures+embedding_column的组合替换为原生的tf.keras.layers.Embedding层,DenseFeatures为了兼容旧版特征列逻辑做了大量格式转换,处理大hash桶嵌入时性能比原生Embedding低30%以上。 - 开启编译优化:在
model.compile时添加jit_compile=True参数(TF2.5及以上版本支持),开启XLA编译优化,能进一步提升静态图的执行效率。 - 减少冗余开销:如果不需要精细的进度展示,可以在
fit中设置verbose=0关闭进度条输出,减少日志和界面渲染的开销。 - 版本适配:如果没有必须使用TF2.6的需求,可以降级到TF2.5版本,避免feature_column的性能退化问题。
内容的提问来源于stack exchange,提问作者PWZER
相关产品推荐
相关产品推荐

