本地构建CatBoost Classifier部署至Amazon SageMaker后如何获取端点?
解决方案:本地CatBoost模型部署到Amazon SageMaker并获取端点
核心原因
本地训练的CatBoostClassifier是原生CatBoost库的类,并未集成SageMaker的deploy方法。要部署到SageMaker,需先将模型打包为符合平台要求的格式,再通过SageMaker SDK创建模型资源并部署端点。
具体步骤&代码
1. 保存本地训练好的CatBoost模型
先将训练完成的模型保存为CatBoost原生格式:
from catboost import CatBoostClassifier # 假设已完成模型训练 model = CatBoostClassifier() model.fit(X_train, y_train) # 保存模型到本地 model.save_model("catboost_model.cbm")
2. 准备SageMaker部署所需的模型包
SageMaker要求模型文件放在.tar.gz压缩包内,同时需配套推理脚本处理预测请求。
编写推理脚本inference.py
import catboost import numpy as np import os def model_fn(model_dir): # 加载CatBoost模型 model = catboost.CatBoostClassifier() model.load_model(os.path.join(model_dir, "catboost_model.cbm")) return model def predict_fn(input_data, model): # 处理输入并返回预测结果 predictions = model.predict(input_data) return predictions
打包模型与脚本
将catboost_model.cbm和inference.py放入同一目录(如model_dir),执行压缩命令:
cd model_dir tar -czf catboost_model.tar.gz catboost_model.cbm inference.py
3. 上传模型包到S3
使用SageMaker SDK将压缩包上传至你的S3存储桶:
import sagemaker from sagemaker import get_execution_role sagemaker_session = sagemaker.Session() role = get_execution_role() # 上传模型到S3,替换为你的桶名 model_data = sagemaker_session.upload_data( path="catboost_model.tar.gz", bucket="your-bucket-name", key_prefix="catboost-model" )
4. 创建SageMaker模型并部署端点
指定CatBoost推理镜像(根据AWS区域选择对应镜像,以下以us-east-1为例),创建模型并部署:
from sagemaker.model import Model from sagemaker.predictor import Predictor # 官方CatBoost推理镜像(不同区域镜像前缀不同,需对应调整) image_uri = "763104351884.dkr.ecr.us-east-1.amazonaws.com/catboost:latest" # 创建SageMaker模型对象 sagemaker_model = Model( image_uri=image_uri, model_data=model_data, role=role, predictor_cls=Predictor ) # 部署端点,可根据需求调整实例类型和数量 predictor = sagemaker_model.deploy( initial_instance_count=1, instance_type="ml.t2.medium" ) # 输出端点名称 print("已部署端点名称:", predictor.endpoint_name)
5. 使用端点预测
部署完成后,通过predictor对象发送预测请求:
# 示例输入数据 test_data = np.array([[1, 2, 3, 4]]) # 获取预测结果 result = predictor.predict(test_data) print("预测结果:", result)
6. 清理资源(可选)
不再使用端点时,及时删除避免产生不必要费用:
predictor.delete_endpoint()
内容的提问来源于stack exchange,提问作者Sukalpo Banerjee
相关产品推荐
相关产品推荐

