You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.27 07:35:19