PyTorch分布式训练保存的DistributedDataParallel模型加载报错如何解决
报错原因
- 核心问题是保存模型时直接序列化了
DistributedDataParallel(DDP)包装后的完整模型对象,而非模型的状态字典(state_dict)。DDP内部包含大量分布式训练场景下的动态运行时参数(包括报错中提到的ddp_join_throw_on_early_termination),这类参数不会被完整序列化,反序列化时缺少构造参数就会抛出该错误。 - 保存代码存在变量名错误:训练过程中被DDP包装的模型变量名为
model,但保存语句写的是torch.save(maml, model_name),maml在保存代码段中没有被定义,保存的对象本身就不符合预期。 - 若训练和加载模型使用的PyTorch版本不一致,不同版本的DDP内部构造参数定义有差异,也会触发该类参数缺失错误。
解决方案
1. 修正训练阶段的模型保存逻辑
丢弃直接保存整个模型对象的写法,仅保存DDP内部实际业务模型的state_dict,且仅在主进程执行保存操作避免多进程写冲突,修改后的保存代码如下:
torch.distributed.init_process_group(backend="nccl") local_rank = torch.distributed.get_rank() torch.cuda.set_device(local_rank) device = torch.device("cuda", local_rank) save_model = f'./model' Path(save_model).mkdir(parents=True, exist_ok=True) net = Net(args) model_name = f"{save_model}/net.pt" torch.save(net.state_dict(), model_name) model = Model(net, args).to(device) model_name = f"{save_model}/model.pt" if torch.cuda.device_count() > 1: model = nn.parallel.DistributedDataParallel(model, device_ids=[local_rank], output_device=local_rank) model.module.fit(tr_data, val_data, args) # 仅主进程执行保存 if local_rank == 0: # 保存DDP内部的实际模型state_dict,而非整个DDP对象 torch.save(model.module.state_dict(), model_name)
2. 修正加载阶段的逻辑
加载时先实例化业务模型,再加载保存的state_dict即可,不需要直接加载完整模型对象:
save_model = f'./model' net = Net(args) model_name = f"{save_model}/net.pt" net.load_state_dict( torch.load(model_name, map_location=torch.device("cpu"))) # 先实例化Model类 maml = Model(net, args).to(device) model_name = f"{save_model}/model.pt" # 加载之前保存的state_dict到实例中 maml.load_state_dict( torch.load(model_name, map_location=torch.device("cuda" if torch.cuda.is_available() else "cpu")))
3. 旧模型兼容方案
如果已经保存了旧的DDP完整模型对象不想重新训练,可以先初始化分布式环境、保证训练和加载的PyTorch版本完全一致,加载后提取module的state_dict重新保存为标准格式即可:
# 加载旧的DDP模型前先初始化最小分布式环境 import torch.distributed as dist dist.init_process_group(backend="nccl", init_method='tcp://127.0.0.1:23456', rank=0, world_size=1) # 加载旧模型 old_model = torch.load("旧model.pt路径", map_location="cpu") # 提取实际模型的state_dict重新保存 torch.save(old_model.module.state_dict(), "新model.pt路径") # 后续就可以用上述标准方法加载新的model.pt
内容的提问来源于stack exchange,提问作者Zihuan
相关产品推荐
相关产品推荐

