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

TF1格式保存的TF2 NCF模型在SageMaker部署报错:无SavedModel包

解决SageMaker部署TF1 Checkpoint模型报错"no SavedModel bundles found"的问题

问题根源

你当前保存的是TensorFlow 1.x的Checkpoint(ckpt)格式文件,而Amazon SageMaker的TensorFlowModel类(搭配TF2.x框架版本)默认要求模型以SavedModel格式提供——这是TensorFlow官方标准的服务部署格式,包含saved_model.pb和variables文件夹,而非ckpt的.meta/.index/.data文件。


解决方案步骤

1. 将TF1 Checkpoint转换为SavedModel格式

你需要重新加载训练好的ckpt模型,然后导出为SavedModel。以下是适配NCF模型的转换代码(需确保模型结构与训练时完全一致):

import tensorflow as tf
from tensorflow.python.saved_model import signature_constants
from tensorflow.python.saved_model import tag_constants

# ----------------------
# 第一步:复制训练时的NCF模型结构
# ----------------------
def build_ncf_model(user_num, item_num, latent_dim=64):
    # 用户输入层
    user_input = tf.compat.v1.placeholder(tf.int32, shape=(None, 1), name='user_id')
    # 物品输入层
    item_input = tf.compat.v1.placeholder(tf.int32, shape=(None, 1), name='item_id')
    
    # 嵌入层(与训练时参数一致)
    user_embedding = tf.compat.v1.layers.embedding_lookup(
        tf.Variable(tf.random.normal([user_num, latent_dim])), user_input
    )
    item_embedding = tf.compat.v1.layers.embedding_lookup(
        tf.Variable(tf.random.normal([item_num, latent_dim])), item_input
    )
    
    # NCF核心结构(GMF+MLP,示例简化,需与训练时完全匹配)
    gmf = tf.multiply(tf.squeeze(user_embedding), tf.squeeze(item_embedding))
    mlp = tf.concat([tf.squeeze(user_embedding), tf.squeeze(item_embedding)], axis=1)
    mlp = tf.compat.v1.layers.dense(mlp, units=128, activation='relu')
    concat = tf.concat([gmf, mlp], axis=1)
    output = tf.compat.v1.layers.dense(concat, units=1, activation='sigmoid', name='prediction')
    
    return user_input, item_input, output

# ----------------------
# 第二步:加载ckpt并导出SavedModel
# ----------------------
# 替换为你训练时的用户数、物品数
USER_NUM = 10000
ITEM_NUM = 20000

# 禁用eager execution,适配TF1.x环境
tf.compat.v1.disable_eager_execution()
sess = tf.compat.v1.Session()

# 构建模型
user_input, item_input, output = build_ncf_model(USER_NUM, ITEM_NUM)
# 初始化变量
sess.run(tf.compat.v1.global_variables_initializer())

# 加载ckpt文件
saver = tf.compat.v1.train.Saver()
saver.restore(sess, "./00000000/model.ckpt")  # 你的ckpt文件路径

# 导出为SavedModel
export_dir = "./saved_model/1"  # 1是版本号,SageMaker会识别该目录
builder = tf.compat.v1.saved_model.builder.SavedModelBuilder(export_dir)

# 定义服务签名(输入输出需与模型对应)
signature_inputs = {
    'user_id': tf.compat.v1.saved_model.utils.build_tensor_info(user_input),
    'item_id': tf.compat.v1.saved_model.utils.build_tensor_info(item_input)
}
signature_outputs = {
    'prediction': tf.compat.v1.saved_model.utils.build_tensor_info(output)
}

predict_signature = tf.compat.v1.saved_model.signature_def_utils.build_signature_def(
    inputs=signature_inputs,
    outputs=signature_outputs,
    method_name=signature_constants.PREDICT_METHOD_NAME
)

# 保存SavedModel
builder.add_meta_graph_and_variables(
    sess,
    [tag_constants.SERVING],
    signature_def_map={
        signature_constants.DEFAULT_SERVING_SIGNATURE_DEF_KEY: predict_signature
    }
)
builder.save()

2. 调整模型打包结构

转换完成后,你会得到如下结构的SavedModel文件夹:

saved_model/
    1/
        variables/
            variables.data-00000-of-00001
            variables.index
        saved_model.pb

将1/文件夹打包为model.tar.gz,最终解压后结构应为:

model.tar.gz
    1/
        variables/
        saved_model.pb

打包命令示例:

tar -czvf model.tar.gz 1/

然后将该压缩包上传至S3,获取model_data的S3路径。

3. 调整SageMaker推理脚本(可选)

如果你的自定义推理脚本tf_inference.py是针对ckpt编写的,需要修改为适配SavedModel的版本。以下是简化的推理脚本示例:

import tensorflow as tf
import numpy as np

model = None

def load_model(model_dir):
    global model
    # 加载SavedModel
    model = tf.compat.v1.saved_model.load(model_dir, tags=[tf.saved_model.SERVING])
    return model

def predict(input_data, model):
    # 解析输入
    user_ids = np.array(input_data['user_id']).reshape(-1,1)
    item_ids = np.array(input_data['item_id']).reshape(-1,1)
    
    # 获取签名并执行预测
    infer = model.signatures['serving_default']
    result = infer(user_id=tf.convert_to_tensor(user_ids), item_id=tf.convert_to_tensor(item_ids))
    return {'prediction': result['prediction'].numpy().tolist()}

部署代码可保持原结构,确保framework_version="2.7"和model_data指向正确的S3路径即可。


内容的提问来源于stack exchange,提问作者Muhammad Humza

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.06 09:10:32