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

使用自定义算法在SageMaker训练失败:超参数解析错误

SageMaker训练超参数解析失败问题排查与解决

问题背景

使用SageMaker训练任务+Python SDK运行训练,训练脚本依赖自定义库,基于ECR自定义镜像在SageMaker Studio环境中执行,报错Failed to parse hyperparameter。

目录结构

working directory
    —Dockerfile
    —train.py
    —requirements.txt 

Dockerfile

# Use python image as base
FROM python:3.10

# Install system dependencies
RUN apt-get update \
    && apt-get install -y --no-install-recommends \
        libpq-dev \
        gcc \
    && rm -rf /var/lib/apt/lists/*

# Set working directory in container
COPY code /opt/program
WORKDIR /code

# Install Python dependencies
COPY requirements.txt /code/
RUN pip install --no-cache-dir -r requirements.txt
RUN pip install sagemaker-training

# Copies the training code inside the container
COPY train.py /opt/ml/code/train.py

# Defines train.py as script entrypoint
ENV SAGEMAKER_PROGRAM train.py

# Set environment variables
ENV PYTHONUNBUFFERED=TRUE
ENV PYTHONDONTWRITEBYTECODE=TRUE
ENV PATH="/opt/program:${PATH}"

requirements.txt

simpletransformers==0.70.0
pandas==2.1.1
numpy==1.26.0
torch==2.2.1
sklearn-deap==0.3.0
sklearn-genetic-opt==0.10.1
boto3==1.33.3
sagemaker

train.py

import argparse
import os
import logging
import pandas as pd
import numpy as np
from sklearn.model_selection import train_test_split
from sklearn.metrics import classification_report, f1_score
from simpletransformers.classification import ClassificationModel
import torch

from sagemaker_pytorch_estimator.pytorch_estimator import PyTorchModel
from sagemaker_containers.data_instances.data_buffer import BufferDataset, BufferedShuffledDataset

logger = logging.getLogger(__name__)
logger.setLevel(logging.DEBUG)
logger.addHandler(logging.StreamHandler())

if __name__ == "__main__":
    parser = argparse.ArgumentParser()

    parser.add_argument("--batch_size", type=int, default=32)
    parser.add_argument("--test_size", type=float, default=0.2)
    parser.add_argument("--target_column", type=str, default="annotation")
    parser.add_argument("--vertical", type=str, default="some_category")
    parser.add_argument("--model_dir", type=str, default=os.environ.get("SM_MODEL_DIR"))
    parser.add_argument("--train", type=str, default=os.environ.get("SM_CHANNEL_TRAIN"))
    parser.add_argument("--val", type=str, default=os.environ.get("SM_CHANNEL_VAL"))
    parser.add_argument("--test", type=str, default=os.environ.get("SM_CHANNEL_TEST"))

    args, _ = parser.parse_known_args()

    model_data = None
    role = None
    entry_point = None

    ....(script continues)

启动脚本

import sagemaker
from sagemaker.session import TrainingInput
from sagemaker.estimator import Estimator

vertical = 'some_category'
s3_bucket = 'some_bucker'
prefix = 'classification'
instance_type = 'ml.m4.xlarge' 
print("Instance Type: {}".format(instance_type))

region = sagemaker.Session().boto_region_name
print("AWS Region: {}".format(region))

role = sagemaker.get_execution_role()
print("RoleArn: {}".format(role))

s3_output_location='s3://{}/{}/{}'.format(s3_bucket, prefix, 'classifier')
container = '############.###.###.##-####-#.amazonaws.com/some-name/ml-training:latest'
print("Image Container: {}".format(container))

estimator = Estimator(
    image_uri=container,
    role=role,
    instance_count=1,
    instance_type=instance_type,
    volume_size=10,
    output_path=s3_output_location,
    sagemaker_session=sagemaker.Session()
)

estimator.set_hyperparameters(vertical=vertical,
                              s3_bucket=s3_bucket,
                              target_column='annotation',
                              test_size=0.2)

estimator.fit()

错误信息

Failed to parse hyperparameter 

已尝试方案

  • 尝试用函数包装超参数,报错TypeError: Estimator.set_hyperparameters() takes 1 positional argument but 2 were given
  • 参考相关问题建议,但无法适配自身场景
  • 看到有说法称argparse与SageMaker不兼容,但AWS官方文档均使用argparse,解决方案表述模糊无法理解

解决方案

1. 修正Dockerfile路径配置

原Dockerfile存在路径混乱问题,SageMaker训练容器默认从/opt/ml/code目录加载训练脚本,需调整路径配置:

# Use python image as base
FROM python:3.10

# Install system dependencies
RUN apt-get update \
    && apt-get install -y --no-install-recommends \
        libpq-dev \
        gcc \
    && rm -rf /var/lib/apt/lists/*

# 设置容器工作目录为SageMaker默认训练代码目录
WORKDIR /opt/ml/code

# 复制依赖文件并安装
COPY requirements.txt ./
RUN pip install --no-cache-dir -r requirements.txt
RUN pip install sagemaker-training

# 复制训练脚本到容器
COPY train.py ./

# 指定训练入口脚本
ENV SAGEMAKER_PROGRAM train.py

# 设置环境变量
ENV PYTHONUNBUFFERED=TRUE
ENV PYTHONDONTWRITEBYTECODE=TRUE

2. 对齐超参数与脚本参数

启动脚本中传递的s3_bucket超参数在train.py的argparse中未定义,SageMaker会将所有超参数传递给训练脚本,导致解析失败。可二选一处理:

  • 在train.py的argparse部分添加对应参数:
    parser.add_argument("--s3_bucket", type=str, default=os.environ.get("SM_HP_S3_BUCKET"))
    
  • 或修改启动脚本,移除未定义的超参数:
    estimator.set_hyperparameters(vertical=vertical,
                                  target_column='annotation',
                                  test_size=0.2)
    

3. 移除训练脚本中不必要的导入

train.py中导入的PyTorchModel和BufferDataset属于SageMaker部署或容器内部模块,训练阶段无需导入,删除这些导入语句避免潜在依赖问题:

# 删除以下两行
# from sagemaker_pytorch_estimator.pytorch_estimator import PyTorchModel
# from sagemaker_containers.data_instances.data_buffer import BufferDataset, BufferedShuffledDataset

4. 确保超参数类型匹配

传递的超参数类型需与train.py中argparse定义的类型一致,比如test_size是float类型,启动脚本中直接传递0.2而非字符串,当前代码已符合,但需注意避免传递字符串格式的数值。


内容的提问来源于stack exchange,提问作者Cyrus Mohammadian

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.24 00:44:53