You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何加载.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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.21 07:33:38