如何冻结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
相关产品推荐
相关产品推荐

