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

AWS SageMaker部署SKLearn NearestNeighbors模型调用端点报错咨询

解决SageMaker部署NearestNeighbors模型无法调用kneighbors的问题

这个问题我之前也碰到过,核心原因是SageMaker的SKLearn默认推理逻辑只适配带有predict方法的监督模型,而NearestNeighbors这类无监督模型的核心交互方法是kneighbors——既没法通过SKLearnPredictor直接调用该方法,用默认的predict方法也会报错(因为模型本身没有这个方法)。

下面是具体的解决方案:

1. 自定义推理脚本,实现kneighbors逻辑

SageMaker允许你通过自定义inference.py脚本,覆盖默认的模型加载和推理逻辑,让端点支持kneighbors的调用。

创建一个名为inference.py的文件,内容如下:

import joblib
import numpy as np

def model_fn(model_dir):
    # 加载训练好的NearestNeighbors模型
    model = joblib.load(f"{model_dir}/model.joblib")
    return model

def predict_fn(input_data, model):
    # 处理输入向量,调用kneighbors方法(可自定义n_neighbors,或从请求参数中读取)
    distances, indices = model.kneighbors(input_data, n_neighbors=11)
    # 将结果序列化为可返回的格式
    return {"distances": distances.tolist(), "indices": indices.tolist()}

2. 部署模型时指定自定义脚本

在部署模型的时候,需要把这个自定义脚本传给SKLearnModel的entry_point参数,确保SageMaker容器使用你的推理逻辑。

示例部署代码:

from sagemaker.sklearn import SKLearnModel

# 替换为你的模型S3路径、IAM角色和对应SKLearn版本
model = SKLearnModel(
    model_data="s3://your-bucket/path/to/model.tar.gz",
    role="your-sagemaker-execution-role",
    entry_point="inference.py",
    framework_version="0.23-1"  # 请匹配训练时使用的SKLearn版本
)

# 部署端点
predictor = model.deploy(instance_type="ml.t2.medium", initial_instance_count=1)

3. 调用端点获取k近邻结果

部署完成后,你可以通过predictor.predict()方法传入输入向量,获取返回的距离和索引结果:

import numpy as np

# 输入向量需要是二维数组(和训练时的输入格式一致)
sample_vector = np.array([[1.2, 3.4, 5.6]])
result = predictor.predict(sample_vector)

# 解析结果
distances = result["distances"][0]
indices = result["indices"][0]

print(f"Top 11近邻的距离:{distances}")
print(f"Top 11近邻的索引:{indices}")

关键注意事项

  • 打包模型时,要把inference.py和训练好的model.joblib一起压缩成model.tar.gz,上传到S3指定路径。
  • 如果需要动态指定n_neighbors,可以在predict_fn里从请求参数中读取(比如把输入改成包含向量和邻居数的字典),增强灵活性。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.13 08:49:53