如何在AWS SageMaker中集成Clarify可解释性与HPO?
将SageMaker Clarify可解释性与HPO超参数调优结合的实现方案
问题背景
需要将SageMaker Clarify的可解释性分析与HPO(超参数调优)任务结合,已参考过使用experiments.run的单训练作业示例,但不清楚如何适配多作业的HPO场景。
文档中的单训练作业示例:
with Run( experiment_name=experiment_name, run_name="combined-report", sagemaker_session=sagemaker_session, ) as run: # 同一实验运行中生成模型训练和可解释性报告 xgb.fit({"train": train_input}, logs=False) clarify_processor.run_explainability( data_config=explainability_data_config, model_config=model_config, explainability_config=shap_config, )
待结合的HPO代码:
# 需将Clarify的run_explainability与下方HPO结合 optimizer = sagemaker.tuner.HyperparameterTuner( container, hyperparameter_ranges=hp_ranges, strategy='Random', objective_type='Maximize', objective_metric_name='val:auc', metric_definitions=metric_definitions, max_jobs=10, max_parallel_jobs=2, ) optimizer.fit(data_channels, wait=True)
核心解决方案
HPO会启动多个独立的训练作业,每个作业对应一组超参数。要为每个作业生成对应的可解释性报告,需将Clarify的调用逻辑嵌入到HPO的训练脚本内部,而非在HPO外部单独执行。这样每个训练作业完成后,会自动触发对应的Clarify可解释性分析。
具体实现步骤
1. 修改训练脚本(train.py)
在训练脚本中完成模型训练后,直接调用Clarify生成可解释性报告,并关联当前训练作业的资源:
import sagemaker from sagemaker.clarify import ClarifyProcessor, DataConfig, ModelConfig, SHAPConfig import xgboost as xgb import os import pandas as pd def main(): # 加载训练/验证数据(根据实际数据格式调整) train_df = pd.read_csv(os.path.join(os.environ["SM_CHANNEL_TRAIN"], "train.csv")) val_df = pd.read_csv(os.path.join(os.environ["SM_CHANNEL_VAL"], "val.csv")) X_train, y_train = train_df.drop("label", axis=1), train_df["label"] X_val, y_val = val_df.drop("label", axis=1), val_df["label"] # 获取HPO传入的超参数 hyperparams = { "max_depth": int(os.environ.get("SM_HP_MAX_DEPTH", 3)), "learning_rate": float(os.environ.get("SM_HP_LEARNING_RATE", 0.1)), "n_estimators": int(os.environ.get("SM_HP_N_ESTIMATORS", 100)) } # 训练XGBoost模型 model = xgb.XGBClassifier(**hyperparams) model.fit( X_train, y_train, eval_set=[(X_val, y_val)], eval_metric="auc", verbose=False ) # 初始化Clarify处理器 sagemaker_session = sagemaker.Session() clarify_processor = ClarifyProcessor( role=os.environ["SM_ROLE"], instance_count=1, instance_type="ml.m5.xlarge", sagemaker_session=sagemaker_session ) # 配置可解释性数据参数 data_config = DataConfig( s3_data_input_path=os.environ["SM_CHANNEL_VAL"], s3_output_path=f"{os.environ['SM_MODEL_DIR']}/clarify-explain", label="label", headers=val_df.columns.tolist(), dataset_type="text/csv" ) # 配置模型参数 model_config = ModelConfig( model_name=os.environ["SM_TRAINING_JOB_NAME"], instance_type="ml.m5.xlarge", sagemaker_session=sagemaker_session ) # 配置SHAP解释参数 shap_config = SHAPConfig( baseline=X_train.head(100).values.tolist(), num_samples=1000 ) # 生成可解释性报告 clarify_processor.run_explainability( data_config=data_config, model_config=model_config, explainability_config=shap_config ) # 保存训练好的模型 model.save_model(os.path.join(os.environ["SM_MODEL_DIR"], "model.bst")) if __name__ == "__main__": main()
2. 编写HPO启动脚本
定义XGBoost估计器和HPO调优器,指向修改后的训练脚本:
import sagemaker from sagemaker.tuner import HyperparameterTuner, IntegerParameter, ContinuousParameter from sagemaker.xgboost import XGBoost # 初始化会话和角色 sagemaker_session = sagemaker.Session() role = sagemaker.get_execution_role() # 定义XGBoost估计器(确保镜像包含Clarify依赖,或使用自定义镜像) xgb_estimator = XGBoost( entry_point="train.py", role=role, instance_count=1, instance_type="ml.m5.xlarge", framework_version="1.5-1", sagemaker_session=sagemaker_session ) # 定义超参数调优范围 hp_ranges = { "max_depth": IntegerParameter(3, 10), "learning_rate": ContinuousParameter(0.01, 0.3), "n_estimators": IntegerParameter(50, 200) } # 初始化HPO调优器 tuner = HyperparameterTuner( estimator=xgb_estimator, hyperparameter_ranges=hp_ranges, strategy="Random", objective_type="Maximize", objective_metric_name="val:auc", metric_definitions=[ {"Name": "val:auc", "Regex": "validation-auc: ([0-9.]+)"} ], max_jobs=10, max_parallel_jobs=2 ) # 启动HPO任务 data_channels = { "train": "s3://your-bucket/path/to/train-data/", "val": "s3://your-bucket/path/to/val-data/" } tuner.fit(inputs=data_channels, wait=True)
关键注意事项
- 权限配置:确保IAM角色拥有
SageMakerFullAccess或细分权限(如sagemaker:CreateProcessingJob、s3:PutObject等),允许Clarify生成报告并写入S3。 - 资源成本控制:每个HPO训练作业会额外启动一个Clarify处理作业,建议选择合适的实例类型,避免不必要的资源开销。
- 结果关联查看:通过SageMaker控制台的「实验和训练作业」页面,可查看每个HPO作业对应的超参数、模型指标以及Clarify可解释性报告,实现超参数与模型可解释性的关联分析。
内容的提问来源于stack exchange,提问作者Yuri Santos
相关产品推荐
相关产品推荐

