AWS Sagemaker环境PyTorch模型训练异常,本地/Colab环境可正常运行
排查步骤
- 版本对齐校验
你当前指定的PyTorch 1.3.1是2019年发布的旧版本,和Colab/本地常用的新版本存在大量API行为、多进程逻辑的差异,很多兼容性问题不会直接抛出错误,只会出现执行逻辑乱序、进程挂死的现象。先确认本地正常运行的PyTorch版本,将estimator的framework_version参数修改为对应版本后重试。 - 入口文件结构校验
SageMaker PyTorch容器启动时会多次import你的入口文件,如果你的网络白化、数据集初始化等预训练逻辑没有放在if __name__ == "__main__":包裹的代码块内,或者没有全部放在训练函数内部,多进程加载数据时会重复执行这些逻辑,直接导致输出乱序、资源竞争挂死。本地/Colab单进程运行时不会触发该问题,属于SageMaker上运行PyTorch的高频坑点。
可参考下方标准入口文件结构调整代码:import os import torch import argparse from your_utils import build_model, whiten_network, get_train_loader def train(args): # 所有预训练、模型初始化、数据加载逻辑全部放在train函数内 model = build_model() # 预训练步骤:网络白化 whiten_network(model) # 优先将num_workers设为0测试兼容性,跑通后再逐步调大 train_loader = get_train_loader(args.training_dir, num_workers=0) for epoch in range(args.epochs): # 训练迭代逻辑 for batch in train_loader: pass if __name__ == "__main__": parser = argparse.ArgumentParser() # 接收SageMaker自动传入的参数 parser.add_argument('--epochs', type=int, default=5) parser.add_argument('--training_dir', type=str, default=os.environ.get('SM_CHANNEL_TRAINING')) args = parser.parse_args() train(args) - 数据加载配置校验
检查你代码中DataLoader的num_workers配置,若写死了较大的数值,配合旧版PyTorch的多进程数据加载bug,很容易在第一轮训练结束后出现进程挂死。先将num_workers设为0验证训练流程能正常跑完,再根据实例性能逐步调大数值。 - 分布式配置校验
你当前仅使用1台训练实例,不需要指定分布式backend参数,旧版PyTorch在单实例下初始化分布式上下文很容易出现通信挂死。先移除hyperparameters中的backend参数重试,若后续需要多机训练再按需添加。
内容的提问来源于stack exchange,提问作者moger
相关产品推荐
相关产品推荐

