如何在Keras损失函数中调用预训练TensorFlow网络?
问题分析与解决方案
你的这个错误其实很典型——核心问题在于Keras损失函数里的y_true和y_pred是TensorFlow张量对象,而不是numpy数组,但你当前的DeepGaze.get_saliency_map方法是通过tf.Session.run()来计算结果的,这个方法只接受numpy数组、Python标量这类“具体数值”作为feed输入,根本不认识TF张量,所以才会抛出TypeError。
除此之外,你还犯了一个关键设计错误:把预训练的DeepGaze模型加载到了一个独立的TF计算图里,还单独维护了一个会话。Keras(TF后端)的整个训练流程是基于默认计算图和自身管理的会话运行的,跨图的张量无法在另一个会话里被处理,这相当于两个完全独立的计算环境,自然没法互通。
怎么解决?
我们需要把预训练模型完全整合到Keras的计算图中,让损失函数的所有操作都变成纯张量运算(这样才能支持反向传播,同时避免类型错误)。具体步骤如下:
1. 修改DeepGaze类,整合到Keras的计算图与会话
把原来独立的图和会话去掉,直接在Keras的默认图里加载预训练模型,并用Keras的会话来恢复参数:
import tensorflow as tf from keras import backend as K import os class DeepGaze(object): CHECK_POINT = os.path.join(os.path.dirname(__file__), 'DeepGazeII.ckpt') def __init__(self): print('Loading Deep Gaze II...') # 直接在Keras的默认计算图中加载模型 # 1. 读取meta图结构 graph_def = tf.GraphDef() with tf.gfile.GFile(f"{self.CHECK_POINT}.meta", 'rb') as f: graph_def.ParseFromString(f.read()) # 2. 定义输入张量(要和你的y_true/y_pred形状匹配,比如[None, H, W, 3]) self.input_placeholder = tf.placeholder(tf.float32, shape=[None, None, None, 3]) # 3. 导入预训练图,映射输入到我们的占位符,同时获取输出张量 # 注意:input_map的键要和预训练模型中input_tensor的名称完全一致 self.saliency_output = tf.import_graph_def( graph_def, input_map={'input_tensor:0': self.input_placeholder}, return_elements=['log_density_wo_centerbias:0'] )[0] # 4. 用Keras的会话恢复模型参数 self.session = K.get_session() saver = tf.train.Saver() saver.restore(self.session, self.CHECK_POINT) print('Deep Gaze II Loaded') def get_saliency_map(self, input_tensor): # 现在直接返回张量运算结果,不需要session.run return tf.identity(self.saliency_output, name='saliency_map')
2. 修改自定义损失函数,使用张量运算
现在get_saliency_map返回的是TF张量,完全可以和y_true/y_pred进行张量运算,不需要再转换为numpy数组:
def custom_loss_func(y_true, y_pred): # 直接传入张量,得到显著性图张量 sal_true = deep_gaze.get_saliency_map(y_true) sal_pred = deep_gaze.get_saliency_map(y_pred) # 计算MSE损失,全程都是张量操作 return K.mean(K.square(sal_true - sal_pred))
关键注意事项
- 形状匹配:确保预训练模型的输入形状和
y_true/y_pred的形状完全一致,如果不一致,需要用tf.reshape或tf.image.resize先调整形状,比如:y_true_resized = tf.image.resize_images(y_true, [H, W]) # H/W是预训练模型要求的输入尺寸 - 禁止在损失函数中调用
session.run():Keras的损失函数是用来构建计算图的,所有操作必须是张量之间的运算,这样才能自动生成反向传播的梯度计算链路。一旦调用session.run(),就会把张量转换成numpy数组,断开计算图,不仅会报错,还会导致无法训练。 - 共享会话:一定要用
K.get_session()获取Keras的会话,不要自己新建会话,避免多个会话之间的参数冲突。
这样修改后,你的损失函数就能正常运行,预训练模型的计算也会被整合到Keras的训练流程中,支持反向传播。
内容的提问来源于stack exchange,提问作者shaharhoch
相关产品推荐
相关产品推荐

