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

本地训练Sklearn模型部署至SageMaker及预处理方法咨询

本地Sklearn模型部署至SageMaker(含预处理)

一、需准备的内容

  • 本地导出的Sklearn模型文件:通过joblib.dump()或pickle.dump()保存的模型(如model.joblib)
  • 与训练一致的预处理逻辑:包括特征缩放、类别编码、缺失值填充等代码,需保证和训练时的处理逻辑完全匹配
  • SageMaker推理脚本(inference.py):必须包含模型加载、输入预处理、预测执行三个核心函数
  • 依赖包清单(requirements.txt):列出模型及预处理所需的所有Python库和对应版本(如scikit-learn==1.2.2、joblib==1.2.0)
  • AWS IAM角色:拥有SageMaker端点创建权限、S3存储读写权限的角色ARN

二、具体操作步骤

1. 整理本地文件结构

将模型、推理脚本、依赖清单放在同一目录,结构如下:

sklearn-deploy/
├── model.joblib
├── inference.py
└── requirements.txt

如果预处理用的编码器/缩放器是单独保存的(如scaler.joblib),也需放入该目录

2. 编写推理脚本(inference.py)

示例代码(包含预处理逻辑):

import joblib
import numpy as np
import json

def model_fn(model_dir):
    # 加载模型和预处理工具
    model = joblib.load(f"{model_dir}/model.joblib")
    scaler = joblib.load(f"{model_dir}/scaler.joblib")
    return (model, scaler)

def input_fn(request_body, request_content_type):
    # 解析输入数据并执行预处理
    if request_content_type != "application/json":
        raise ValueError("仅支持application/json格式输入")
    
    input_data = json.loads(request_body)
    # 执行训练时的预处理操作(这里以标准化为例)
    scaled_data = scaler.transform(np.array(input_data).reshape(-1, len(input_data[0])))
    return scaled_data

def predict_fn(input_data, model_tuple):
    # 执行预测
    model, _ = model_tuple
    predictions = model.predict(input_data)
    return predictions.tolist()

3. 打包模型并上传至S3

使用SageMaker Python SDK完成打包和上传:

import sagemaker
from sagemaker.sklearn.model import SKLearnModel

sagemaker_session = sagemaker.Session()
role = "你的IAM角色ARN"

# 初始化SKLearn模型对象
sklearn_model = SKLearnModel(
    role=role,
    entry_point="inference.py",
    model_data="s3://你的S3桶名/sklearn-model/model.tar.gz",
    framework_version="1.2-1",  # 匹配本地scikit-learn大版本
    py_version="py39"
)

# 打包本地目录并上传到S3(若未手动上传)
sklearn_model.upload_data(path="./sklearn-deploy", key_prefix="sklearn-model")

4. 创建实时推理端点

调用deploy方法启动端点:

predictor = sklearn_model.deploy(
    initial_instance_count=1,
    instance_type="ml.t2.medium"  # 根据业务需求选择实例类型
)

5. 测试端点功能

用测试数据验证预处理和预测是否正常:

test_data = [[6.2, 3.4, 5.4, 2.3]]  # 示例鸢尾花特征数据
result = predictor.predict(test_data)
print(f"预测结果:{result}")

6. 清理资源(可选)

测试完成后,删除端点避免不必要的费用:

predictor.delete_endpoint()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.13 21:30:01