如何在冻结TensorFlow图时避免变量被转换为常量
解决Freeze Graph时保留特定可变变量的问题
我完全理解你的困扰——freeze_graph.py默认会把所有变量一股脑转成常量,但你这个变量需要保持可变状态来支持tf.assign或者喂入数据,官方文档缺失相关说明确实让人头疼。不过别担心,这个工具其实支持通过白名单/黑名单参数来控制变量是否被冻结,我来给你详细说明怎么用:
直接使用freeze_graph的参数
freeze_graph.py有两个关键参数可以控制变量冻结行为:
--exclude_var_list:黑名单,指定哪些变量要排除在冻结之外(也就是保持可变状态),多个变量用逗号分隔。--freeze_var_list:白名单,指定哪些变量需要被冻结成常量,其余变量会保持可变。
举个实际的命令例子,假设你要保留的变量名叫dynamic_input_var,运行freeze_graph时可以这么写:
python freeze_graph.py \ --input_graph=your_raw_graph.pb \ --input_checkpoint=your_model.ckpt \ --output_graph=frozen_with_dynamic_var.pb \ --output_node_names=your_model_output_node \ --exclude_var_list=dynamic_input_var
如果你的变量在命名空间下(比如用tf.variable_scope包裹的),要写完整的变量名,比如model_scope/dynamic_input_var。
小技巧:确认变量名
如果你不确定变量的准确名称,可以用下面的代码列出检查点里的所有变量:
import tensorflow as tf print(tf.train.list_variables("your_model.ckpt"))
手动编写冻结脚本(更灵活)
如果觉得命令行参数不够直观,你也可以自己写Python脚本精确控制冻结逻辑,避免参数使用出错:
import tensorflow as tf # 加载原始图和检查点 with tf.Graph().as_default() as graph: saver = tf.train.import_meta_graph("your_graph.meta") with tf.Session() as sess: saver.restore(sess, "your_model.ckpt") # 筛选出要保留的可变变量(这里替换成你的变量名) keep_variable_names = {"dynamic_input_var"} keep_vars = [v for v in tf.global_variables() if v.name.split(":")[0] in keep_variable_names] # 剩下的变量都是要冻结成常量的 freeze_vars = [v for v in tf.global_variables() if v not in keep_vars] # 执行冻结操作,只把指定变量转成常量 frozen_graph_def = tf.graph_util.convert_variables_to_constants( sess, graph.as_graph_def(), output_node_names=["your_model_output_node"], # 替换成你的输出节点名 variable_names_whitelist=[v.name.split(":")[0] for v in freeze_vars] ) # 保存最终的冻结图 with tf.gfile.GFile("custom_frozen_graph.pb", "wb") as f: f.write(frozen_graph_def.SerializeToString())
这个脚本的优势是你可以用代码逻辑灵活筛选变量,比如按命名空间、变量类型批量选择,比命令行参数更可控。
不管用哪种方法,处理后的模型都会保留你需要的可变变量,既可以用tf.assign赋值,也能在推理前喂入数据,同时把其他不需要可变的变量转成常量来减小模型体积、提升推理速度。
内容的提问来源于stack exchange,提问作者Lukeyb
相关产品推荐
相关产品推荐

