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

如何在执行时将超参数传入SageMaker Pipeline训练步骤?

在SageMaker Pipeline中动态传递训练超参数的解决方案

问题根源

你定义的input_hyperparams是ParameterString类型,本质是管道执行时的占位符,并非Python字典,因此直接用**解包会触发类型错误;而将其转成字符串传给单个参数,会导致训练容器收到嵌套的参数结构,与默认容器的解析逻辑冲突。

正确实现方式

方法1:在训练脚本中解析JSON格式的超参数字符串

这种方式无需修改管道结构,只需在训练脚本中处理传入的JSON字符串,适合超参数较多或频繁变动的场景:

  1. 保持管道参数定义不变(可将默认值改为标准JSON格式):
input_hyperparams = ParameterString(
    name="input_hyperparams",
    default_value='{"learning_rate":0.01,"batch_size":32}'
)
  1. 给训练估算器传递一个固定名称的参数,绑定这个ParameterString:
train_estimator.set_hyperparameters(hyperparams_json=input_hyperparams)
  1. 在训练脚本(如train.py)中解析JSON字符串,提取所需超参数:
import json
import argparse

if __name__ == "__main__":
    parser = argparse.ArgumentParser()
    # 接收传入的JSON字符串参数
    parser.add_argument("--hyperparams_json", type=str, default="{}")
    args, _ = parser.parse_known_args()
    
    # 解析为字典并提取参数
    hyperparams = json.loads(args.hyperparams_json)
    learning_rate = hyperparams.get("learning_rate", 0.001)
    batch_size = hyperparams.get("batch_size", 64)
    
    # 后续训练逻辑使用这些参数
    print(f"Training with learning rate: {learning_rate}, batch size: {batch_size}")
  1. 启动管道时,传入自定义的超参数字符串即可:
pipeline.start(
    parameters={
        "input_hyperparams": '{"learning_rate":0.005,"batch_size":64}'
    }
)

方法2:为单个超参数定义管道参数(适合少量超参数)

如果超参数数量固定且较少,可直接为每个超参数定义对应类型的管道参数,无需额外解析:

  1. 定义单个超参数的管道参数:
from sagemaker.workflow.parameters import ParameterFloat, ParameterInteger

learning_rate = ParameterFloat(
    name="learning_rate",
    default_value=0.01
)
batch_size = ParameterInteger(
    name="batch_size",
    default_value=32
)
  1. 直接将参数传递给估算器:
train_estimator.set_hyperparameters(
    learning_rate=learning_rate,
    batch_size=batch_size
)
  1. 启动管道时修改参数值:
pipeline.start(
    parameters={
        "learning_rate": 0.005,
        "batch_size": 64
    }
)

注意事项

  • 两种方式都无需重建管道,只需在启动时传入新参数值即可重新运行训练;
  • 方法1的扩展性更强,新增超参数时无需修改管道结构,只需调整JSON字符串内容;
  • 方法2的参数类型更明确,训练脚本无需额外解析逻辑,更易维护。

内容的提问来源于stack exchange,提问作者L Xandor

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.24 00:00:12