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

如何针对二分类问题使用SageMaker的ModelQualityCheckConfig类

解决XGBoost二分类概率输出的模型质量检查配置问题

核心思路

你的问题根源是模型输出的连续概率值被默认当作离散类别处理,导致质量检查任务误判为多分类场景。解决关键是先通过预处理将概率转换为0/1离散标签,再传入ModelQualityCheckConfig进行二分类模型质量评估。

具体实现步骤

1. 编写预处理脚本,将概率转换为离散标签

创建一个Python脚本(比如preprocess_prob.py),读取批量转换的输出数据,对概率列应用阈值生成离散预测标签:

import pandas as pd
import sys

def main():
    # 读取输入数据(批量转换输出:第一列标签,第二列概率)
    input_path = sys.argv[1]
    output_path = sys.argv[2]
    
    df = pd.read_csv(input_path, header=None, names=['label', 'probability'])
    # 应用阈值生成离散预测(这里用0.5,可根据业务调整)
    df['prediction'] = (df['probability'] >= 0.5).astype(int)
    # 保留标签和转换后的预测列,用于质量检查
    df[['label', 'prediction']].to_csv(output_path, header=False, index=False)

if __name__ == "__main__":
    main()

2. 在SageMaker Pipeline中添加预处理步骤

在流水线中加入一个ProcessingStep,调用上述脚本处理批量转换的输出:

from sagemaker.processing import ScriptProcessor, ProcessingInput, ProcessingOutput
from sagemaker.workflow.steps import ProcessingStep

# 定义脚本处理器(使用兼容的Scikit-learn镜像)
script_processor = ScriptProcessor(
    image_uri="763104351884.dkr.ecr.us-east-1.amazonaws.com/sagemaker-scikit-learn:0.23-1-cpu-py3",
    command=["python3"],
    instance_type="ml.t3.medium",
    instance_count=1,
    role=sagemaker_role
)

# 定义预处理步骤
preprocess_step = ProcessingStep(
    name="ConvertProbToDiscretePrediction",
    processor=script_processor,
    inputs=[
        ProcessingInput(
            source=batch_transform_step.properties.TransformOutput.S3OutputPath,
            destination="/opt/ml/processing/input"
        )
    ],
    outputs=[
        ProcessingOutput(
            source="/opt/ml/processing/output",
            destination=f"s3://{bucket}/processed_predictions"
        )
    ],
    code="preprocess_prob.py"
)

3. 配置ModelQualityCheckConfig使用处理后的数据

修改ModelQualityCheckConfig的输入为预处理后的离散标签数据,同时指定问题类型为二分类:

from sagemaker.workflow.quality_check_step import ModelQualityCheckConfig, ModelQualityCheckStep
from sagemaker.model_monitor import ModelQualityMonitor

# 创建模型质量监控器
model_quality_monitor = ModelQualityMonitor(
    role=sagemaker_role,
    instance_count=1,
    instance_type="ml.m5.xlarge",
    volume_size_in_gb=10
)

# 配置模型质量检查
model_quality_check_config = ModelQualityCheckConfig(
    baseline_dataset=f"s3://{bucket}/ground_truth_data",  # 原始标签数据集
    dataset_format={"csv": {"header": False}},
    input_dataset=preprocess_step.properties.ProcessingOutputConfig.Outputs["output"].S3Output.S3Uri,
    problem_type="BinaryClassification",  # 指定为二分类问题
    ground_truth_attribute=0,  # 处理后数据的第一列是真实标签
    inference_attribute=1  # 处理后数据的第二列是离散预测标签
)

# 创建模型质量检查步骤
model_quality_check_step = ModelQualityCheckStep(
    name="ModelQualityCheck",
    monitor=model_quality_monitor,
    quality_check_config=model_quality_check_config
)

关键注意事项

  • 阈值选择:根据业务需求调整阈值(0.5为通用值,部分场景可改用0.3、0.7等),确保离散预测符合业务逻辑。
  • 数据集格式:确保预处理后的输出和基准数据集格式一致(均为无表头CSV,第一列是标签、第二列是预测值)。
  • 问题类型指定:必须显式设置problem_type="BinaryClassification",让质量检查任务正确识别二分类场景。

内容的提问来源于stack exchange,提问作者Elkin Javier Guerra Galeano

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.26 03:23:21