PyTorch云端训练保活与断点续训实战:tmux与checkpoint完整方案

发布时间:2026/9/26 5:57:14
PyTorch云端训练保活与断点续训实战:tmux与checkpoint完整方案
1. 云端训练任务为什么会“跑着跑着就没了”但凡把 PyTorch 训练任务放到云端 GPU 上跑过的人大概率都经历过这种崩溃时刻模型训到第 8 个 epochloss 曲线正漂亮地往下走你关掉浏览器去吃了个饭回来一看——SSH 会话断了进程没了显存也释放了日志停在最后一个 step 上前面几个小时的算力全打了水漂。更气人的是你甚至不知道它是被 OOM 杀的、被会话超时断的还是被平台调度回收的。这个问题的本质其实不是 PyTorch 的锅而是训练进程的生命周期管理没做好。云端 GPU 任务面临三重威胁第一是 SSH 会话本身不持久网络抖动、本地休眠、终端关闭都会触发 SIGHUP 把前台进程带走第二是训练脚本自身没有容错一旦中途异常退出只能从头再来第三是平台侧的资源回收策略比如抢占式实例、配额冻结、节点维护这些你控制不了但可以提前防御。所以一套能“稳定跑完”的方案必须同时解决两件事进程要活得住后台保活状态要存得下断点续训。前者靠 tmux 这类会话保持工具后者靠 checkpoint 机制。这两件事单独看都不复杂但真正落地到实际项目里坑非常多——checkpoint 存什么、多久存一次、存到哪、怎么保证原子性、恢复时怎么对齐 optimizer 状态和 scheduler 步数每一步都有讲究。这篇文章面向的是已经在用 PyTorch 做训练、并且需要把任务放到云端 GPU 上长时间运行的开发者。不管你是微调大模型、跑强化学习、还是训一个中等规模的 CV 模型只要单次训练超过半小时这套东西就值得认真搭一遍。下面我会从进程保活讲到 checkpoint 设计再到恢复逻辑的完整实现最后聊聊那些只有踩过才知道的细节。2. tmux 保活让训练进程脱离 SSH 会话独立存活2.1 为什么 nohup 不够用tmux 才是正解很多人第一反应是用nohup python train.py 把任务丢到后台。这招在简单场景下能用但一旦你需要中途查看训练状态、临时敲个命令、或者进程卡住时进去调试nohup 就抓瞎了——你没法重新“attach”回那个进程的终端上下文只能靠日志文件盲猜。tmux 的核心价值在于它提供了一个持久化的伪终端会话。你在这个会话里启动的进程其父进程是 tmux server而不是你的 SSH 连接。SSH 断了tmux server 还在进程就还在。等你重新连上服务器tmux attach一下屏幕上的输出原封不动地回来了就像从没离开过。我实测下来tmux 相比 screen 更值得推荐原因是它的配置更灵活、窗口分屏更顺手、脚本化能力更强。下面是一套我常用的配置和操作流程。2.2 从零搭一个抗断线的训练会话先装 tmuxUbuntu/Debian 系直接apt install tmuxCentOS 系yum install tmux。装完之后建议先写一份~/.tmux.conf把一些默认反人类的快捷键改掉# ~/.tmux.conf set -g mouse on # 允许鼠标滚轮翻页、点击切窗格 set -g history-limit 50000 # 回滚缓冲区加大方便翻训练日志 set -g base-index 1 # 窗口编号从1开始 setw -g pane-base-index 1 set -g renumber-windows on set -g escape-time 10 # 降低ESC延迟vim用户友好 bind r source-file ~/.tmux.conf \; display config reloaded配好之后标准操作流程是这样的# 1. 新建一个名为 train 的会话直接进入 tmux new -s train # 2. 在会话里激活环境并启动训练 conda activate myenv cd /path/to/project python train.py --config configs/base.yaml 21 | tee logs/train_$(date %m%d_%H%M).log # 3. 按 Ctrlb 然后按 ddetach 出来进程继续跑 # 4. 下次连上服务器后重新进入 tmux attach -t train这里有个细节值得说21 | tee这个组合把 stderr 合并进 stdout同时输出到终端和日志文件。为什么要这么做因为 PyTorch 的 tqdm 进度条、警告信息很多是走 stderr 的不合并的话日志文件里会缺东西。而 tee 保证你既能在 tmux 里实时看又有落盘记录可查。2.3 会话管理里那些容易翻车的点第一个坑是会话名冲突。如果你习惯用tmux new -s train第二次执行会报 duplicate session。正确做法是先tmux ls看看有没有同名会话有就 attach没有才 new。我一般写个小函数丢进.bashrctm() { tmux attach -t $1 2/dev/null || tmux new -s $1 }这样tm train一条命令搞定“有则进、无则建”。第二个坑是tmux 里的进程被 OOM Killer 干掉。tmux 只能防会话断开防不了系统内存不足时内核杀进程。这种情况日志里通常看不到 Python 的 traceback进程直接消失。判断方法dmesg | grep -i killed process如果看到你的 python 进程被 kill那就是系统内存问题得从 batch size、数据加载器 worker 数量上找原因。第三个坑是服务器重启。tmux 会话不会在重启后自动恢复这是物理限制。如果你的平台会不定期维护重启那必须配合 checkpoint 才能扛过去——这也正好引出下一部分。提示tmux 会话是绑定到具体机器的。如果你用的是 K8s 或容器化平台Pod 重建后 tmux 会话同样会丢这时候保活要靠平台自身的重启策略而状态恢复只能靠 checkpoint。3. Checkpoint 到底该存什么不只是 model.state_dict()3.1 一个“能真正恢复”的 checkpoint 包含哪些字段新手最容易犯的错就是只存model.state_dict()。结果恢复训练时发现optimizer 的动量没了、学习率调度器的步数归零了、epoch 计数从头开始了。这样恢复出来的训练loss 会有一个明显的跳变甚至可能发散。一个完整的、可无缝续训的 checkpoint至少应该包含以下内容字段作用不存的后果model_state_dict模型权重等于从头训optimizer_state_dict优化器动量/二阶矩loss 跳变收敛变慢scheduler_state_dict学习率调度状态学习率重置训练节奏乱epoch/global_step进度计数日志、断点判断错乱best_metric历史最优指标无法正确保存最优模型rng_state可选随机数状态数据增强、dropout 不可复现scaler_state_dictAMP混合精度缩放因子AMP 训练恢复后可能溢出我一般会封装成一个函数把该存的都存进去import torch import random import numpy as np def save_checkpoint(path, model, optimizer, scheduler, epoch, global_step, best_metric, scalerNone): ckpt { model: model.state_dict(), optimizer: optimizer.state_dict(), scheduler: scheduler.state_dict() if scheduler else None, epoch: epoch, global_step: global_step, best_metric: best_metric, rng_state: { torch: torch.get_rng_state(), cuda: torch.cuda.get_rng_state_all(), numpy: np.random.get_state(), python: random.getstate(), }, } if scaler is not None: ckpt[scaler] scaler.state_dict() # 先写临时文件再原子重命名防止写一半崩溃导致文件损坏 tmp_path path .tmp torch.save(ckpt, tmp_path) os.replace(tmp_path, path)注意最后那个os.replace的写法。这是原子写的关键如果直接torch.save(ckpt, path)在写入过程中进程被 kill你会得到一个半截的、无法加载的损坏文件而它恰好覆盖了上一个好的 checkpoint。用临时文件 原子重命名就能保证磁盘上永远是一个完整可用的版本。3.2 存盘频率太勤伤 IO太懒丢进度checkpoint 的保存间隔是个权衡。存太频繁每个 epoch 都写几个 GBIO 成为瓶颈训练速度明显下降存太稀疏比如 10 个 epoch 才存一次一旦崩溃就丢 10 个 epoch 的进度。我的经验是分两档周期性 checkpoint和最优 checkpoint。周期性的一般按时间或步数触发比如每 30 分钟或每 1000 个 step 存一次覆盖写同一个文件或保留最近 2-3 个滚动版本最优 checkpoint 则在验证指标刷新时保存单独命名永不覆盖。SAVE_INTERVAL_SEC 1800 # 30分钟 last_save_time time.time() for epoch in range(start_epoch, num_epochs): for step, batch in enumerate(loader): train_step(batch) global_step 1 if time.time() - last_save_time SAVE_INTERVAL_SEC: save_checkpoint(ckpt/last.pt, model, optimizer, scheduler, epoch, global_step, best_metric, scaler) last_save_time time.time()按时间触发比按 step 触发更稳因为不同 batch 的处理耗时可能差异很大尤其是变长序列任务按 step 算不准真实的时间成本。3.3 大模型场景下的存储优化如果你在微调大模型一个 checkpoint 动辄几十 GB存盘本身就是个负担。几个实用的优化方向只存可训练参数LoRA 微调时冻结的主干权重根本不用存只存 adapter 的几十 MB 即可。分片保存用torch.save配合state_dict分片或者直接用safetensors格式加载更快也更安全。异步写盘把torch.save丢到一个后台线程里训练主循环不阻塞。但要注意异步写盘期间如果进程崩溃可能丢最后一次 checkpoint所以关键节点还是同步写。import threading def async_save(path, ckpt): def _save(): tmp path .tmp torch.save(ckpt, tmp) os.replace(tmp, path) t threading.Thread(target_save) t.start() return t注意异步写盘时ckpt里的 tensor 如果在写盘期间被训练循环修改会存出错乱的数据。稳妥做法是先copy.deepcopy或者把 state_dict 转到 CPU 再交给后台线程。4. 恢复逻辑让训练从断点“无缝”接上4.1 加载顺序错了恢复就白搭恢复训练时加载顺序是有讲究的。正确的顺序是先建模型和优化器实例再加载 state_dict最后把 RNG 状态也恢复上。很多人把 RNG 恢复漏了导致数据增强的随机序列和中断前对不上虽然不影响收敛但严格来说不算“无缝”。def load_checkpoint(path, model, optimizer, scheduler, scalerNone): ckpt torch.load(path, map_locationcpu) model.load_state_dict(ckpt[model]) optimizer.load_state_dict(ckpt[optimizer]) if scheduler and ckpt.get(scheduler): scheduler.load_state_dict(ckpt[scheduler]) if scaler and ckpt.get(scaler): scaler.load_state_dict(ckpt[scaler]) # 恢复随机数状态 rng ckpt.get(rng_state) if rng: torch.set_rng_state(rng[torch]) torch.cuda.set_rng_state_all(rng[cuda]) np.random.set_state(rng[numpy]) random.setstate(rng[python]) return ckpt[epoch], ckpt[global_step], ckpt[best_metric]这里有个非常隐蔽的坑torch.load默认会把 tensor 加载到保存时的设备上。如果你在 GPU 上存的 checkpoint换到另一台机器或 CPU 上加载会直接报错。所以务必加map_locationcpu加载完再model.to(device)。这个参数看起来不起眼但在跨环境恢复时能救命。4.2 数据加载器的“断点对齐”问题模型状态恢复了但数据加载器不一定能对齐。如果你用的是DataLoader的随机 sampler恢复后它会从头开始 shuffle导致某些样本被重复训练、某些被跳过。对于 epoch 级别的恢复这个问题影响不大但如果你做的是 step 级别的精确恢复就需要保存 sampler 的状态。PyTorch 的RandomSampler支持通过set_epoch配合DistributedSampler来保证每个 epoch 的 shuffle 可复现。单机场景下更简单的做法是记录global_step恢复时用itertools.islice跳过已经训练过的 batchfrom itertools import islice for epoch in range(start_epoch, num_epochs): loader_iter iter(loader) if epoch start_epoch and global_step 0: # 跳过本 epoch 内已经训练过的 step steps_done_in_epoch global_step % len(loader) loader_iter islice(loader_iter, steps_done_in_epoch, None) for batch in loader_iter: ...这个逻辑在单机训练里够用分布式场景下要复杂一些需要每个 rank 各自对齐。4.3 自动检测并恢复让脚本自己决定从哪开始最省心的方案是让训练脚本启动时自动检查有没有可用的 checkpoint有就从最新的恢复没有就从头开始。这样配合 tmux即使进程意外退出你只要重新跑一遍同样的命令它自己就接上了。def find_latest_checkpoint(ckpt_dir): if not os.path.isdir(ckpt_dir): return None files [f for f in os.listdir(ckpt_dir) if f.endswith(.pt)] if not files: return None files.sort(keylambda f: os.path.getmtime(os.path.join(ckpt_dir, f))) return os.path.join(ckpt_dir, files[-1]) # 主流程 ckpt_path find_latest_checkpoint(ckpt) if ckpt_path: start_epoch, global_step, best_metric load_checkpoint( ckpt_path, model, optimizer, scheduler, scaler) print(fResumed from {ckpt_path} at epoch {start_epoch}, step {global_step}) else: start_epoch, global_step, best_metric 0, 0, float(inf)配合一个while循环包住训练主体还能实现“崩溃自动重启”while true; do python train.py --config configs/base.yaml code$? if [ $code -eq 0 ]; then echo Training finished normally. break fi echo Training crashed with code $code, restarting in 10s... sleep 10 done这段 shell 逻辑放在 tmux 会话里跑就形成了一个相当健壮的“自愈”训练环境进程崩了自动重启重启后自动从最新 checkpoint 恢复tmux 保证会话不断。我实测下来这套组合能扛住绝大多数非硬件级别的意外。5. 那些只有踩过才知道的细节5.1 checkpoint 文件损坏与版本兼容前面提到的原子写能防大部分损坏但还有一种情况磁盘写满。当磁盘空间不足时torch.save可能写出一个截断的文件而os.replace依然会成功于是好的 checkpoint 被坏文件覆盖了。防御方法是保存前检查磁盘剩余空间import shutil def safe_save(path, ckpt, min_free_gb5): free_gb shutil.disk_usage(os.path.dirname(path)).free / 1e9 if free_gb min_free_gb: print(fWARNING: only {free_gb:.1f}GB free, skip saving {path}) return False tmp path .tmp torch.save(ckpt, tmp) os.replace(tmp, path) return True另一个坑是PyTorch 版本兼容。用 2.0 存的 checkpoint在 1.13 上加载可能因为weights_only默认值变化而报错。跨版本恢复时建议显式指定torch.load(path, map_locationcpu, weights_onlyFalse)并确保两边的模型定义代码一致。5.2 GPU 显存碎片与长时间运行的稳定性训练跑得越久越容易遇到显存碎片问题。表现是明明还有几 GB 空闲显存但分配一个小 tensor 就 OOM。这通常是因为频繁的变长分配导致显存池碎片化。缓解手段有两个一是设置PYTORCH_CUDA_ALLOC_CONFexpandable_segments:True让分配器用可扩展段减少碎片二是定期在验证阶段torch.cuda.empty_cache()但注意这个操作会拖慢速度别在训练循环里频繁调用。还有一个容易被忽略的点长时间训练后 GPU 温度墙导致的降频。如果你发现训练速度在几个小时后明显变慢但 loss 正常很可能是 GPU 过热降频了。用nvidia-smi -q -d TEMPERATURE看看温度必要时改善散热或降低功耗上限。5.3 日志与监控出问题时能快速定位训练崩了不可怕可怕的是不知道为啥崩。我的习惯是至少记录三类信息训练指标loss、lr、grad_norm、系统指标GPU 利用率、显存、温度、异常堆栈完整 traceback 落盘。grad_norm 尤其值得记录它是判断训练是否健康的早期信号。如果 grad_norm 突然飙升到几百甚至上千往往预示着即将发散这时候如果刚好有个 checkpoint就能回滚到健康状态重来。total_norm torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) writer.add_scalar(train/grad_norm, total_norm, global_step) if total_norm 100: print(fWARNING: large grad norm {total_norm:.2f} at step {global_step})配合 TensorBoard 或 WandB把指标实时可视化即使你不在服务器前也能通过手机看一眼训练是否正常。5.4 平台配额与抢占的应对策略云端 GPU 平台经常有配额限制或抢占机制。热词里提到的“配额已不够预冻结”就是典型场景——你的任务可能因为配额不足被暂停甚至回收。应对策略是把 checkpoint 存到持久化存储上而不是容器本地盘。容器一旦被回收本地盘的数据全没了只有挂载的网络存储或对象存储能保住 checkpoint。另外如果你的任务跑在抢占式实例上要假设它随时可能被中断。这种情况下 checkpoint 频率要调高比如每 10 分钟一次并且恢复逻辑要足够健壮能在新实例上自动拉起。有些平台提供了中断通知信号比如收到 SIGTERM 后给你 30 秒善后可以注册一个信号处理器在收到信号时立刻存一次 checkpointimport signal def handle_sigterm(signum, frame): print(Received SIGTERM, saving emergency checkpoint...) save_checkpoint(ckpt/emergency.pt, model, optimizer, scheduler, epoch, global_step, best_metric, scaler) sys.exit(0) signal.signal(signal.SIGTERM, handle_sigterm)这个“临终存盘”机制在抢占式环境里价值极高能把损失从几十分钟压缩到几十秒。6. 一套可以直接抄的完整骨架把上面的东西串起来一个能扛住云端环境的训练脚本骨架大概长这样import os, sys, time, signal, random import numpy as np import torch def main(): device torch.device(cuda if torch.cuda.is_available() else cpu) model build_model().to(device) optimizer build_optimizer(model) scheduler build_scheduler(optimizer) scaler torch.cuda.amp.GradScaler() ckpt_path find_latest_checkpoint(ckpt) if ckpt_path: start_epoch, global_step, best_metric load_checkpoint( ckpt_path, model, optimizer, scheduler, scaler) else: start_epoch, global_step, best_metric 0, 0, float(inf) def emergency_save(signum, frame): save_checkpoint(ckpt/emergency.pt, model, optimizer, scheduler, epoch, global_step, best_metric, scaler) sys.exit(0) signal.signal(signal.SIGTERM, emergency_save) last_save time.time() for epoch in range(start_epoch, NUM_EPOCHS): for batch in train_loader: loss train_step(model, batch, optimizer, scaler) global_step 1 if time.time() - last_save 1800: save_checkpoint(ckpt/last.pt, model, optimizer, scheduler, epoch, global_step, best_metric, scaler) last_save time.time() metric validate(model, val_loader) if metric best_metric: best_metric metric save_checkpoint(ckpt/best.pt, model, optimizer, scheduler, epoch, global_step, best_metric, scaler) if __name__ __main__: main()启动命令放在 tmux 里外面套一层自动重启的 shell 循环checkpoint 目录挂到持久化存储上。这套组合我在多个项目里用过从单卡微调到多卡预训练都扛得住最长的一次连续跑了十几天没出问题。最后分享一个我踩过的坑别把 checkpoint 和日志存在同一个目录下用通配符清理。我曾经写过一个清理脚本rm ckpt/*.pt想删旧 checkpoint结果手滑把best.pt也删了而那次训练的最优模型恰好只存在这一个文件里。从那以后我养成了习惯——最优模型单独放一个目录清理脚本只针对last_*.pt这类滚动文件并且清理前先ls确认一遍。这种低级错误往往比技术难题更让人肉疼。