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

如何将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'
  • 原代码存在语法错误(如语句末尾多余点号、变量名拼写错误)

错误核心原因

  1. SageMaker v2中RealTimePredictor已被废弃,需使用Predictor类替代
  2. v2中不再支持直接设置content_type、serializer等属性,需通过初始化参数或序列化器类配置
  3. 原批量预测逻辑存在变量名错误、字符串拼接逻辑问题,且未适配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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.04 03:15:46