如何将Checkpoint文件转换为TensorFlow.js?新手求详细分步指导
从TensorFlow Checkpoint转TensorFlow.js 分步指南
一、准备环境
- 安装TensorFlow 1.x(因为你用的是
tf.compat.v1,推荐1.15.5版本):pip install tensorflow==1.15.5 - 安装TensorFlow.js转换器:
pip install tensorflowjs
二、处理转换代码
1. 代码写在哪里?
这段代码是Python脚本,不能直接在命令提示符里运行。你需要:
- 新建一个文本文件,命名为
convert_to_pb.py - 把你提供的代码粘贴进去,再根据实际情况修改参数。
2. 修改代码中的关键参数
(1)meta_path:指定.meta文件路径
Checkpoint文件夹里会有一个后缀为.meta的文件,比如model.ckpt-1000.meta(数字是训练步数)。你需要把meta_path改成这个文件的实际路径:
# 示例:如果.meta文件是./newcheckpoint/model.ckpt-1000.meta meta_path = './newcheckpoint/model.ckpt-1000.meta'
找不到的话,直接打开
./newcheckpoint文件夹,找带.meta后缀的文件即可。
(2)output_node_names:指定模型输出节点名称
这是新手最容易卡壳的地方,你需要找到模型最后输出结果的节点名称:
- 方法一:查看原训练代码,找模型最后输出张量的
name属性。比如训练时写了y_pred = tf.nn.softmax(logits, name='predictions'),那输出节点就是['predictions']。 - 方法二:用代码打印所有节点名称,找到输出节点:
在你的代码里,saver.restore(sess, ...)之后添加这段代码:
运行脚本后,在控制台里找和输出相关的名称(比如# 打印所有节点名称,方便找输出节点 for node in tf.get_default_graph().as_graph_def().node: print(node.name)output、predictions、logits等),把它放进output_node_names列表里。
修改后的完整代码示例:
import tensorflow.compat.v1 as tf # 修改为你的.meta文件实际路径 meta_path = './newcheckpoint/model.ckpt-1000.meta' # 修改为你的模型输出节点名称 output_node_names = ['predictions'] with tf.Session() as sess: # 恢复图结构 saver = tf.train.import_meta_graph(meta_path) # 加载权重 saver.restore(sess, tf.train.latest_checkpoint('./newcheckpoint/')) # 可选:打印所有节点名称,找输出节点 # for node in tf.get_default_graph().as_graph_def().node: # print(node.name) # 冻结图(把变量转成常量) frozen_graph_def = tf.graph_util.convert_variables_to_constants( sess, sess.graph_def, output_node_names) # 保存冻结后的.pb文件 # 确保freeze文件夹存在,否则先手动创建 with open('./freeze/output_graph.pb', 'wb') as f: f.write(frozen_graph_def.SerializeToString())
三、生成冻结图(.pb文件)
- 先手动创建
freeze文件夹(如果代码里的路径是./freeze) - 打开命令提示符,进入脚本所在的文件夹,运行:
运行成功后,python convert_to_pb.py./freeze文件夹里会生成output_graph.pb文件。
四、把.pb文件转成TensorFlow.js格式
在命令提示符里运行tensorflowjs_converter命令,注意替换参数:
tensorflowjs_converter --input_format=tf_frozen_model --output_node_names="predictions" ./freeze/output_graph.pb ./tfjs_model
--output_node_names的值要和之前代码里的output_node_names一致(多个节点用逗号分隔)./tfjs_model是输出文件夹,运行后会生成model.json和多个权重文件,这就是TensorFlow.js可以直接加载的模型文件。
内容的提问来源于stack exchange,提问作者Christopher Koh
相关产品推荐
相关产品推荐

