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

SageMaker实时端点推理失败求助(StableBaselines3 MaskablePPO)

问题排查与修复方案

核心错误点分析

你的推理代码和调用代码存在几处关键问题,直接导致了500服务器错误:

1. 缺失必要模块导入

推理代码中使用了json模块做序列化/反序列化,但完全没导入,会触发NameError,直接中断推理流程。

2. predict_fn引用未定义变量

predict_fn里写了body['observation'],但实际传入的参数是input_data(由input_fn返回的解析后请求数据),body变量根本没定义,必然抛出异常。

3. output_fn序列化逻辑错误

当前output_fn把预测结果直接转成字符串,要是返回的是字典(比如异常时的{"error": "Terrible"}),会变成类似"{'error': 'Terrible'}"的格式,不符合JSON响应规范,SageMaker无法正确处理这种响应。

4. 调用端序列化冗余

用PyTorchPredictor时,predict方法默认会自动处理JSON序列化,你手动调用json.dumps(request)会导致双层序列化,后端解析时直接格式错误。

修正后的推理代码(inference.py)

import os
import json
import sb3_contrib
from sb3_contrib import MaskablePPO

def model_fn(model_dir):
    """加载推理模型"""
    model_path = os.path.join(model_dir, "final_model.zip")
    # 可选:如果实例支持GPU,可指定device='cuda',默认会自动检测设备
    model = MaskablePPO.load(model_path)
    return model

def predict_fn(input_data, model):
    """执行推理逻辑"""
    try:
        observation = input_data['observation']
        action, _ = model.predict(observation, action_masks=None, deterministic=True)
        # 返回可直接序列化的对象,交给output_fn统一处理
        return {"action": action.tolist()}
    except Exception as e:
        # 捕获具体异常信息,方便后续排查
        return {"error": str(e)}

def input_fn(request_body, request_content_type):
    """解析请求输入"""
    if request_content_type == "application/json":
        return json.loads(request_body)
    else:
        raise ValueError(f"不支持的内容类型: {request_content_type}")

def output_fn(prediction, response_content_type):
    """序列化响应输出"""
    if response_content_type == "application/json":
        return json.dumps(prediction), response_content_type
    else:
        raise ValueError(f"不支持的内容类型: {response_content_type}")

修正后的调用代码

model = PyTorchModel(
    entry_point="inference.py",
    model_data="<redacted bucket name>",
    role=get_execution_role(),
    framework_version="1.5",
    py_version="py3",
)

predictor = model.deploy(
    initial_instance_count=1, instance_type="ml.m5.xlarge", endpoint_name=endpoint_name
)

request = {
   "observation": [0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 1, 0, 0, 0, 0, 1, 0, 0, 0, 0, 1, 0, 0, 0, 0]
}

# 直接传入字典,predictor会自动处理JSON序列化
response = predictor.predict(request)
print(response)

额外排查建议

  1. 验证模型文件路径:确认final_model.zip确实在模型包的根目录,要是不在,得调整model_fn里的路径。
  2. 查看CloudWatch日志细节:去CloudWatch找对应端点的日志流(路径:/aws/sagemaker/Endpoints/<你的端点名称>),里面会有具体的堆栈报错信息,这是排查SageMaker推理错误的核心依据。
  3. 对齐依赖版本:确保推理环境的sb3-contrib、stable-baselines3版本和训练时完全一致,可以在模型包根目录加requirements.txt指定版本。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.16 13:26:03