如何加载.ckpt预训练模型,通过SavedModelBuilder转存为protobuf且无需声明tf.Variables
解决预训练ResNetV2-50转SavedModel的方案
嘿,我之前刚好帮朋友处理过类似的预训练模型转SavedModel的需求,针对你的场景(用Go部署、不想手动声明变量),给你一套完整的可行步骤:
核心思路
预训练的resnet_v2_50.ckpt是基于TensorFlow Slim实现的,所以我们可以直接用Slim提供的预定义网络结构重建计算图,自动复用里面的所有变量(不用手动写tf.Variables),然后加载ckpt权重,最后用SavedModelBuilder导出成protobuf格式。
完整代码示例
import tensorflow as tf from tensorflow.contrib import slim from tensorflow.contrib.slim.nets import resnet_v2 # 1. 定义模型输入(匹配ResNetV2默认格式:224x224x3的RGB图像,像素值0-255) input_tensor = tf.placeholder(tf.float32, shape=[None, 224, 224, 3], name='input_image') # 2. 用Slim自动构建ResNetV2-50计算图,内部已包含所有所需变量 with slim.arg_scope(resnet_v2.resnet_arg_scope()): logits, _ = resnet_v2.resnet_v2_50(input_tensor, num_classes=1000, is_training=False) # 若你的场景需要调整分类数,可修改num_classes,之后按需微调权重即可 predictions = tf.nn.softmax(logits, name='predictions') # 3. 加载预训练ckpt的权重 saver = tf.train.Saver() with tf.Session() as sess: # 初始化变量(Slim已定义变量,这里初始化后用ckpt覆盖权重) sess.run(tf.global_variables_initializer()) saver.restore(sess, './resnet_v2_50.ckpt') # 替换为你的ckpt文件路径(无需加后缀) # 4. 用SavedModelBuilder导出为protobuf格式 builder = tf.saved_model.builder.SavedModelBuilder('./resnet_v2_50_savedmodel') # 定义输入输出签名(Go端调用时需要通过签名定位节点) input_signature = tf.saved_model.utils.build_tensor_info(input_tensor) output_signature = tf.saved_model.utils.build_tensor_info(predictions) prediction_signature = ( tf.saved_model.signature_def_utils.build_signature_def( inputs={'input_image': input_signature}, outputs={'predictions': output_signature}, method_name=tf.saved_model.signature_constants.PREDICT_METHOD_NAME)) # 添加签名并完成保存 builder.add_meta_graph_and_variables( sess, [tf.saved_model.tag_constants.SERVING], signature_def_map={ 'predict_images': prediction_signature }) builder.save()
关键细节说明
- 无需手动声明变量的原因:TensorFlow Slim的
resnet_v2_50函数内部已经封装了所有卷积、批归一化等模块的变量定义,调用函数即可自动创建所有所需变量,省去手动编写大量变量代码的麻烦。 - 输入格式适配:预训练ResNetV2-50默认输入是224x224的RGB图像,若你需要适配不同尺寸,可修改
input_tensor的shape,但建议微调权重以保证推理效果。 - 签名的作用:
predict_images这个签名是给Go端调用用的,Go代码中可以通过该签名快速定位输入输出节点,实现模型推理。
注意事项
- 确保安装依赖:
pip install tensorflow tensorflow-slim(若使用TensorFlow 2.x,需改用tf.compat.v1兼容模式编写代码) - 你的
resnet_v2_50.ckpt必须包含三个配套文件:resnet_v2_50.ckpt.data-00000-of-00001、resnet_v2_50.ckpt.index、resnet_v2_50.ckpt.meta,且需放在同一目录下。 - 导出后可通过以下代码验证SavedModel是否正常:
import tensorflow as tf with tf.Session() as sess: tf.saved_model.loader.load(sess, [tf.saved_model.tag_constants.SERVING], './resnet_v2_50_savedmodel') # 测试推理节点是否可正常获取 input_tensor = sess.graph.get_tensor_by_name('input_image:0') output_tensor = sess.graph.get_tensor_by_name('predictions:0') # 可传入测试图像进行推理验证
内容的提问来源于stack exchange,提问作者mofury
相关产品推荐
相关产品推荐

