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

本地构建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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.19 07:20:36