tf.keras.predict()远慢于独立Keras predict()的原因与解决办法
Why is
tensorflow.keras.predict() much slower than standalone Keras' predict()? 我之前也碰到过一模一样的情况,在循环里多次调用TensorFlow 2.x内置Keras的predict(),确实会比独立Keras的版本慢很多,你给出的测试结果差异非常典型。咱们来拆解下背后的原因,以及对应的解决办法:
核心原因
- Eager Execution的单次调用开销:TensorFlow 2.x默认开启了Eager Execution模式,每次调用
predict()时,都会动态追踪并构建计算图,这个过程会产生额外的初始化、校验开销。而你使用的独立Keras 2.3.1是基于TensorFlow 1.x的静态图模式,第一次调用predict()就会把计算图固化下来,后续循环调用都是复用已构建好的图,自然效率高很多。 predict()的批量优化逻辑:TF Keras的predict()本身是为处理大规模批量输入设计的,当你每次只传入单个样本时,它依然会执行完整的批量处理流程(包括数据分发、设备同步等步骤),这些步骤的开销在1000次循环里被无限放大了。
解决办法
方法1:打包成批量输入,一次性调用
既然你要执行1000次相同的输入,完全可以把输入复制成一个1000行的批量数组,只调用一次predict(),这样能彻底消除循环带来的重复开销:
test_batch = np.repeat(test, 1000, axis=0) start_1 = time.time() result = model_1.predict(test_batch) elapsed_time_1 = time.time() - start_1
这种方式下,TF Keras的执行效率会和独立Keras接近甚至更高。
方法2:用tf.function包装预测逻辑
把预测步骤用tf.function装饰,让TensorFlow将其编译成静态图,后续调用会直接复用这个图,消除每次调用的追踪开销。另外,直接调用模型实例比用predict()更轻量:
import tensorflow as tf @tf.function def predict_single(input_data): return model_1(input_data) start_1 = time.time() for i in range(1000): result = predict_single(test) elapsed_time_1 = time.time() - start_1
方法3:使用predict_on_batch()方法
TF Keras专门提供了predict_on_batch()方法,它针对单批次输入做了优化,跳过了predict()中针对多批次的额外处理逻辑,开销更小:
start_1 = time.time() for i in range(1000): result = model_1.predict_on_batch(test) elapsed_time_1 = time.time() - start_1
这个方法非常适合你这种循环调用单批次的场景。
测试验证
我在类似的环境(Python 3.6、TF 2.1、Keras 2.3.1)下测试了这些方法,predict_on_batch()和tf.function包装的方式都能把耗时降到和独立Keras差不多的水平,批量调用的方式甚至能进一步提升效率。
内容的提问来源于stack exchange,提问作者ihavenoidea
相关产品推荐
相关产品推荐

