如何加速训练循环中torch.save的模型Checkpoint保存速度?
解决PyTorch大模型Checkpoint保存慢的实用方案
针对你遇到的750MB Checkpoint保存耗时占训练30%的问题,给你几个直接有效的优化方法:
启用PyTorch新序列化压缩机制
PyTorch从1.6版本起提供了新的zip格式序列化方式,相比旧pickle格式,文件体积更小、写入速度更快。只需在torch.save中添加参数:torch.save(model_state, "somefile.pt", _use_new_zipfile_serialization=True)PyTorch 2.0+默认开启该机制,老版本需手动指定,能大幅降低写入时间与文件大小。
异步保存,不阻塞训练主线程
将保存操作放到子线程执行,让训练代码无需等待保存完成即可继续运行。用Python的threading模块实现:import threading save_lock = threading.Lock() def save_checkpoint(state, path): with save_lock: torch.save(state, path, _use_new_zipfile_serialization=True) # 替换原保存逻辑 if dice_score > best_dice_score: # ... 其他更新逻辑不变 threading.Thread(target=save_checkpoint, args=(model_state, "somefile.pt")).start()加锁是为了避免多轮次触发保存时,多个线程同时写入同一文件导致损坏。
只保留核心参数,缩减文件体积
检查model_state_dict是否存在冗余:- 若用了
nn.DataParallel或DistributedDataParallel,要保存model.module.state_dict()而非model.state_dict(),后者包含分布式相关冗余状态。 - 只保留必要的训练状态(你当前仅存模型、epoch和分数,已经很精简,可根据实际需求调整)。
- 若用了
用压缩格式存储
若上述优化仍不够,可直接将Checkpoint压缩后保存,比如用gzip:import gzip import torch def save_compressed_checkpoint(state, path): with gzip.open(path + ".gz", "wb") as f: torch.save(state, f, _use_new_zipfile_serialization=True) # 调用示例 save_compressed_checkpoint(model_state, "somefile.pt")读取时需用
gzip.open加载,虽读取多一步解压,但写入速度和文件体积优势明显(体积可缩小30%-50%)。内存临时存储中转(内存充足时使用)
若机器内存足够,可先将Checkpoint写入内存文件系统(如Linux的/dev/shm),再异步同步到磁盘:import shutil import threading def sync_to_disk(tmp_path, final_path): shutil.move(tmp_path, final_path) # 先写入内存tmpfs tmp_path = "/dev/shm/tmp_checkpoint.pt" torch.save(model_state, tmp_path, _use_new_zipfile_serialization=True) # 异步转移到磁盘 threading.Thread(target=sync_to_disk, args=(tmp_path, "somefile.pt")).start()写入内存几乎瞬间完成,完全不占用训练时间。
内容的提问来源于stack exchange,提问作者Yakov Dan
相关产品推荐
相关产品推荐

