如何针对二分类问题使用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
相关产品推荐
相关产品推荐

