如何在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
相关产品推荐
相关产品推荐

