云端GPU训练保活与断点续训:tmux与PyTorch checkpoint实战
发布时间:2026/9/26 12:46:43来源:尧图网络
1. 云端训练任务为什么总是“跑不完”1.1 一个让所有炼丹师都头疼的场景你花了大半天时间调试模型结构、清洗数据、配置环境终于把训练脚本跑通了。看着终端里 loss 一点点下降你心满意足地去吃饭、睡觉或者关掉笔记本去忙别的事。结果第二天早上打开终端一看——连接断了进程没了日志停在半夜两点半loss 曲线戛然而止。这种情况在云端 GPU 训练里太常见了。不管是租用的 GPU 服务器、实验室的共享集群还是公司内部的训练平台只要训练任务超过几个小时就一定会遇到连接中断、会话超时、进程被意外杀掉的问题。尤其是当你用 SSH 连到远程机器上直接跑python train.py的时候网络一抖动终端一关闭SIGHUP 信号就把你的训练进程带走了。更让人崩溃的是有些任务已经跑了十几个 epochcheckpoint 也没存一切从头再来。如果是小模型还好几十分钟能重跑但如果是在微调大模型、训练 LSTM 序列模型、或者跑强化学习里的 TD3 这类需要大量交互的算法重跑一次可能就是十几个小时甚至几天的代价。所以这篇文章要解决的核心问题就两个第一怎么让训练进程在断开连接后继续活着第二怎么让训练任务在意外中断后能从最近的进度恢复而不是从头开始。前者靠 tmux 这类终端复用工具做后台保活后者靠 PyTorch 的 checkpoint 机制做断点续训。两个结合起来才能让云端 GPU 任务真正“稳定跑完”。这篇文章适合所有需要在远程 GPU 上跑 PyTorch 训练的人——不管你是刚接触深度学习的新手还是已经在调大模型的老手只要你的训练任务超过半小时这套方案就值得你花二十分钟配置好。1.2 断点续训和后台保活到底在解决什么问题先把问题拆开来看。云端训练任务中断的原因大致可以分成三类连接层中断SSH 会话超时、本地网络波动、笔记本合盖休眠、终端窗口被误关。这类中断的特点是训练进程本身还在跑但因为你跟它之间的“管道”断了进程收到 SIGHUP 信号后被杀掉。进程层中断训练脚本本身报错崩溃OOM、CUDA error、数据加载异常、被系统 OOM killer 杀掉、或者被集群调度器抢占。这类中断是进程真的没了。硬件层中断GPU 掉卡、驱动崩溃、机器重启、租用的实例被回收。这类中断最彻底内存里的所有状态全部丢失。后台保活tmux解决的是第一类问题断点续训checkpoint解决的是第二类和第三类问题。两者配合才能覆盖绝大多数中断场景。我自己的习惯是只要训练预计超过 30 分钟一律在 tmux 里跑并且每隔一定步数或 epoch 存一次 checkpoint。这个习惯帮我省下了无数次重跑的时间。有一次在租用的 GPU 上微调一个模型跑到第 8 个小时的时候实例突然被回收了但因为每 500 步存了一次 checkpoint重新开实例后从最近的 checkpoint 恢复只损失了不到 20 分钟的计算量。2. 整体方案设计与工具选型思路2.1 为什么选 tmux 而不是 nohup 或 screen后台保活的工具主要有三个选择nohup、screen、tmux。我三个都用过最后稳定用 tmux原因如下。nohup 最简单nohup python train.py 就能让进程忽略 SIGHUP 信号在后台跑。但它的问题是你没法方便地“回到”这个进程看实时输出。你只能通过重定向的日志文件去 tail交互性很差。如果训练中途你想看看 GPU 利用率、想手动调个参数、想确认一下当前进度nohup 就很别扭。screen 是 tmux 的老前辈功能上够用但它的分屏和配置体验比 tmux 差不少。tmux 的窗口管理、面板分割、复制模式、脚本化配置都更现代而且几乎每台 Linux 服务器都预装了或者能一行命令装上。tmux 的核心价值在于它在你和训练进程之间维持了一个独立的会话session。你的 SSH 连接只是“附着”到这个会话上断开 SSH 只是“分离”detach会话和里面的进程继续在服务器上跑。下次你 SSH 上来tmux attach就能回到原来的界面看到训练日志还在滚动。这个机制用生活化的类比就是tmux 相当于在服务器上开了一个“虚拟终端房间”你人走了房间里的机器还在运转你回来推门进去就行。2.2 checkpoint 存什么、多久存一次checkpoint 的本质是把训练过程中的关键状态序列化到磁盘上中断后能重新加载回来。一个完整的 PyTorch checkpoint 通常包含以下几部分内容作用是否必须model.state_dict()模型权重参数必须optimizer.state_dict()优化器状态如 Adam 的动量强烈建议scheduler.state_dict()学习率调度器状态建议epoch / global_step当前训练进度必须loss / metric 记录用于日志和恢复判断可选scaler.state_dict()AMP 混合精度缩放器状态用 AMP 时必须random seed 状态保证可复现性可选很多人存 checkpoint 只存model.state_dict()这是不够的。如果你用 Adam 优化器它的动量状态不恢复续训后的前几百步 loss 会明显抖动相当于优化器“失忆”了。学习率调度器同理不恢复的话学习率会从初始值重新开始可能直接破坏已经收敛的状态。存储频率怎么定我的经验是按步数存比按 epoch 存更合理。因为 epoch 的长度可能很长一个 epoch 跑两小时中途崩了就损失两小时。一般设成每 500 到 2000 步存一次具体看单步耗时。如果单步 0.5 秒1000 步就是 8 分钟左右损失可控。同时保留最近 N 个 checkpoint比如 3 个避免磁盘被撑爆。2.3 恢复逻辑的设计从哪个 checkpoint 恢复恢复逻辑要解决一个问题程序启动时怎么知道该从哪个 checkpoint 恢复常见做法有两种。第一种是固定路径覆盖每次存 checkpoint 都覆盖同一个文件比如checkpoint_latest.pt。恢复时直接加载这个文件。优点是简单缺点是如果这个文件在写入过程中崩溃文件可能损坏导致无法恢复。解决办法是“先写临时文件再原子重命名”这个后面会讲。第二种是带步数编号的多文件存成checkpoint_step_1000.pt、checkpoint_step_2000.pt这样恢复时扫描目录找步数最大的那个。优点是安全缺点是文件多、占空间需要定期清理。我一般用混合方案存带步数的文件同时维护一个latest.pt软链接或副本指向最新的。恢复时优先读latest.pt读失败就扫描目录找最大的步数文件。这样兼顾了方便和安全。3. 核心细节解析与实操要点3.1 tmux 会话管理的关键操作tmux 的常用操作其实就几个但每个都有坑我逐个说。创建会话tmux new -s train。-s后面是会话名起个有意义的名字比如train、finetune、rl_td3。不要用默认的编号名不然开多了分不清哪个是哪个。分离会话在 tmux 里按Ctrlb然后按d。注意是先按 Ctrlb 松开再按 d不是同时按。这个快捷键是 tmux 的“前缀键”机制所有 tmux 命令都要先按前缀键。默认前缀是Ctrlb我习惯改成Ctrla因为 b 离得远a 更顺手。改的话在~/.tmux.conf里写set -g prefix C-a。重新附着tmux attach -t train或者简写tmux a -t train。如果只有一个会话直接tmux a就行。查看所有会话tmux ls。会列出会话名、窗口数、创建时间。如果看到某个会话的窗口数是 0说明里面的进程都退出了可以tmux kill-session -t 名字清理掉。注意tmux 会话里的进程是 tmux 的子进程。如果你在 tmux 里跑训练然后又开了一个 tmux 会话两个会话是独立的。关掉一个不影响另一个。但如果你在 tmux 里手动kill了训练进程那进程就真没了tmux 救不了。还有一个容易踩的坑tmux 里的环境变量可能和你登录 shell 的不一样。尤其是 conda 环境有时候tmux new出来的会话没有激活 conda导致python指向系统自带的版本。解决办法是在 tmux 里显式conda activate 你的环境或者在~/.tmux.conf里配置自动激活。我一般是在训练脚本外面套一个 shell 脚本脚本里先source activate再跑 python这样最稳。3.2 checkpoint 保存的原子性与安全性前面提到直接覆盖写 checkpoint 有损坏风险。假设你正在写latest.pt写到一半进程被 kill 了这个文件就是半个残废下次加载直接报错。解决办法是原子写入先写到临时文件写完再os.replace重命名。import os import torch def save_checkpoint(state, path): tmp_path path .tmp torch.save(state, tmp_path) os.replace(tmp_path, path) # 原子操作要么成功要么保持原文件os.replace在同一个文件系统内是原子操作这意味着不会出现“写了一半”的中间状态。要么新文件完整替换旧文件要么旧文件保持不变。这个技巧在存任何重要文件时都适用不只是 checkpoint。另外如果你存的是带步数的多文件记得加一个清理逻辑只保留最近 N 个。不然磁盘满了训练照样崩。清理逻辑很简单列出目录下所有checkpoint_step_*.pt按步数排序删掉最老的几个。import glob import re def cleanup_checkpoints(save_dir, keep3): files glob.glob(os.path.join(save_dir, checkpoint_step_*.pt)) # 从文件名提取步数 def get_step(f): m re.search(rstep_(\d), f) return int(m.group(1)) if m else 0 files.sort(keyget_step) for f in files[:-keep]: os.remove(f)3.3 恢复训练时的状态一致性恢复训练最容易出问题的地方是状态不一致。比如模型权重恢复了但优化器没恢复或者数据加载器的位置没恢复导致重复训练同一批数据。这里逐项说。模型和优化器必须成对恢复。加载的时候用load_state_dict注意strict参数。如果模型结构改过strictTrue会报错这时候要么改回原结构要么用strictFalse但清楚自己在做什么。学习率调度器的恢复经常被忽略。如果你用CosineAnnealingLR或者OneCycleLR不恢复 scheduler 状态的话学习率会从初始值重新走一遍可能直接把模型带偏。恢复方法和模型一样scheduler.load_state_dict(ckpt[scheduler])。数据加载器的恢复比较麻烦。如果你用DataLoader的 shuffle每个 epoch 的数据顺序是随机的没法精确恢复。一般做法是记录当前 epoch 和 epoch 内已完成的步数恢复时跳过已经训练过的数据。或者用sampler设置固定的 seed让每个 epoch 的顺序可复现。我一般用后者配合torch.manual_seed和numpy.random.seed保证可复现性。混合精度训练AMP的GradScaler状态也要存。不存的话恢复后 scaler 的缩放因子从头开始可能导致前几步梯度溢出或者缩放不合适。存法很简单scaler.state_dict()加进去就行。实操心得我习惯在 checkpoint 里额外存一个rng_state包括 Python、NumPy、PyTorch 的随机数状态。这样恢复后连 dropout 的随机性都能接上对于需要严格复现的实验特别有用。虽然大多数时候用不上但存着不亏。4. 完整实操流程与核心环节实现4.1 环境准备与 tmux 配置假设你已经有一台带 GPU 的远程服务器能 SSH 上去conda 环境也配好了。第一步是确认 tmux 装了tmux -V # 如果没装Ubuntu/Debian 下 sudo apt-get install tmux然后配置~/.tmux.conf我常用的配置如下# 改前缀键为 Ctrla set -g prefix C-a unbind C-b bind C-a send-prefix # 开启鼠标支持方便滚动和选面板 set -g mouse on # 设置窗口编号从 1 开始 set -g base-index 1 setw -g pane-base-index 1 # 增大回滚缓冲区方便看历史日志 set -g history-limit 50000 # 状态栏显示更丰富的信息 set -g status-right #[fggreen]#(whoami)#H #[fgyellow]%Y-%m-%d %H:%M配置改完tmux kill-server重启生效。鼠标支持这个特别实用开了之后可以直接用滚轮翻看训练日志不用记复杂的复制模式快捷键。4.2 训练脚本的断点续训改造下面是一个完整的训练脚本骨架包含 checkpoint 保存和恢复逻辑。我以图像分类任务为例其他任务改改数据加载部分就行。import os import glob import re import torch import torch.nn as nn from torch.utils.data import DataLoader from torch.optim import Adam from torch.optim.lr_scheduler import CosineAnnealingLR from torch.cuda.amp import GradScaler, autocast # 配置 SAVE_DIR ./checkpoints os.makedirs(SAVE_DIR, exist_okTrue) SAVE_EVERY 1000 # 每 1000 步存一次 KEEP_LAST 3 # 保留最近 3 个 checkpoint RESUME True # 是否自动恢复 device torch.device(cuda if torch.cuda.is_available() else cpu) # 模型、优化器、调度器 model MyModel().to(device) optimizer Adam(model.parameters(), lr1e-4) scheduler CosineAnnealingLR(optimizer, T_max10000) scaler GradScaler() # 恢复逻辑 start_step 0 start_epoch 0 def find_latest_checkpoint(save_dir): latest os.path.join(save_dir, latest.pt) if os.path.exists(latest): return latest files glob.glob(os.path.join(save_dir, checkpoint_step_*.pt)) if not files: return None def get_step(f): m re.search(rstep_(\d), f) return int(m.group(1)) if m else 0 files.sort(keyget_step) return files[-1] if RESUME: ckpt_path find_latest_checkpoint(SAVE_DIR) if ckpt_path: print(fResuming from {ckpt_path}) ckpt torch.load(ckpt_path, map_locationdevice) model.load_state_dict(ckpt[model]) optimizer.load_state_dict(ckpt[optimizer]) scheduler.load_state_dict(ckpt[scheduler]) scaler.load_state_dict(ckpt[scaler]) start_step ckpt[step] start_epoch ckpt[epoch] print(fResumed at step {start_step}, epoch {start_epoch}) else: print(No checkpoint found, training from scratch) # 保存函数 def save_checkpoint(step, epoch): state { model: model.state_dict(), optimizer: optimizer.state_dict(), scheduler: scheduler.state_dict(), scaler: scaler.state_dict(), step: step, epoch: epoch, } # 带步数的文件 path os.path.join(SAVE_DIR, fcheckpoint_step_{step}.pt) tmp path .tmp torch.save(state, tmp) os.replace(tmp, path) # 更新 latest latest os.path.join(SAVE_DIR, latest.pt) tmp_latest latest .tmp torch.save(state, tmp_latest) os.replace(tmp_latest, latest) # 清理旧文件 cleanup_checkpoints(SAVE_DIR, KEEP_LAST) print(fCheckpoint saved at step {step}) def cleanup_checkpoints(save_dir, keep): files glob.glob(os.path.join(save_dir, checkpoint_step_*.pt)) def get_step(f): m re.search(rstep_(\d), f) return int(m.group(1)) if m else 0 files.sort(keyget_step) for f in files[:-keep]: os.remove(f) # 训练循环 global_step start_step for epoch in range(start_epoch, NUM_EPOCHS): for batch in train_loader: if global_step start_step: global_step 1 continue # 跳过已训练的步 model.train() optimizer.zero_grad() with autocast(): loss model(batch) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() scheduler.step() global_step 1 if global_step % SAVE_EVERY 0: save_checkpoint(global_step, epoch) # epoch 结束也存一次 save_checkpoint(global_step, epoch 1)这个脚本有几个关键点。第一恢复时start_step之前的步会被跳过避免重复训练。第二latest.pt和带步数的文件都存双保险。第三清理逻辑保证磁盘不会爆。第四AMP 的 scaler 状态也存了混合精度训练能无缝续上。4.3 在 tmux 里启动训练并验证保活脚本准备好后完整启动流程如下# 1. SSH 登录服务器 ssh usergpu-server # 2. 创建 tmux 会话 tmux new -s train # 3. 在 tmux 里激活环境 conda activate myenv # 4. 确认 GPU 可用 nvidia-smi # 5. 启动训练 python train.py 21 | tee train.log # 6. 按 Ctrla 然后按 d 分离会话分离之后你可以直接关掉 SSH 终端甚至关掉本地电脑。训练进程在服务器上继续跑。下次想看进度ssh usergpu-server tmux attach -t train就能回到训练界面看到日志还在滚动。如果想在附着状态下快速看一眼 GPU 利用率可以按Ctrla然后按%分一个垂直面板在面板里跑watch -n 1 nvidia-smi这样一边看日志一边看 GPU 状态。验证保活是否生效可以做个简单测试启动一个每 5 秒打印一次时间的脚本分离会话关掉 SSH等几分钟再连回来 attach看时间是否连续。如果连续说明保活成功。注意tee train.log这个操作很重要。它把输出同时写到终端和日志文件。万一 tmux 会话因为某种原因挂了你还能从日志文件里看到最后的输出判断训练到哪一步了。我一般还会在训练脚本里用logging模块单独写一份结构化日志方便后续分析。5. 常见问题与排查技巧实录5.1 tmux 相关的高频问题问题一tmux attach报错 “no sessions”。这通常是因为会话已经退出了。原因可能是训练进程崩溃导致 tmux 里没有活动进程tmux 自动关闭了会话。解决办法是tmux ls确认如果确实没了检查训练日志找崩溃原因。预防措施是在 tmux 里跑一个不会退出的 shell比如训练脚本外面套bash -c python train.py; bash这样即使 python 挂了bash 还在会话不会消失你能进去看现场。问题二tmux 里中文乱码。这是 locale 设置问题。在~/.tmux.conf里加set -g default-terminal screen-256color并在 shell 的~/.bashrc里设置export LANGen_US.UTF-8或zh_CN.UTF-8。如果服务器没装中文 locale用英文也行日志里的中文可能显示成方块但不影响训练。问题三tmux 会话里的进程看不到 GPU。有时候nvidia-smi在 tmux 里报 “No devices found”但在外面正常。这通常是环境变量CUDA_VISIBLE_DEVICES没传进去。解决办法是在 tmux 里显式export CUDA_VISIBLE_DEVICES0或者检查~/.bashrc里的相关设置是否在 tmux 启动时被加载。5.2 checkpoint 加载失败的排查问题一RuntimeError: Error(s) in loading state_dict。这是模型结构不匹配。常见原因是你改了模型定义但想加载旧 checkpoint。解决办法是用strictFalse加载然后手动检查哪些层没加载上。或者打印两边的state_dict的 key 对比。model_dict model.state_dict() ckpt_dict ckpt[model] # 找出不匹配的 key missing [k for k in model_dict if k not in ckpt_dict] unexpected [k for k in ckpt_dict if k not in model_dict] print(Missing:, missing) print(Unexpected:, unexpected)问题二CUDA out of memory在恢复后出现。这通常是因为恢复时同时加载了模型和优化器状态到 GPU显存占用比训练时高。解决办法是先用map_locationcpu加载再逐个移到 GPU。或者恢复完成后手动torch.cuda.empty_cache()。问题三恢复后 loss 突然飙升。这是状态不一致的典型表现。检查优化器、调度器、scaler 是否都恢复了。如果都恢复了还飙升可能是数据加载器的位置不对重复训练了某些数据。检查start_step的跳过逻辑是否正确。5.3 常见问题速查表现象可能原因排查方法解决SSH 断开后进程消失没用 tmux/nohupps aux | grep python用 tmux 重跑tmux attach 无会话会话已退出tmux ls检查日志套 bash 保活checkpoint 加载报错结构不匹配对比 state_dict keystrictFalse 或改回结构恢复后 loss 抖动优化器状态没恢复检查 ckpt 内容补上 optimizer/scheduler磁盘写满checkpoint 太多du -sh checkpoints/加清理逻辑GPU 不可见环境变量问题echo $CUDA_VISIBLE_DEVICES显式 export训练速度变慢数据加载瓶颈nvidia-smi看利用率加 num_workers5.4 几个我踩过的坑第一个坑是在 tmux 里用Ctrlc中断训练。这个操作会直接杀掉 python 进程tmux 会话还在但训练没了。正确的做法是让训练脚本自己处理信号收到 SIGINT 时先存 checkpoint 再退出。可以在脚本里注册信号处理import signal import sys def signal_handler(sig, frame): print(Interrupted, saving checkpoint...) save_checkpoint(global_step, current_epoch) sys.exit(0) signal.signal(signal.SIGINT, signal_handler)这样按Ctrlc会先存 checkpoint 再退出不会丢进度。第二个坑是checkpoint 存到网络文件系统NFS上。有些集群的 home 目录是 NFS写入速度慢而且os.replace在 NFS 上不一定是原子的。解决办法是把 checkpoint 存到本地磁盘训练完再拷贝到 NFS。或者至少确认 NFS 支持原子重命名。第三个坑是恢复时忘了设置随机种子。如果训练里有 dropout、数据 shuffle 等随机操作不设种子的话每次恢复后的行为都不一样实验没法复现。我一般在脚本开头就设好import random import numpy as np def set_seed(seed42): random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed) set_seed(42)第四个坑是tmux 会话名冲突。如果你开了多个训练任务都用train这个名字tmux new -s train会报错说会话已存在。解决办法是用带任务标识的名字比如train_resnet、train_lstm、rl_td3。或者用tmux new -s train_$(date %m%d_%H%M)自动加时间戳。6. 进阶技巧与规模化建议6.1 多任务并行时的 tmux 窗口管理当你同时跑多个训练任务时一个 tmux 会话里可以开多个窗口。Ctrla然后c创建新窗口Ctrla然后n/p切换下一个/上一个窗口Ctrla然后数字键直接跳到对应窗口。每个窗口独立跑一个训练互不干扰。我一般这样组织窗口 0 是主训练任务窗口 1 是验证/评估任务窗口 2 是 TensorBoard 或者日志监控窗口 3 是备用 shell 用来跑临时命令。这样所有相关的东西都在一个会话里attach 一次全能看到。如果任务特别多可以按项目分会话。比如tmux new -s project_a跑 A 项目的所有任务tmux new -s project_b跑 B 项目的。tmux ls一眼看清所有项目状态。6.2 自动监控与异常重启tmux 保活解决的是连接断开问题但如果训练进程本身崩溃了tmux 不会自动重启它。对于需要长时间无人值守的任务可以加一层自动重启逻辑。最简单的做法是写一个 shell 循环while true; do python train.py exit_code$? if [ $exit_code -eq 0 ]; then echo Training completed successfully break else echo Training crashed with code $exit_code, restarting in 30s... sleep 30 fi done这个循环放在 tmux 里跑。训练正常结束exit code 0就退出循环异常崩溃就等 30 秒重启重启后脚本会自动从最近的 checkpoint 恢复。这样就实现了“崩溃自愈”。更完善的方案是加一个监控脚本定期检查训练进程是否还在、GPU 是否还在工作、日志是否还在更新。如果发现异常发邮件或发消息通知。不过对于大多数个人项目上面的 shell 循环已经够用了。6.3 checkpoint 的版本管理与实验追踪当你跑了很多组实验checkpoint 文件会越来越多容易搞混哪个是哪个。我的做法是在 checkpoint 目录里放一个meta.json记录这次训练的超参数、git commit、开始时间等信息。每次存 checkpoint 时更新一下。import json import time meta { lr: 1e-4, batch_size: 32, model: resnet50, start_time: time.strftime(%Y-%m-%d %H:%M:%S), git_commit: os.popen(git rev-parse HEAD).read().strip(), } with open(os.path.join(SAVE_DIR, meta.json), w) as f: json.dump(meta, f, indent2)这样即使过了几个月回头看也能知道每个 checkpoint 对应的实验配置。如果配合 TensorBoard 或 WandB 这类工具把 checkpoint 路径和实验记录关联起来管理起来更清晰。6.4 从单卡到多卡的注意事项如果你从单卡扩展到多卡DDPcheckpoint 的保存和恢复有几个额外注意点。保存时只需要在主进程rank 0保存避免多个进程同时写同一个文件。恢复时每个进程都要加载但map_location要设成对应的 GPU。if dist.get_rank() 0: save_checkpoint(...) # 恢复时 ckpt torch.load(path, map_locationfcuda:{local_rank}) model.load_state_dict(ckpt[model])另外 DDP 的模型是DistributedDataParallel包装的保存时要model.module.state_dict()加载时先加载到原始模型再包装。这些细节不注意的话多卡续训很容易出问题。我个人在实际操作中的体会是断点续训这套东西配置一次可能花半小时但它带来的安心感是巨大的。以前跑长任务总是提心吊胆怕断线怕崩溃现在设好 tmux 和 checkpoint该睡觉睡觉该出门出门回来 attach 一下看进度就行。尤其是租用 GPU 按小时计费的时候能续训意味着不用为中断的那部分时间重复付费省下的都是真金白银。最后再分享一个小技巧把常用的 tmux 启动命令和训练启动命令写成一个start.sh脚本每次新任务直接bash start.sh省得每次手敲一堆命令还容易敲错。
网站建设高端定制企业官网