Amazon SageMaker Python SDK v3 重新设计了 script mode:训练和部署流程分别围绕统一的 ModelTrainer 与 ModelBuilder 组织,而 SourceCode 可以在任务运行时把本地代码同步进指定容器。对于已有自定义训练脚本、又不想每次改一行代码就重新构建 Docker 镜像的团队,这会明显缩短迭代回路。
从“镜像承载代码”转向“容器承载运行环境”
传统的 bring-your-own-container 流程,常把依赖、训练代码和启动逻辑全部打入镜像。这样做适合运行环境稳定、发布节奏较慢的工作负载,但在模型实验阶段成本很高:改动参数解析、数据预处理或 checkpoint 策略后,都可能需要重新构建、推送和验证镜像。
SDK v3 的 script mode 将职责拆开:
- 自定义容器继续负责 Python、CUDA、系统库和基础依赖。
- 本地目录中的训练代码由
SourceCode在运行时同步到容器。 ModelTrainer统一描述训练任务。ModelBuilder统一描述模型构建和后续部署所需配置。
这并不意味着镜像不再重要。GPU 驱动兼容性、PyTorch 或 Diffusers 版本、系统级编译依赖仍应固定在镜像中;更适合频繁变化的,是业务训练脚本、配置文件和轻量辅助模块。
两类工作负载为何都适用
来源示例覆盖了两个相差很大的场景:scikit-learn Random Forest,以及多 GPU 的 Stable Diffusion 3.5 LoRA 微调。这说明 script mode 的重点不是绑定某个框架,而是提供一致的“代码进入容器”方式。
对于 Random Forest,这种模式适合快速调整特征工程、验证集切分和超参数。训练代码通常很小,但变更频繁。
对于 Stable Diffusion 3.5 LoRA 微调,镜像通常已经包含 CUDA、PyTorch、分布式训练工具和模型训练依赖。此时把 LoRA 配置、数据集处理、恢复训练逻辑放在 SourceCode 中,可以避免每次实验都制作大型 GPU 镜像。多 GPU 训练仍需由容器和训练脚本正确处理进程启动、GPU 可见性及 checkpoint 写入;script mode 不会替你修复分布式训练本身的问题。
可以这样实践:用 SourceCode 提交一个本地随机森林项目
下面是一个可改造的最小项目结构。示例假定你已经有可用的 SageMaker 执行角色、S3 输出位置,以及一个包含 scikit-learn 的训练镜像。SDK v3 的具体构造参数可能随已安装版本略有差异,因此应以当前 SDK 文档和 IDE 类型提示为准,但目录划分与训练脚本入口模式可以直接复用。
rf-project/
├── train.py
└── requirements.txt
train.py 接收 SageMaker 常见的训练通道和模型输出目录,读取 CSV 后训练并保存模型:
import argparse
import os
import joblib
import pandas as pd
from sklearn.ensemble import RandomForestClassifier
from sklearn.model_selection import train_test_split
parser = argparse.ArgumentParser()
parser.add_argument("--n-estimators", type=int, default=200)
parser.add_argument("--max-depth", type=int, default=None)
parser.add_argument("--data-dir", default=os.environ.get("SM_CHANNEL_TRAINING", "/opt/ml/input/data/training"))
parser.add_argument("--model-dir", default=os.environ.get("SM_MODEL_DIR", "/opt/ml/model"))
args = parser.parse_args()
df = pd.read_csv(os.path.join(args.data_dir, "train.csv"))
y = df.pop("label")
X_train, X_test, y_train, y_test = train_test_split(
df, y, test_size=0.2, random_state=42, stratify=y
)
model = RandomForestClassifier(
n_estimators=args.n_estimators,
max_depth=args.max_depth,
random_state=42,
n_jobs=-1,
)
model.fit(X_train, y_train)
print(f"validation_accuracy={model.score(X_test, y_test):.4f}")
os.makedirs(args.model_dir, exist_ok=True)
joblib.dump(model, os.path.join(args.model_dir, "model.joblib"))
本地先验证脚本,确认 CSV 至少包含 label 列和若干数值特征列:
python -m pip install -r requirements.txt
python train.py --data-dir ./data --model-dir ./artifacts --n-estimators 300
接着,在提交训练任务的 Python 文件中,将 rf-project 作为源码目录交给 SDK v3 的 script mode。下面的对象名和资源值应替换为你的账户配置:
from sagemaker.train import ModelTrainer
from sagemaker.core.resources import SourceCode
source_code = SourceCode(
source_dir="./rf-project",
entry_point="train.py",
)
trainer = ModelTrainer(
role_arn="arn:aws:iam::123456789012:role/SageMakerExecutionRole",
image_uri="123456789012.dkr.ecr.us-east-1.amazonaws.com/sklearn-runtime:latest",
instance_type="ml.m5.xlarge",
instance_count=1,
source_code=source_code,
hyperparameters={
"n-estimators": 300,
"max-depth": 12,
},
output_path="s3://my-ml-artifacts/random-forest/output/",
)
trainer.fit({
"training": "s3://my-ml-data/random-forest/train/"
})
这里最值得保留的边界是:不要把访问密钥、临时 token 或数据集密钥写入 SourceCode 目录。同步前还应检查 .gitignore、构建上下文和源码目录,避免把本地虚拟环境、大型 checkpoint、.env 文件或实验数据意外上传。
多 GPU LoRA 微调时要明确责任归属
Stable Diffusion 3.5 的 LoRA 微调比随机森林复杂得多,但拆分原则相同。可以把基础镜像固定为经过验证的 GPU 运行时,把以下高频变化内容保留在本地源码目录:
train_lora.py:参数解析、训练循环和 checkpoint 恢复逻辑。configs/:batch size、学习率、rank、保存间隔等实验配置。dataset.py:caption 清洗、图像预处理和数据校验。launch.sh:调用torchrun或框架分布式启动器的入口。
多 GPU 场景要在提交前确认三件事:容器内的 CUDA、PyTorch 和 NCCL 组合已验证;训练脚本会让每个进程只处理自己的数据分片;模型和日志最终写入 SageMaker 约定的输出路径。否则,即使代码同步成功,也可能遇到 GPU 利用率低、多个进程覆盖同一 checkpoint 或训练结束后没有可导出模型的问题。
采用前的检查清单
script mode 最适合“环境变化慢、训练代码变化快”的团队。将镜像版本作为可追溯的运行时基线,将 SourceCode 与 Git commit、训练参数和数据版本一并记录,才能让快速迭代不牺牲可复现性。
上线前可以检查:
- 镜像是否固定了经过验证的框架和 CUDA 版本,而非只依赖浮动标签。
- 本地源码目录是否只包含任务真正需要的文件。
- 训练入口能否同时在本地和 SageMaker 约定目录下运行。
- 训练产物、日志和 checkpoint 是否写入可持久化的位置。
- 多 GPU 作业是否验证了进程数、数据分片和断点恢复。
当这些边界明确后,SDK v3 的 ModelTrainer、ModelBuilder 与 SourceCode 能把自带模型流程从“改代码就重发镜像”变成“保留稳定运行时,直接提交最新脚本”。