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)
额外排查建议
- 验证模型文件路径:确认
final_model.zip确实在模型包的根目录,要是不在,得调整model_fn里的路径。 - 查看CloudWatch日志细节:去CloudWatch找对应端点的日志流(路径:
/aws/sagemaker/Endpoints/<你的端点名称>),里面会有具体的堆栈报错信息,这是排查SageMaker推理错误的核心依据。 - 对齐依赖版本:确保推理环境的
sb3-contrib、stable-baselines3版本和训练时完全一致,可以在模型包根目录加requirements.txt指定版本。
内容的提问来源于stack exchange,提问作者imconfused
相关产品推荐
相关产品推荐

