使用自定义算法在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
相关产品推荐
相关产品推荐

