如何为SageMaker训练作业适配自定义Docker镜像
适配自定义Docker镜像到SageMaker训练作业的步骤
要让你的机器学习Docker镜像能在SageMaker训练作业中运行,需遵循SageMaker的容器规范,调整镜像结构和代码,具体步骤如下:
1. 适配SageMaker环境变量规范
SageMaker会在启动训练容器时自动注入一系列环境变量,你的训练代码必须依赖这些变量来获取数据路径、模型输出路径等,不能硬编码路径。核心变量包括:
SM_CHANNEL_TRAIN:训练数据所在的目录路径(对应你在训练作业中指定的训练通道)SM_MODEL_DIR:模型文件的输出目录,训练完成后需将模型保存到这里,SageMaker会自动把该目录内容同步到指定S3路径SM_OUTPUT_DATA_DIR:训练过程中生成的额外数据(如日志、中间结果)的输出目录SM_NUM_GPUS:容器可用的GPU数量(如果使用GPU实例)
2. 调整Dockerfile对齐SageMaker目录结构
SageMaker有默认的容器目录约定,建议在Dockerfile中对齐这些路径,避免额外配置:
/opt/ml/code/:存放训练脚本及相关代码,SageMaker会将此目录视为训练入口/opt/ml/model/:对应SM_MODEL_DIR,用于保存最终模型/opt/ml/output/data/:对应SM_OUTPUT_DATA_DIR,用于存放训练衍生数据
示例Dockerfile:
# 基于你现有的机器学习镜像构建 FROM your-ml-image:latest # 设置工作目录为SageMaker默认代码目录 WORKDIR /opt/ml/code # 复制训练脚本和依赖文件到容器内 COPY train.py requirements.txt ./ # 安装依赖(如果你的基础镜像未包含所需依赖) RUN pip install --no-cache-dir -r requirements.txt # 设置容器启动时的入口命令,直接执行训练脚本 ENTRYPOINT ["python", "train.py"]
3. 修改训练代码读取环境变量
训练脚本中需通过环境变量获取路径,示例Python代码片段:
import os import pandas as pd from sklearn.ensemble import RandomForestClassifier import joblib # 从环境变量获取训练数据路径和模型输出路径 train_data_dir = os.environ["SM_CHANNEL_TRAIN"] model_save_path = os.environ["SM_MODEL_DIR"] # 加载训练数据(根据你的数据格式调整读取逻辑) train_df = pd.read_csv(os.path.join(train_data_dir, "train.csv")) X_train = train_df.drop("target", axis=1) y_train = train_df["target"] # 训练模型 model = RandomForestClassifier(n_estimators=100) model.fit(X_train, y_train) # 将模型保存到SageMaker指定的目录 joblib.dump(model, os.path.join(model_save_path, "model.joblib"))
4. 推送镜像到Amazon ECR
SageMaker仅支持从Amazon ECR拉取自定义容器镜像,需完成以下操作:
- 登录到你的AWS账号对应的ECR仓库:
aws ecr get-login-password --region <你的AWS区域> | docker login --username AWS --password-stdin <你的AWS账号ID>.dkr.ecr.<你的AWS区域>.amazonaws.com - 创建ECR仓库(如果还没有):
aws ecr create-repository --repository-name <你的镜像仓库名称> --region <你的AWS区域> - 给本地镜像打标签:
docker tag your-ml-image:latest <你的AWS账号ID>.dkr.ecr.<你的AWS区域>.amazonaws.com/<你的镜像仓库名称>:latest - 推送镜像到ECR:
docker push <你的AWS账号ID>.dkr.ecr.<你的AWS区域>.amazonaws.com/<你的镜像仓库名称>:latest
5. 提交SageMaker训练作业
可以通过AWS控制台或SDK提交训练作业,指定你的ECR镜像:
控制台方式
- 进入SageMaker控制台,创建新的训练作业
- 在“算法来源”中选择“自定义容器”,填入你的ECR镜像完整URI
- 配置训练数据通道(指向S3中的训练数据)、输出路径(S3路径)和实例类型等参数,启动作业
Python SDK方式示例
import sagemaker from sagemaker.estimator import Estimator # 获取SageMaker执行角色和会话 role = sagemaker.get_execution_role() sagemaker_session = sagemaker.Session() # 定义自定义容器的训练器 estimator = Estimator( image_uri="<你的ECR镜像完整URI>", role=role, instance_count=1, instance_type="ml.m5.xlarge", output_path="s3://<你的S3存储桶>/training-output/", sagemaker_session=sagemaker_session ) # 启动训练作业,指定训练数据的S3路径 estimator.fit({"train": "s3://<你的S3存储桶>/training-data/"})
内容的提问来源于stack exchange,提问作者ryfeus
相关产品推荐
相关产品推荐

