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

使用SageMaker Pipeline+RegisterModel部署PyTorch模型缺失inference.py报错求助

解决方案

你遇到的报错是因为PyTorch模型在SageMaker部署时,需要配套的推理脚本被打包到模型制品中,你当前的训练配置仅指定了训练入口文件,未将推理脚本同步打包到模型包,因此部署时找不到对应文件。

1. 编写推理脚本inference.py

将该文件和你的训练脚本train.py放在同一个本地目录下,示例代码适配你使用的CSV输入输出格式:

import os
import torch
import pandas as pd
import io

# 加载训练好的模型,必填函数
def model_fn(model_dir):
    model = YourPyTorchModelClass() # 替换为你自己定义的模型类,要和训练时的结构完全一致
    with open(os.path.join(model_dir, "model.pth"), "rb") as f:
        model.load_state_dict(torch.load(f))
    model.eval()
    return model

# 解析输入数据,必填函数
def input_fn(input_data, content_type):
    if content_type == "text/csv":
        # 读取CSV格式输入,如有表头可自行调整参数
        df = pd.read_csv(io.StringIO(input_data), header=None)
        return torch.tensor(df.values).float()
    else:
        raise ValueError(f"不支持的Content类型: {content_type}")

# 执行推理,必填函数
def predict_fn(input_data, model):
    with torch.no_grad():
        output = model(input_data)
    return output.numpy()

# 格式化输出结果,必填函数
def output_fn(prediction, accept):
    if accept == "text/csv":
        # 将预测结果转成CSV格式返回
        return pd.DataFrame(prediction).to_csv(header=False, index=False), accept
    else:
        raise ValueError(f"不支持的Accept类型: {accept}")

如果推理有额外依赖包,可在同目录下新建requirements.txt写入依赖项,部署时会自动安装。

2. 修改PyTorch Estimator配置

给Estimator添加source_dir参数,指向你存放train.py和inference.py的本地目录,SageMaker训练完成后会自动将整个目录打包到模型制品的code子目录中,部署时可直接读取:

pytorch_estimator = PyTorch(
    entry_point= 'train.py', # 已指定source_dir,这里直接写文件名即可
    source_dir= BASE_DIR, # 替换为你存放train.py和inference.py的目录路径
    instance_type= "ml.m5.xlarge",
    instance_count=1,
    role=role,
    framework_version='1.8.0',
    py_version='py3',
    hyperparameters = {'epochs': 5, 'batch-size': 64, 'learning-rate': 0.1}
)

3. 确认RegisterModel配置

你当前的RegisterModel步骤已经关联了修改后的pytorch_estimator,无需额外修改,会自动继承推理相关的配置。

修改完成后重新运行流水线即可,训练完成后生成的模型包会自动包含推理脚本,部署时不会再报找不到inference.py的错误。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.30 19:54:03