TensorFlow2调用多个Keras模型predict触发tf.function重追踪警告如何解决?
在项目中以TensorFlow为后端,基于Keras训练了一系列二分类器,输入数据为图像集合,每个二分类器都需要对输入图像执行预测,最终将预测结果保存到CSV文件中。
前几个二分类器调用predict方法预测时没有任何警告,等到第5或第6个二分类器调用predict方法对输入数据预测时,会弹出如下警告:
WARNING:tensorflow:5 out of the last 5 calls to <function Model.make_predict_function..predict_function at 0x2b280ff5c158> triggered tf.function retracing. Tracing is expensive and the excessive number of tracings could be due to (1) creating @tf.function repeatedly in a loop, (2) passing tensors with different shapes, (3) passing Python objects instead of tensors. For (1), please define your @tf.function outside of the loop. For (2), @tf.function has experimental_relax_shapes=True option that relaxes argument shapes that can avoid unnecessary retracing. For (3), please refer to TensorFlow官方文档获取更多详情。
针对警告中列出的三类可能原因,逐一核对情况如下:
- predict方法确实是在for循环中被调用
- 未传入张量,传入的是灰度图像对应的NumPy数组列表,所有图像的宽高尺寸完全一致,唯一可能变化的是批次大小,因为列表中可能只有1张或多张图像
- 传入的参数是NumPy数组列表
经调试确认,警告每次都是在调用predict方法时触发,简化版复现代码如下:
import cv2 as cv import tensorflow as tf from tensorflow.keras.models import load_model # 加载模型 binary_classifiers = [load_model(path) for path in path2models] # 读取图像 images = [# 用OpenCV加载图像] # 对图像执行 resize 和 reshape 操作 my_list = list() for image in images: image_reworked = # 对图像做尺寸调整和维度调整 my_list.append(image_reworked) # 每个模型分别执行预测,警告在此处触发 predictions = [model.predict(x=my_list,verbose=0) for model in binary_classifiers]
尝试定义被tf.function装饰的函数,将预测逻辑放在函数内部,代码如下:
@tf.function def testing(models, faces): return [model.predict(x=faces,verbose=0) for model in models]
运行后报出如下错误:
RuntimeError: Detected a call to
Model.predictinside atf.function. Model.predict is a high-level endpoint that manages its owntf.function. Please move the call toModel.predictoutside of all enclosingtf.functions. Note that you can call aModeldirectly on Tensors inside atf.functionlike:model(x).
由此可知predict方法本身已经封装了tf.function,额外套一层tf.function没有作用,警告本身就来自predict方法内部的tf.function。查询了相关技术社区讨论和官方文档后,仍未找到解决办法。
希望能够消除该警告,同时优化当前程序预测耗时过长的问题,使用的运行环境为:
- Python 3.6.13
- TensorFlow 2.3.0
多次尝试压制predict方法的警告未果后,查阅TensorFlow官方文档发现,TensorFlow默认运行在eager模式下,该模式适合模型测试和调试场景,而模型已经经过多次测试验证,不需要eager模式的调试能力,只需要添加一行代码关闭eager模式即可解决问题:tf.compat.v1.disable_eager_execution()
添加该行代码后,警告不再出现,预测效率也得到了提升。
内容的提问来源于stack exchange,提问作者Simone Starace

