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

如何冻结TensorFlow模型?测试阶段固定RNN权重仅获预测结果的方法

如何冻结TensorFlow模型并在测试阶段固定权重

嘿,我来帮你搞定这个问题——在实际项目里,不管是测试时不想动权重,还是要导出冻结模型部署,这俩需求都挺常见的,分两种场景给你讲清楚:

一、测试阶段阻止权重更新(代码层面快速实现)

你有model.py里的RNN类,在main.py调用测试,不想让权重更新?其实不用复杂操作,核心就是让模型知道"现在是推理模式",同时关闭梯度追踪就行:

方法1:调用模型时指定training=False(最推荐)

TensorFlow的Keras层(包括你自定义的RNN类)都支持training参数,这个参数不仅会让Dropout、BatchNorm这类层切换到推理模式,还会自动阻止权重更新。代码大概是这样:

# main.py 测试阶段代码
from model import YourRNNModel
import tensorflow as tf

# 实例化模型并加载训练好的权重
model = YourRNNModel()
model.load_weights("你的权重文件路径.h5")  # 或者用SavedModel加载

# 测试预测
test_input = tf.random.normal([1, 10, 32])  # 替换成你的测试输入格式
predictions = model(test_input, training=False)  # 关键就在这里!

如果你的RNN类里自定义了call方法,记得要把training参数传递给内部的层(比如LSTM层),比如:

# model.py 里的RNN类
class YourRNNModel(tf.keras.Model):
    def __init__(self):
        super().__init__()
        self.lstm = tf.keras.layers.LSTM(64)
        self.dense = tf.keras.layers.Dense(10)
    
    def call(self, inputs, training=False):
        x = self.lstm(inputs, training=training)  # 把training传进去
        return self.dense(x)

方法2:手动设置变量不可训练(更彻底)

如果你担心不小心触发训练逻辑,可以加载权重后直接把所有可训练变量设为不可训练:

# 加载权重后执行
for var in model.trainable_variables:
    var.trainable = False

# 之后不管怎么调用模型,权重都不会更新
predictions = model(test_input, training=False)

二、导出冻结模型(适合部署场景)

如果需要把模型导出成一个独立的、权重和图绑定的文件(比如.pb格式),方便部署到生产环境或者离线使用,这才是真正意义上的"冻结模型":

方式1:保存为SavedModel格式(TensorFlow 2.x首选)

SavedModel是TensorFlow的标准持久化格式,已经包含了冻结的权重和模型结构,直接保存和加载就行:

# 训练完成后保存
model.save("saved_model_dir")  # 会生成一个文件夹,包含模型文件

# 测试/部署时加载
loaded_model = tf.keras.models.load_model("saved_model_dir")
predictions = loaded_model(test_input, training=False)

这个方式最省心,不需要手动处理图和变量,TensorFlow会帮你搞定一切。

方式2:转换为传统冻结图(.pb文件)

如果需要适配旧的部署框架,比如TensorFlow Lite或者一些嵌入式设备,可以把SavedModel转成冻结图(权重嵌入到计算图里):

import tensorflow as tf

# 加载SavedModel
loaded = tf.saved_model.load("saved_model_dir")
infer = loaded.signatures["serving_default"]  # 获取默认的推理签名

# 获取输入输出张量名称
input_name = infer.inputs[0].name.split(":")[0]
output_name = infer.outputs[0].name.split(":")[0]

# 转换为冻结图
graph = tf.compat.v1.get_default_graph()
with tf.compat.v1.Session(graph=graph) as sess:
    frozen_graph_def = tf.compat.v1.graph_util.convert_variables_to_constants(
        sess,
        graph.as_graph_def(),
        [output_name]
    )

# 保存冻结图到文件
with open("frozen_model.pb", "wb") as f:
    f.write(frozen_graph_def.SerializeToString())

加载冻结图推理的代码:

def load_frozen_model(pb_path):
    with tf.io.gfile.GFile(pb_path, "rb") as f:
        graph_def = tf.compat.v1.GraphDef()
        graph_def.ParseFromString(f.read())
    
    with tf.Graph().as_default() as graph:
        tf.import_graph_def(graph_def, name="")
    
    # 替换成你实际的输入输出张量名称
    input_tensor = graph.get_tensor_by_name(f"{input_name}:0")
    output_tensor = graph.get_tensor_by_name(f"{output_name}:0")
    return graph, input_tensor, output_tensor

# 调用推理
graph, input_tensor, output_tensor = load_frozen_model("frozen_model.pb")
with tf.compat.v1.Session(graph=graph) as sess:
    result = sess.run(output_tensor, feed_dict={input_tensor: test_input})

最后总结一下

  • 只是自己测试用:**直接加training=False**就够了,简单高效,还能保证层的推理行为正确。
  • 需要部署:优先用SavedModel,兼容性好;如果必须用冻结图,再用第二种转换方式。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.19 08:33:06