如何将SageMaker 1.x批量预测旧代码升级至v2版本?
SageMaker 1.x到v2 XGBoost端点批量预测代码更新方案
问题背景
原基于SageMaker 1.x的批量预测代码在升级到v2后出现多个错误:
'deprecated_class..DeprecatedClass'对象的'content_type'属性无setter- 删除content_type定义后出现
AttributeError: 'NoneType' object has no attribute 'ACCEPT' - 原代码存在语法错误(如语句末尾多余点号、变量名拼写错误)
错误核心原因
- SageMaker v2中
RealTimePredictor已被废弃,需使用Predictor类替代 - v2中不再支持直接设置
content_type、serializer等属性,需通过初始化参数或序列化器类配置 - 原批量预测逻辑存在变量名错误、字符串拼接逻辑问题,且未适配v2的预测返回格式
完整更新代码
import numpy as np import pandas as pd from sagemaker.predictor import Predictor from sagemaker.serializers import CSVSerializer from sagemaker.deserializers import CSVDeserializer def batch_predict(data, predictor, rows=500): # 将数据按指定行数拆分 split_array = np.array_split(data, int(np.ceil(data.shape[0] / rows))) predictions_list = [] for chunk in split_array: # 调用v2的predict方法,直接获取解析后的结果 chunk_predictions = predictor.predict(chunk) # 将二维数组展平为一维,添加到结果列表 predictions_list.extend(chunk_predictions.flatten()) return np.array(predictions_list) def get_predictions(ordered_data, predictor): # 调用批量预测函数,传入DataFrame的数值数组 predictions = batch_predict(ordered_data.values, predictor) # 转换为DataFrame return pd.DataFrame(predictions, columns=['score']) # 1. 创建v2版本的Predictor实例,指定序列化/反序列化器 xgb_predictor = Predictor( endpoint_name='sagemaker-xgboost-2023-01-18', serializer=CSVSerializer(), deserializer=CSVDeserializer() ) # 2. 获取预测结果 predictions = get_predictions(ordered_data, xgb_predictor) # 3. 拼接order_id列(修复原代码的双层中括号错误) predictions2 = pd.concat([predictions, raw_data[['order_id']]], axis=1)
关键改动说明
- 替换废弃类:用
Predictor替代RealTimePredictor,初始化时直接指定serializer和deserializer,无需后续修改属性 - 简化预测逻辑:v2的
predict方法会自动处理序列化/反序列化,无需手动解码字符串和拼接处理,避免格式错误 - 修复语法错误:
- 移除原代码中语句末尾的多余点号(如
decode('utf-8').、'text/csv'.) - 修正变量名错误(
predicates→predictions_list,predictions_new→chunk_predictions) - 修复DataFrame拼接的双层中括号错误(
raw_data[[['order_id']]]→raw_data[['order_id']])
- 移除原代码中语句末尾的多余点号(如
- 优化批量拆分:用
np.ceil替代原有的浮点计算,确保拆分逻辑更严谨
内容的提问来源于stack exchange,提问作者Citronela
相关产品推荐
相关产品推荐

