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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.12 03:35:30