You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

PyTorch分布式训练保存的DistributedDataParallel模型加载报错如何解决

报错原因
  1. 核心问题是保存模型时直接序列化了DistributedDataParallel(DDP)包装后的完整模型对象,而非模型的状态字典(state_dict)。DDP内部包含大量分布式训练场景下的动态运行时参数(包括报错中提到的ddp_join_throw_on_early_termination),这类参数不会被完整序列化,反序列化时缺少构造参数就会抛出该错误。
  2. 保存代码存在变量名错误:训练过程中被DDP包装的模型变量名为model,但保存语句写的是torch.save(maml, model_name),maml在保存代码段中没有被定义,保存的对象本身就不符合预期。
  3. 若训练和加载模型使用的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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.10.02 08:15:04