使用freeze_graph冻结LinearClassifier训练的图时触发TypeError
兄弟,我之前在用Estimator训练的模型做冻结移植时也踩过这个一模一样的坑!这个错误的核心原因很明确:你传给保存器的不是Variable对象,而是一个Tensor实例,而tf.train.Saver要求的是变量(或者字符串到变量的字典映射),自然就报错了。
问题本质
LinearClassifier是TensorFlow Estimator API的封装实现,它内部的权重、偏置这些参数都是以Variable的形式存在的,但你直接取"linear/linear_model/0/weights:0"这个节点时,拿到的是变量的输出张量(Tensor),不是变量本身,这就触发了那个TypeError。
靠谱的解决方案
下面给你三种可行的解决办法,按推荐程度排序:
1. 用Estimator自带的SavedModel导出(最省心)
Estimator本身就支持直接导出适合部署的SavedModel格式,完全不需要手动处理变量和冻结,强烈推荐用这种方式:
# 训练完成后,定义输入接收函数 def serving_input_receiver_fn(): # 替换成你训练时用的特征列 feature_spec = tf.feature_column.make_parse_example_spec(your_feature_columns) return tf.estimator.export.build_parsing_serving_input_receiver_fn(feature_spec)() # 导出SavedModel export_dir = "./my_saved_model" estimator.export_saved_model(export_dir_base=export_dir, serving_input_receiver_fn=serving_input_receiver_fn)
之后你可以直接用TensorFlow Lite Converter把这个SavedModel转换成适合移动设备的.tflite格式,一步到位。
2. 手动获取Estimator的变量再保存
如果你一定要手动处理冻结流程,可以先获取模型的所有可训练变量,再传入Saver:
with tf.Session() as sess: # 加载训练好的模型权重 estimator._load_model(sess) # 获取所有可训练变量 trainable_vars = tf.trainable_variables() # 构建变量名到变量的字典(去掉张量名的:0后缀) name_to_var_map = {var.name.split(':')[0]: var for var in trainable_vars} # 初始化Saver并保存检查点 saver = tf.train.Saver(name_to_var_map) saver.save(sess, "./model_ckpt") # 接下来就可以用freeze_graph脚本处理这个检查点了
这里的关键是把变量名的:0后缀去掉,因为变量本身的名字不带这个后缀,:0是张量节点的标识。
3. 直接用graph_util冻结变量
如果你是用tf.graph_util.convert_variables_to_constants来冻结图,要确保传入的是变量列表,而不是张量:
with tf.Session() as sess: estimator._load_model(sess) # 获取所有全局变量 all_vars = tf.global_variables() # 把变量转换为常量,替换成你的模型输出节点名 frozen_graph = tf.graph_util.convert_variables_to_constants( sess, sess.graph_def, output_node_names=["linear/head/predictions/probabilities"] # 按需修改 ) # 保存冻结后的.pb文件 with open("./frozen_model.pb", "wb") as f: f.write(frozen_graph.SerializeToString())
这里的output_node_names需要你确认模型的输出节点名,你可以用TensorBoard打开训练时的日志,查看图结构找到准确的输出节点。
小提示
- 尽量用第一种方法,官方封装的流程最稳定,不容易出错
- 要是不确定节点名,用TensorBoard可视化一下模型图,能帮你快速定位变量和输出节点
内容的提问来源于stack exchange,提问作者kojanck

