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

如何在TensorFlow 2中用Keras Lambda层封装冻结图?

在TensorFlow 2.x中正确封装冻结PB模型为Keras Model

你遇到的这个错误,核心原因是TensorFlow 2.x的执行模式和1.x有本质区别:TF2默认启用即刻执行(Eager Execution),而且Keras的predict方法会自动将模型逻辑转换为tf.function(图模式优化),但你原来的代码直接沿用了TF1.x的tf.Session和Graph张量,这些张量属于旧的Graph上下文,无法在TF2的函数式执行环境中正常工作,所以抛出了"Graph tensor"相关的错误。

下面是适配TF2.x的完整解决方案,分为三个步骤:

1. 生成测试用的冻结图(兼容TF2)

这部分代码可以稍作调整,用TF2的API来生成冻结图,当然原来的代码用tf.compat.v1也能运行,这里给出更贴合TF2的写法:

import numpy as np
import tensorflow as tf

def dump_model():
    # 在TF2中,我们可以直接构建计算图,然后导出为GraphDef
    x = tf.compat.v1.placeholder(tf.float32, shape=(None, 123), name='x')
    c = tf.constant(100, dtype=tf.float32, name='C')
    y = tf.multiply(x, c, name='y')
    z = tf.add(y, x, name='z')
    
    # 获取当前图的GraphDef并序列化保存
    graph_def = tf.compat.v1.get_default_graph().as_graph_def()
    with tf.io.gfile.GFile("tmp_net.pb", "wb") as f:
        raw = graph_def.SerializeToString()
        print(type(raw), len(raw))
        f.write(raw)

dump_model()

2. 加载冻结图并封装为Keras Model(TF2正确方式)

这里我们不再使用tf.Session,而是用TF2的API将冻结图导入到当前执行上下文,然后通过tf.keras.layers.Input和tf.keras.layers.Lambda正确关联输入输出:

import tensorflow as tf

# 1. 读取并解析冻结图
with tf.io.gfile.GFile("./tmp_net.pb", 'rb') as f:
    graph_def = tf.compat.v1.GraphDef()
    graph_def.ParseFromString(f.read())

# 2. 将冻结图导入到当前TF2上下文,指定输入输出的映射
# 这里我们先定义Keras的输入层,作为模型的入口
input_x = tf.keras.layers.Input(shape=(123,), dtype=tf.float32, name='x')

# 导入GraphDef,将输入层的张量映射到冻结图的输入占位符
# 返回的是冻结图的输出张量列表
[y_tensor, z_tensor] = tf.import_graph_def(
    graph_def,
    input_map={'x:0': input_x},
    return_elements=['y:0', 'z:0'],
    name=''
)

# 3. 构建Keras Model
base_model = tf.keras.Model(inputs=input_x, outputs=[y_tensor, z_tensor])

# 可选:打印模型结构,确认输入输出正确
base_model.summary()

3. 测试模型

现在可以正常使用Keras的predict方法测试:

import numpy as np

y_out, z_out = base_model.predict(np.ones((3, 123), dtype=np.float32))
print(y_out.shape, z_out.shape)  # 输出:((3, 123), (3, 123))

为什么这个方法能解决问题?

  • 我们用TF2的tf.import_graph_def直接将冻结图的计算逻辑和Keras的输入层绑定,避免了旧的tf.Session和Graph上下文的冲突。
  • 导入后的输出张量属于当前TF2的执行上下文,能被tf.function正确处理,所以predict方法可以正常工作。
  • 整个流程完全适配TF2的即刻执行和函数式API特性,没有Graph张量泄漏的问题。

内容的提问来源于stack exchange,提问作者GPhilo

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.08 13:17:27