本地训练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
相关产品推荐
相关产品推荐

