在tf.data流水线中调用tf.keras.Model.predict报错ValueError,求排查方案
问题分析与解决
你遇到的ValueError: Unknown graph. Aborting.错误,核心原因是在tf.data.Dataset.map()中使用了model.predict()方法。
为什么会出错?
model.predict()是Keras提供的高层API,它的设计目标是在Eager模式下对批量数据进行预测,内部会创建独立的执行函数和计算图。但tf.data的map()函数运行在TensorFlow的计算图上下文环境中,当你在map里调用predict时,它会尝试创建新的图,这就和当前tf.data所在的图环境产生冲突,导致"未知图"的错误。
修正方案
把model.predict()替换为直接调用模型(即model(img)),这是TensorFlow计算图友好的方式,会把模型的计算逻辑整合到当前tf.data的图中,而不是创建新图。
另外要注意:map()处理的是单个样本,而模型通常期望输入带有批量维度(形状为(batch_size, H, W, C)),所以需要先给单个图像增加批量维度,预测后再去掉多余的维度。
修正后的代码如下:
def load_data(..., model): # code to load an image files = tf.data.Dataset.from_tensor_slices(file_list) images = files.map(load_image_from_file) def pass_image_through_model(img): # 给单个图像增加批量维度 img_with_batch = tf.expand_dims(img, axis=0) # 直接调用模型进行预测 pred = model(img_with_batch) # 去除批量维度,返回单个样本的预测结果 return tf.squeeze(pred, axis=0) dataset = images.map(pass_image_through_model) return dataset
额外提示
如果你的模型是在Eager模式外训练的(比如旧版TensorFlow的图模式),确保在使用前已经正确加载,并且和当前TensorFlow环境的图上下文兼容。直接调用模型的方式在Eager和图模式下都能正常工作,比predict()更适合嵌入到tf.data流水线中。
内容的提问来源于stack exchange,提问作者Monte
相关产品推荐
相关产品推荐

