跳转到主内容
websoft网络软件专家 - 深耕网络技术,打造实用软件!

Python中PyTorch分布式保存模型_仅在主进程进行保存避免冲突

多进程下torch.save会因并发写同一文件导致错乱或报错,应仅由rank=0主进程在dist.barrier()同步后保存state_dict,使用绝对路径并确保目录权限一致。 为什么
torch.save
在多进程下会报错或覆盖? 分布式训练时多个进程(比如 rank=0/1/2/3)同时调用
torch.save
,会导致文件被反复打开写入,轻则保存内容错乱,重则触发
OSError: [Errno 17] File exists
或直接损坏模型文件。根本原因不是 PyTorch 不支持并发写,而是模型保存本质是写磁盘文件——而文件系统不保证多进程同时写同一路径的安全性。 多卡训练中,每个
rank
都有完整模型副本,但只需保存一次
rank == 0
(主进程)通常负责日志、验证和保存,这是约定俗成的协调点 千万别用
if dist.get_rank() == 0:
包住
torch.save
就以为万事大吉——得确认分布式环境已正确初始化 怎么确保只在主进程保存且不漏掉同步? 主进程保存前,必须等所有进程完成当前训练步(尤其是梯度同步和模型状态更新),否则可能保存到未收敛或不一致的参数。 先调用
dist.barrier()
,强制所有进程在此处等待,避免 rank 0 过早保存 再用
if dist.get_rank() == 0:
判断并执行
torch.save
保存路径建议用绝对路径,避免各进程工作目录不同导致写入位置意外 如果用了
torch.compile
或 FSDP,要先
model.state_dict()
或调用对应获取原始参数的方法,不能直接传编译后模型对象
dist.barrier() # 等所有进程跑完这一步 if dist.get_rank() == 0: torch.save({ 'model_state_dict': model.state_dict(), 'optimizer_state_dict': optimizer.state_dict(), }, "/path/to/checkpoint.pt")
用
torch.save
保存时哪些参数容易出问题?
torch.save
默认用 pickle 序列化,但分布式训练中模型可能含非 pickleable 对象(如某些自定义 hook、CUDA stream、未 detach 的计算图引用)。 避免保存整个
model
实例,只保存
state_dict()
—— 它是纯 tensor 字典,安全且轻量 不要保存
optimizer
以外的运行时对象(如
lr_scheduler
若含 lambda 函数,可能无法反序列化) 如果用了混合精度(AMP),
scaler.state_dict()
可以保存,但需确保加载时也启用 AMP 保存时加
_use_new_zipfile_serialization=True
(PyTorch ≥1.6 默认开启),否则老格式在大模型上易出
RuntimeError: unable to open file
保存后其它进程怎么知道文件已就绪? 严格来说,其它进程不需要“知道”——它们本就不该读这个文件。但如果要做 checkpoint 恢复或评估,常见做法是: Python 3.14.3 微软官方的 Python 扩展,是 VS Code 安装量最高的扩展(209M+)。集成 IntelliSense(通过 Pylance)、调试(通过 Python Debugger)、代码检查、格式化、重构和单元测试等功能。支持 Jupyter Notebook、虚拟环境管理和多 Python 版本切换。 下载 立即学习 “ Python免费学习笔记(深入) ”; 主进程保存完成后,可选地广播一个标志(比如用
torch.tensor([1], device='cuda')
+
dist.broadcast
),但多数场景没必要 更稳妥的做法是:恢复逻辑永远由 rank 0 加载,再用
model.load_state_dict()
后调用
dist.broadcast_parameters
(DDP)或
FSDP.load_state_dict
(FSDP)分发 注意:不要让非 rank 0 进程尝试
torch.load
同一文件——虽然不会报错,但浪费 IO 且破坏职责分离 文件系统本身不提供跨进程的“保存完成”通知,依赖 barrier 和显式控制流就够了。真正容易被忽略的是:保存路径的父目录权限是否对所有进程都可写(尤其在 SLURM 或 Kubernetes 环境下,挂载路径权限可能不一致)。

相关文章