云端GPU训练总断?tmux+checkpoint保活与断点续训实战
发布时间:2026/9/26 12:46:43来源:尧图网络
1. 云端GPU训练为什么总在半夜断掉跑深度学习训练的人大概都经历过这种崩溃晚上十一点挂上一个模型训练设好学习率、batch size看着loss曲线平稳下降安心去睡觉。第二天早上打开终端一看训练进程没了日志停在凌晨两点十七分loss曲线戛然而止。更气人的是GPU显存已经释放机器还在计费但训练成果全没了。这个问题在云端GPU租用场景下尤其普遍。你租的是一台远程服务器通过SSH连上去跑PyTorch训练脚本。只要SSH连接一断——不管是本地网络抖动、笔记本合盖休眠、还是公司网络策略踢人——前台进程就会收到SIGHUP信号默认行为是直接终止。PyTorch训练进程一死所有在显存里的模型参数、优化器状态、学习率调度器状态全部蒸发。解决这个问题需要两条腿走路第一条腿是让训练进程脱离终端独立存活这样SSH断了训练也不断第二条腿是让训练过程能定期保存状态这样即使进程真的挂了比如GPU OOM、机器重启、被抢占也能从最近的存档点恢复而不是从头再来。前者靠tmux这类终端复用工具后者靠PyTorch的checkpoint机制。两者配合才能让一个动辄跑几十小时的训练任务真正稳定跑完。这篇文章面向的是所有在远程GPU上跑PyTorch训练的人——不管你是租的云GPU、用的实验室服务器、还是公司内网的训练集群。只要你的训练时间超过一次SSH会话的稳定时长这套方案就值得你花二十分钟配置好。下面我会从整体设计思路讲起然后拆解checkpoint的保存与加载细节再讲tmux保活的实操最后把我踩过的坑和排查经验整理出来。2. 整体方案设计与核心思路拆解2.1 为什么是tmux加checkpoint这个组合让训练在后台稳定运行市面上有几种思路。第一种是nohup加把进程放到后台并忽略SIGHUP信号。这个方案最简单但有个致命缺陷你没法再看到训练的实时输出只能重定向到日志文件里tail。而且一旦进程真的挂了你没法在同一个会话里重新拉起并观察。第二种是用systemd做成服务这个适合生产环境但配置繁琐调试阶段改个参数就要重载服务不够灵活。第三种就是tmux或screen这类终端复用器它创建一个持久的会话你断开SSH后会话还在重新连上attach回去就能看到完整的终端状态包括正在滚动的训练日志。我选tmux而不是screen原因很实际tmux的分屏和滚动缓冲更好用配置更现代社区活跃度也更高。screen虽然更老牌但默认的滚动回看体验差快捷键也反直觉。对于需要长时间盯着loss曲线、偶尔翻看前面日志的训练场景tmux的体验明显更顺。但光有tmux不够。tmux解决的是“SSH断了进程不死”它解决不了“进程自己崩了怎么办”。GPU训练进程崩溃的原因太多了显存碎片导致OOM、某条数据异常导致loss变成NaN、CUDA context丢失、甚至宿主机被云厂商热迁移。这些情况下进程是真的死了tmux也救不回来。这时候就需要checkpoint——定期把模型参数、优化器状态、当前epoch和step保存到磁盘进程重启后从最近的存档继续。所以这套方案的核心逻辑是tmux负责对抗外部中断网络、终端checkpoint负责对抗内部崩溃进程、硬件。两层防护叠加才能让一个长任务真正跑完。2.2 checkpoint到底该存什么很多人第一次写checkpoint只存了model.state_dict()结果恢复训练后发现loss曲线跳变模型效果变差。原因是漏存了优化器状态。Adam优化器内部维护着每个参数的一阶矩和二阶矩估计这些动量信息对训练连续性至关重要。只加载模型参数而不加载优化器状态相当于让优化器“失忆”重新开始累积动量loss自然会抖。一个完整的训练checkpoint应该包含以下内容保存项作用是否必须model.state_dict()模型权重参数必须optimizer.state_dict()优化器动量状态必须lr_scheduler.state_dict()学习率调度状态用了调度器就必须epoch / global_step训练进度必须scaler.state_dict()AMP混合精度缩放因子用了AMP就必须best_metric历史最优指标建议rng_state随机数生成器状态追求严格复现时需要args / config超参数配置建议这里重点说几个容易漏的。lr_scheduler的状态如果你用的是CosineAnnealing或StepLR调度器内部有last_epoch计数不保存的话恢复后学习率会从初始值重新开始等于白跑。AMP的scaler混合精度训练时GradScaler维护着一个动态缩放因子不保存的话恢复后缩放因子重置可能导致前期梯度下溢。RNG状态PyTorch的随机种子、CUDA的随机状态、numpy的随机状态如果训练里有数据增强或dropout不保存这些就无法严格复现但对大多数场景来说影响不大可以按需取舍。2.3 保存频率的权衡checkpoint存太频繁IO开销大训练速度被拖慢存太稀疏崩溃时丢的进度多。这里有个经验公式保存间隔应该约等于你能接受的最大重跑时间。比如你预计训练总共要20小时能接受崩溃后最多重跑30分钟那就每30分钟存一次。但按时间存有个问题不同step的耗时不一样前期数据加载快后期可能因为显存碎片变慢。更稳妥的做法是按step数存比如每500或1000个step存一次同时记录时间戳方便你事后分析实际间隔。另外建议保留最近N个checkpoint比如3个采用滚动覆盖策略既防止最新存档损坏导致无档可用又不会把磁盘撑爆。云GPU的磁盘通常不大一个7B模型的完整checkpoint含优化器状态可能就有几十GB存太多直接爆盘。3. Checkpoint保存与加载的核心细节3.1 保存函数的完整实现先看一个我实际在用的保存函数它把上面说的所有状态都打包进去了import torch import os import random import numpy as np def save_checkpoint(state, save_dir, filenamecheckpoint.pth, max_keep3): os.makedirs(save_dir, exist_okTrue) filepath os.path.join(save_dir, filename) # 先存到临时文件再原子重命名防止存到一半崩溃导致文件损坏 tmp_path filepath .tmp torch.save(state, tmp_path) os.replace(tmp_path, filepath) # 滚动清理旧checkpoint ckpts sorted( [f for f in os.listdir(save_dir) if f.startswith(checkpoint_epoch)], keylambda x: os.path.getmtime(os.path.join(save_dir, x)) ) while len(ckpts) max_keep: os.remove(os.path.join(save_dir, ckpts.pop(0)))这里有两个关键技巧。第一是原子写入直接torch.save到目标路径如果保存过程中进程被杀比如云厂商抢占实例会留下一个半截的损坏文件下次加载直接报错。先写.tmp再os.replacereplace是原子操作要么成功要么保持原文件不变安全得多。第二是滚动清理按修改时间排序只保留最近max_keep个防止磁盘写满。磁盘写满的后果很严重不仅checkpoint存不了连日志都写不进去训练直接卡死。调用的时候这样组织statestate { epoch: epoch, global_step: global_step, model: model.state_dict(), optimizer: optimizer.state_dict(), scheduler: scheduler.state_dict() if scheduler else None, scaler: scaler.state_dict() if scaler else None, best_metric: best_metric, rng_state: torch.get_rng_state(), cuda_rng_state: torch.cuda.get_rng_state_all(), numpy_rng_state: np.random.get_state(), python_rng_state: random.getstate(), } save_checkpoint(state, save_dir, fcheckpoint_epoch{epoch}.pth)3.2 加载函数的容错处理加载比保存更需要小心因为你要处理各种异常情况文件不存在、文件损坏、版本不匹配。下面是我用的加载函数def load_checkpoint(model, optimizer, scheduler, scaler, load_path, device): if not os.path.exists(load_path): print(f[Checkpoint] 未找到 {load_path}从头开始训练) return 0, 0, float(inf) try: ckpt torch.load(load_path, map_locationdevice) except Exception as e: print(f[Checkpoint] 加载失败: {e}尝试回退到上一个存档) # 这里可以实现回退逻辑找次新的checkpoint return 0, 0, float(inf) model.load_state_dict(ckpt[model]) if optimizer and optimizer in ckpt: 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]) # 恢复随机状态 if rng_state in ckpt: torch.set_rng_state(ckpt[rng_state]) if cuda_rng_state in ckpt and torch.cuda.is_available(): torch.cuda.set_rng_state_all(ckpt[cuda_rng_state]) epoch ckpt.get(epoch, 0) global_step ckpt.get(global_step, 0) best_metric ckpt.get(best_metric, float(inf)) print(f[Checkpoint] 已从 {load_path} 恢复epoch{epoch}, step{global_step}) return epoch, global_step, best_metric注意torch.load的map_location参数一定要显式指定。如果你在GPU上保存的checkpoint拿到只有CPU的机器上加载不指定map_location会直接报错。反过来从CPU存档加载到GPU时map_locationcuda能省去手动搬移的麻烦。3.3 训练循环里怎么嵌入checkpoint逻辑把保存和加载嵌进训练循环需要处理好几个边界。下面是一个精简但完整的训练循环骨架start_epoch, global_step, best_metric load_checkpoint( model, optimizer, scheduler, scaler, resume_path, device ) for epoch in range(start_epoch, num_epochs): model.train() for batch in dataloader: # 如果是从checkpoint恢复跳过已经训练过的step if global_step resume_step: global_step 1 continue inputs, labels batch inputs, labels inputs.to(device), labels.to(device) optimizer.zero_grad() with torch.cuda.amp.autocast(enableduse_amp): outputs model(inputs) loss criterion(outputs, labels) if use_amp: scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() else: loss.backward() optimizer.step() if scheduler: scheduler.step() global_step 1 # 按step保存 if global_step % save_every_steps 0: state build_state(...) save_checkpoint(state, save_dir, fcheckpoint_step{global_step}.pth) # 每个epoch结束也存一次 state build_state(...) save_checkpoint(state, save_dir, fcheckpoint_epoch{epoch}.pth)这里有个细节值得展开恢复时如何跳过已训练的step。如果你的dataloader是从头遍历的恢复后需要跳过前面已经跑过的batch。但更优雅的做法是用可恢复的sampler让dataloader直接从断点位置开始。PyTorch的DistributedSampler支持set_epoch但普通Sampler没有内置的断点续传。一个实用技巧是记录global_step然后在循环开头用continue跳过虽然会浪费一点数据加载时间但实现简单不易出错。对于大数据集这点浪费可以忽略。4. tmux后台保活的完整实操4.1 tmux的安装与基础配置大多数Linux发行版自带tmux没有的话一条命令搞定# Ubuntu/Debian sudo apt-get install tmux # CentOS/RHEL sudo yum install tmux装好后建议先写个配置文件把一些反直觉的默认设置改掉。在~/.tmux.conf里加上# 开启鼠标支持方便滚动和选择窗格 set -g mouse on # 增加滚动缓冲行数默认2000行看训练日志不够 set -g history-limit 50000 # 设置前缀键为Ctrla比默认的Ctrlb顺手 set -g prefix C-a unbind C-b bind C-a send-prefix # 窗口编号从1开始 set -g base-index 1 setw -g pane-base-index 1 # 关闭自动重命名窗口方便自己命名 set -g allow-rename off改完配置后在tmux里按前缀键加:进入命令模式输入source-file ~/.tmux.conf重载或者直接重启tmux。4.2 创建训练会话的标准流程我的习惯是给每个训练任务建一个独立的tmux会话会话名带上项目名和日期方便管理# 创建名为train_bert_0415的会话 tmux new -s train_bert_0415 # 在会话里激活环境并启动训练 conda activate myenv cd /workspace/project python train.py --config config.yaml 21 | tee logs/train_$(date %m%d_%H%M).log这里用tee把输出同时写到终端和日志文件好处是既能在tmux里实时看又能事后grep日志排查问题。日志文件名带时间戳多次重启不会互相覆盖。启动训练后按前缀键加ddetach退出会话训练继续在后台跑。想回来看的时候# 列出所有会话 tmux ls # 重新连接到指定会话 tmux attach -t train_bert_0415 # 或者简写 tmux a -t train_bert_04154.3 让训练脚本自己处理信号tmux虽然能保住进程但有个坑如果你在tmux里直接跑python train.py然后不小心按了CtrlC训练照样会停。更稳妥的做法是在训练脚本里捕获SIGINT和SIGTERM信号收到信号时先保存checkpoint再退出import signal import sys class GracefulKiller: def __init__(self): self.kill_now False signal.signal(signal.SIGINT, self.exit_gracefully) signal.signal(signal.SIGTERM, self.exit_gracefully) def exit_gracefully(self, signum, frame): print(f\n[Signal] 收到信号 {signum}准备保存checkpoint后退出...) self.kill_now True killer GracefulKiller() for epoch in range(start_epoch, num_epochs): for batch in dataloader: # ... 训练代码 ... if killer.kill_now: save_checkpoint(build_state(...), save_dir, checkpoint_interrupt.pth) print([Signal] checkpoint已保存安全退出) sys.exit(0)这样即使你手滑按了CtrlC或者云厂商发送了终止信号训练也能优雅地存盘退出而不是硬生生被杀掉。这个技巧在竞价实例spot instance场景下特别有用因为云厂商会在回收前30秒发SIGTERM你有机会保存最后的进度。4.4 自动重启的守护脚本tmux加checkpoint已经能覆盖大部分场景但如果进程因为OOM被系统杀掉tmux会话还在但里面的python进程没了你需要手动重新attach再启动。如果想更省心可以写一个守护脚本检测到进程退出就自动重新拉起#!/bin/bash # run_train.sh - 自动重启训练脚本 MAX_RETRIES10 RETRY0 while [ $RETRY -lt $MAX_RETRIES ]; do echo [$(date)] 第 $((RETRY1)) 次启动训练... python train.py --config config.yaml --resume auto EXIT_CODE$? if [ $EXIT_CODE -eq 0 ]; then echo [$(date)] 训练正常结束 break fi echo [$(date)] 训练异常退出退出码 $EXIT_CODE10秒后重启... RETRY$((RETRY1)) sleep 10 done把这个脚本放到tmux里跑就实现了“tmux保活 checkpoint续训 自动重启”的三层防护。训练脚本的--resume auto参数让它自动寻找最新的checkpoint恢复不需要人工干预。5. 常见问题与排查技巧实录5.1 checkpoint加载后loss跳变这是最常见的问题八成是优化器状态没加载对。排查步骤先确认保存时optimizer.state_dict()确实在state字典里再确认加载时optimizer.load_state_dict()被调用了最后检查优化器状态里的参数顺序是否和模型参数顺序一致。如果中途改过模型结构比如加了层旧checkpoint的优化器状态就对不上了这种情况只能丢弃优化器状态重新累积或者用strictFalse加载模型参数。另一个可能原因是学习率调度器。如果你用的是带warmup的调度器恢复后last_epoch没对上学习率会突然跳变。检查方法是在恢复后打印当前学习率和保存时的学习率对比不一致就说明调度器状态有问题。5.2 磁盘写满导致训练卡死云GPU的磁盘通常只有几十到几百GB一个带优化器状态的大模型checkpoint可能就占几十GB。如果max_keep设得太大或者日志文件没做轮转磁盘很快写满。磁盘满的表现很隐蔽训练不报错但loss不再更新因为写日志和存checkpoint都失败了进程卡在IO等待上。排查方法训练前用df -h看一眼可用空间训练中定期检查。预防措施有三个一是max_keep设小一点2到3个足够二是日志用logrotate或自己写清理逻辑三是checkpoint存到独立的数据盘而不是系统盘避免影响系统运行。5.3 tmux会话莫名消失tmux会话消失通常有两个原因。一是服务器重启了tmux会话不会持久化到磁盘重启后全没了。这种情况只能靠checkpoint恢复tmux本身救不了。二是有人执行了tmux kill-server把所有会话都杀了。多人共用的服务器上要小心这个。如果服务器会定期重启比如云厂商的维护窗口建议把训练启动脚本加到crontab的reboot里开机自动拉起tmux会话和训练。或者用systemd服务来管理比tmux更抗重启。5.4 CUDA OOM后如何恢复OOM是GPU训练最常见的崩溃原因。进程被OOM杀掉后显存会被释放但tmux会话还在。重新attach后直接重新运行训练脚本即可它会从最近的checkpoint恢复。但要注意如果OOM是因为batch size太大导致的恢复后还会再OOM。这时候需要调小batch size或者开启梯度累积来等效大batch。一个实用技巧是在训练脚本里加显存监控接近上限时主动保存checkpoint并退出而不是等系统杀def check_memory(threshold0.95): allocated torch.cuda.memory_allocated() total torch.cuda.get_device_properties(0).total_memory if allocated / total threshold: print(f[Memory] 显存使用率 {allocated/total:.1%}主动保存并退出) save_checkpoint(build_state(...), save_dir, checkpoint_oom.pth) sys.exit(1)5.5 常见问题速查表现象可能原因排查方法解决措施恢复后loss跳变优化器/调度器状态未加载打印学习率和优化器state补全state_dict加载训练卡住无输出磁盘写满df -h检查空间清理旧checkpoint和日志tmux会话消失服务器重启uptime查看运行时间用crontab或systemd自启反复OOMbatch size过大nvidia-smi看显存调小batch或梯度累积checkpoint文件损坏保存时进程被杀torch.load报错原子写入保留多个存档恢复后step不对global_step未保存打印恢复的step值在state里加global_step6. 几个我踩过的坑和实操心得第一个坑是关于checkpoint文件大小的。我一开始把整个模型对象torch.save(model)存下来结果文件巨大而且加载慢。正确做法是只存state_dict它是纯字典体积小、加载快、跨版本兼容性好。模型结构用代码定义checkpoint只存参数这是PyTorch的推荐实践。第二个坑是多卡训练的checkpoint。如果你用DataParallel或DistributedDataParallelmodel.state_dict()保存的是带module.前缀的键。单卡加载时会报键不匹配。解决方法是在保存时用model.module.state_dict()或者在加载时用strictFalse并手动处理前缀。DDP场景下建议只在rank 0上保存避免多个进程同时写同一个文件导致损坏。第三个心得是checkpoint的命名要带足够信息。我现在的命名格式是checkpoint_epoch{epoch}_step{step}_loss{loss:.4f}.pth这样不用加载就能从文件名看出这个存档的训练状态。配合滚动清理时按文件名里的step排序比按修改时间更可靠因为修改时间可能被复制操作改变。第四个心得是定期验证checkpoint可加载。我遇到过存了几十个小时的checkpoint真正需要恢复时才发现加载报错原因是保存时某个tensor是NaNtorch.save能存但加载后模型输出全NaN。现在我养成了习惯每次保存后立刻用一个小脚本验证能否加载并跑一次前向def verify_checkpoint(path, model, device): try: ckpt torch.load(path, map_locationdevice) model.load_state_dict(ckpt[model]) dummy torch.randn(1, 3, 224, 224).to(device) with torch.no_grad(): out model(dummy) assert not torch.isnan(out).any(), 输出包含NaN print(f[Verify] {path} 验证通过) return True except Exception as e: print(f[Verify] {path} 验证失败: {e}) return False这个验证花不了几秒钟但能在关键时刻救你一命。毕竟训练了几十小时的成果不能赌在“应该没问题”上。最后说一个关于tmux的小技巧如果你在tmux里跑训练想同时监控GPU状态可以分屏。按前缀键加%左右分屏在右边跑watch -n 5 nvidia-smi左边继续看训练日志。这样一眼就能看到显存占用和GPU利用率不用来回切换窗口。如果GPU利用率长期低于50%说明数据加载是瓶颈该考虑增加num_workers或者把数据预处理做缓存了。
网站建设高端定制企业官网