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

如何在Amazon SageMaker端点内访问模型注册表的模型质量指标

解决方案

有两种可行的方法实现你的需求,优先推荐方法一,简单高效:

方法一:将R2指标内嵌到模型文件中(推荐)

在训练阶段计算完R2值后,直接把指标附加到模型对象上再保存,推理时无需额外API调用即可读取。

训练代码示例:

import joblib
from sklearn.metrics import r2_score

# 训练模型并计算R2值
model = 你的训练函数()
y_pred = model.predict(X_test)
r2 = r2_score(y_test, y_pred)

# 将R2指标绑定到模型对象
model.model_metrics = {"r2": r2}

# 保存模型
joblib.dump(model, "model.joblib")

推理Handler修改:

直接读取模型内嵌的指标,和你最初的思路一致:

def default_predict_fn(self, input_data, model):
    output = model.predict(input_data)
    # 读取模型内嵌的R2值
    model_metric = model.model_metrics['r2']
    # 保证model_r2列长度与预测结果一致
    return pd.DataFrame({
        'model_r2': [model_metric] * len(output),
        'model_prediction': output
    })

方法二:从SageMaker模型注册表获取指标

如果需要读取模型注册表中存储的官方模型质量指标,可通过SageMaker API调用实现:

步骤1:打包模型时添加元数据

训练完成后,创建metadata.json文件记录模型名称和版本号,与model.joblib一起打包成模型压缩包:

{
  "model_name": "你的模型名称",
  "model_version": 1
}

步骤2:修改InferenceHandler

读取元数据并调用API获取指标:

import json
import boto3
import pandas as pd

class InferenceHandler(DefaultInferenceHandler):
    def default_model_fn(self, model_dir):
        model = joblib.load(f"{model_dir}/model.joblib")
        # 读取模型元数据
        with open(f"{model_dir}/metadata.json", "r") as f:
            self.model_metadata = json.load(f)
        return model

    def default_predict_fn(self, input_data, model):
        sm_client = boto3.client("sagemaker")
        # 查询模型版本详情
        response = sm_client.describe_model_version(
            ModelName=self.model_metadata["model_name"],
            Version=self.model_metadata["model_version"]
        )
        # 提取模型质量中的R2值(需根据实际指标存储结构调整)
        model_quality = response.get("ModelMetrics", {}).get("ModelQuality", {})
        r2_score = model_quality.get("Statistics", {}).get("r2")
        
        output = model.predict(input_data)
        return pd.DataFrame({
            'model_r2': [r2_score] * len(output),
            'model_prediction': output
        })

步骤3:配置权限

确保端点执行角色拥有sagemaker:DescribeModelVersion权限,否则API调用会失败。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.09 20:20:36