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

调用仅含.pb文件的目标检测模型AWS SageMaker端点报错

AWS SageMaker OD模型调用报错排查解决步骤

核心问题定位

你当前的报错是因为本地导出的目标检测模型不是SageMaker默认TensorFlow推理容器支持的标准SavedModel格式:仅含.pb的冻结计算图(Frozen Graph)缺少必须的variables权重目录,无法被TFServing直接加载。

解决方案步骤

1. 确认正确的SavedModel目录结构

标准可被SageMaker加载的TensorFlow SavedModel必须包含以下结构:

saved_model/
├── saved_model.pb          # 计算图结构文件
└── variables/              # 权重变量目录
    ├── variables.data-00000-of-00001
    └── variables.index

2. 重新导出标准SavedModel

如果你用TensorFlow Object Detection API训练模型:

使用官方导出接口生成标准SavedModel:

from object_detection import export_tflite_graph_lib_tf2

# 替换为你本地的对应路径
pipeline_config = "./pipeline.config"
ckpt_path = "./training_checkpoints/model.ckpt-10000" # 替换为你的最新checkpoint编号
export_dir = "./exported_saved_model"

export_tflite_graph_lib_tf2.export_tflite_model(
    pipeline_config_path=pipeline_config,
    trained_checkpoint_dir=ckpt_path,
    output_directory=export_dir,
    max_detections=100
)

如果你是自定义训练的OD模型:

直接用TF原生接口导出,不要手动执行图冻结操作:

import tensorflow as tf

# 替换为你训练完成的模型实例
trained_od_model = load_your_trained_model()
tf.saved_model.save(trained_od_model, export_dir="./exported_saved_model")

3. 本地验证模型有效性

导出完成后先在本地测试加载,避免上传后再报错:

import tensorflow as tf

loaded_model = tf.saved_model.load("./exported_saved_model")
# 模拟推理输入测试
test_img = tf.random.uniform(shape=[1, 640, 640, 3], minval=0, maxval=255, dtype=tf.uint8)
infer = loaded_model.signatures["serving_default"]
output = infer(test_img)
# 确认输出包含目标检测所需字段
print(output.keys()) # 正常输出应包含detection_boxes、detection_scores、detection_classes等

4. 重新上传到SageMaker并测试

  • 将导出的exported_saved_model目录下的所有文件打包为model.tar.gz,注意打包时不要嵌套多余目录(解压后必须直接看到saved_model.pb和variables目录)
  • 上传到SageMaker多模型端点对应的S3模型存储路径
  • 刷新端点模型缓存后重新调用测试即可

特殊适配方案(仅当无法重新导出模型时使用)

如果受限于训练环境无法重新导出标准SavedModel,可以自定义SageMaker推理脚本,手动加载冻结图:

# inference.py 模型加载逻辑
import tensorflow as tf

def model_fn(model_dir):
    pb_file_path = f"{model_dir}/saved_model.pb"
    with tf.io.gfile.GFile(pb_file_path, "rb") as f:
        graph_def = tf.compat.v1.GraphDef()
        graph_def.ParseFromString(f.read())
    tf.import_graph_def(graph_def, name="")
    session = tf.compat.v1.Session()
    return session

该方案仅兼容TensorFlow 1.x版本的SageMaker推理容器,不推荐生产环境长期使用。

内容的提问来源于stack exchange,提问作者Manjunath Sudheer

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.03 22:24:05