如何在SageMaker自定义Docker镜像中配置检查点本地路径?
在SageMaker自定义镜像中通过环境变量配置检查点路径的方案
完全可行,而且完全符合AWS推荐的/opt/ml/checkpoints路径规范,具体实现步骤如下:
1. 在Docker镜像中配置环境变量与目录
在Dockerfile中设置环境变量指定检查点路径,同时提前创建该目录确保容器启动时路径存在:
# 基于合适的基础镜像构建 FROM python:3.9-slim # 设置环境变量,定义检查点存储路径 ENV CHECKPOINT_DIR="/opt/ml/checkpoints" # 创建检查点目录,避免训练时路径不存在报错 RUN mkdir -p ${CHECKPOINT_DIR} # 复制训练脚本到容器指定目录 COPY train.py /opt/ml/code/train.py # 设置训练入口命令 ENTRYPOINT ["python", "/opt/ml/code/train.py"]
2. 训练脚本读取环境变量保存检查点
在训练代码中读取镜像里定义的环境变量,将检查点保存到指定路径,避免硬编码:
import os # 以PyTorch为例,其他框架逻辑一致 import torch # 从环境变量获取检查点路径,同时设置默认值兼容未配置的情况 checkpoint_dir = os.environ.get("CHECKPOINT_DIR", "/opt/ml/checkpoints") # 训练逻辑示例 model = torch.nn.Linear(10, 2) # ... 训练过程 ... # 保存检查点到指定路径 torch.save(model.state_dict(), os.path.join(checkpoint_dir, "epoch_10.pt"))
3. 简化Estimator初始化配置
因为我们用了AWS推荐的默认检查点路径/opt/ml/checkpoints,初始化Estimator时无需再指定checkpoint_local_path,只需要配置S3存储路径即可:
estimator = Estimator( image_uri="<ecr_path>/<algorithm-name>:<tag>", output_path=bucket, base_job_name=base_job_name, # 仅需指定检查点要同步到的S3路径 checkpoint_s3_uri=checkpoint_s3_bucket )
额外说明
- 如果后续需要修改检查点路径,只需调整Dockerfile中的
CHECKPOINT_DIR环境变量并重新构建镜像,无需修改Estimator代码,配置更灵活 - 确保容器内的检查点目录有读写权限,上述Dockerfile中用
mkdir -p创建的目录默认有足够权限,若使用非root用户需额外设置权限
内容的提问来源于stack exchange,提问作者user1315621
相关产品推荐
相关产品推荐

