1. 云端训练任务为什么总是“跑不完”
1.1 一个让所有炼丹师都头疼的场景
你花了大半天时间调试模型结构、清洗数据、配置环境,终于把训练脚本跑通了。看着终端里 loss 一点点下降,你心满意足地去吃饭、睡觉,或者关掉笔记本去忙别的事。结果第二天早上打开终端一看——连接断了,进程没了,日志停在半夜两点半,loss 曲线戛然而止。
这种情况在云端 GPU 训练里太常见了。不管是租用的 GPU 服务器、实验室的共享集群,还是公司内部的训练平台,只要训练任务超过几个小时,就一定会遇到连接中断、会话超时、进程被意外杀掉的问题。尤其是当你用 SSH 连到远程机器上直接跑python train.py的时候,网络一抖动,终端一关闭,SIGHUP 信号就把你的训练进程带走了。
更让人崩溃的是,有些任务已经跑了十几个 epoch,checkpoint 也没存,一切从头再来。如果是小模型还好,几十分钟能重跑;但如果是在微调大模型、训练 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 里按Ctrl+b然后按d。注意是先按 Ctrl+b 松开,再按 d,不是同时按。这个快捷键是 tmux 的“前缀键”机制,所有 tmux 命令都要先按前缀键。默认前缀是Ctrl+b,我习惯改成Ctrl+a,因为 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, keep=3): files = glob.glob(os.path.join(save_dir, "checkpoint_step_*.pt")) # 从文件名提取步数 def get_step(f): m = re.search(r"step_(\d+)", f) return int(m.group(1)) if m else 0 files.sort(key=get_step) for f in files[:-keep]: os.remove(f)3.3 恢复训练时的状态一致性
恢复训练最容易出问题的地方是状态不一致。比如模型权重恢复了,但优化器没恢复;或者数据加载器的位置没恢复,导致重复训练同一批数据。这里逐项说。
模型和优化器必须成对恢复。加载的时候用load_state_dict,注意strict参数。如果模型结构改过,strict=True会报错,这时候要么改回原结构,要么用strict=False但清楚自己在做什么。
学习率调度器的恢复经常被忽略。如果你用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,我常用的配置如下:
# 改前缀键为 Ctrl+a 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 "#[fg=green]#(whoami)@#H #[fg=yellow]%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_ok=True) 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(), lr=1e-4) scheduler = CosineAnnealingLR(optimizer, T_max=10000) 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(r"step_(\d+)", f) return int(m.group(1)) if m else 0 files.sort(key=get_step) return files[-1] if RESUME: ckpt_path = find_latest_checkpoint(SAVE_DIR) if ckpt_path: print(f"Resuming from {ckpt_path}") ckpt = torch.load(ckpt_path, map_location=device) 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(f"Resumed 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, f"checkpoint_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(f"Checkpoint 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(r"step_(\d+)", f) return int(m.group(1)) if m else 0 files.sort(key=get_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 user@gpu-server # 2. 创建 tmux 会话 tmux new -s train # 3. 在 tmux 里激活环境 conda activate myenv # 4. 确认 GPU 可用 nvidia-smi # 5. 启动训练 python train.py 2>&1 | tee train.log # 6. 按 Ctrl+a 然后按 d 分离会话分离之后,你可以直接关掉 SSH 终端,甚至关掉本地电脑。训练进程在服务器上继续跑。下次想看进度:
ssh user@gpu-server tmux attach -t train就能回到训练界面,看到日志还在滚动。如果想在附着状态下快速看一眼 GPU 利用率,可以按Ctrl+a然后按%分一个垂直面板,在面板里跑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 LANG=en_US.UTF-8或zh_CN.UTF-8。如果服务器没装中文 locale,用英文也行,日志里的中文可能显示成方块,但不影响训练。
问题三:tmux 会话里的进程看不到 GPU。有时候nvidia-smi在 tmux 里报 “No devices found”,但在外面正常。这通常是环境变量CUDA_VISIBLE_DEVICES没传进去。解决办法是在 tmux 里显式export CUDA_VISIBLE_DEVICES=0,或者检查~/.bashrc里的相关设置是否在 tmux 启动时被加载。
5.2 checkpoint 加载失败的排查
问题一:RuntimeError: Error(s) in loading state_dict。这是模型结构不匹配。常见原因是你改了模型定义但想加载旧 checkpoint。解决办法是用strict=False加载,然后手动检查哪些层没加载上。或者打印两边的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_location="cpu"加载,再逐个移到 GPU。或者恢复完成后手动torch.cuda.empty_cache()。
问题三:恢复后 loss 突然飙升。这是状态不一致的典型表现。检查优化器、调度器、scaler 是否都恢复了。如果都恢复了还飙升,可能是数据加载器的位置不对,重复训练了某些数据。检查start_step的跳过逻辑是否正确。
5.3 常见问题速查表
| 现象 | 可能原因 | 排查方法 | 解决 |
|---|---|---|---|
| SSH 断开后进程消失 | 没用 tmux/nohup | ps aux | grep python | 用 tmux 重跑 |
| tmux attach 无会话 | 会话已退出 | tmux ls | 检查日志,套 bash 保活 |
| checkpoint 加载报错 | 结构不匹配 | 对比 state_dict key | strict=False 或改回结构 |
| 恢复后 loss 抖动 | 优化器状态没恢复 | 检查 ckpt 内容 | 补上 optimizer/scheduler |
| 磁盘写满 | checkpoint 太多 | du -sh checkpoints/ | 加清理逻辑 |
| GPU 不可见 | 环境变量问题 | echo $CUDA_VISIBLE_DEVICES | 显式 export |
| 训练速度变慢 | 数据加载瓶颈 | nvidia-smi看利用率 | 加 num_workers |
5.4 几个我踩过的坑
第一个坑是在 tmux 里用Ctrl+c中断训练。这个操作会直接杀掉 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)这样按Ctrl+c会先存 checkpoint 再退出,不会丢进度。
第二个坑是checkpoint 存到网络文件系统(NFS)上。有些集群的 home 目录是 NFS,写入速度慢,而且os.replace在 NFS 上不一定是原子的。解决办法是把 checkpoint 存到本地磁盘,训练完再拷贝到 NFS。或者至少确认 NFS 支持原子重命名。
第三个坑是恢复时忘了设置随机种子。如果训练里有 dropout、数据 shuffle 等随机操作,不设种子的话每次恢复后的行为都不一样,实验没法复现。我一般在脚本开头就设好:
import random import numpy as np def set_seed(seed=42): 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 会话里可以开多个窗口。Ctrl+a然后c创建新窗口,Ctrl+a然后n/p切换下一个/上一个窗口,Ctrl+a然后数字键直接跳到对应窗口。每个窗口独立跑一个训练,互不干扰。
我一般这样组织:窗口 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, indent=2)这样即使过了几个月回头看,也能知道每个 checkpoint 对应的实验配置。如果配合 TensorBoard 或 WandB 这类工具,把 checkpoint 路径和实验记录关联起来,管理起来更清晰。
6.4 从单卡到多卡的注意事项
如果你从单卡扩展到多卡(DDP),checkpoint 的保存和恢复有几个额外注意点。保存时只需要在主进程(rank 0)保存,避免多个进程同时写同一个文件。恢复时每个进程都要加载,但map_location要设成对应的 GPU。
if dist.get_rank() == 0: save_checkpoint(...) # 恢复时 ckpt = torch.load(path, map_location=f"cuda:{local_rank}") model.load_state_dict(ckpt["model"])另外 DDP 的模型是DistributedDataParallel包装的,保存时要model.module.state_dict(),加载时先加载到原始模型再包装。这些细节不注意的话,多卡续训很容易出问题。
我个人在实际操作中的体会是,断点续训这套东西,配置一次可能花半小时,但它带来的安心感是巨大的。以前跑长任务总是提心吊胆,怕断线怕崩溃,现在设好 tmux 和 checkpoint,该睡觉睡觉,该出门出门,回来 attach 一下看进度就行。尤其是租用 GPU 按小时计费的时候,能续训意味着不用为中断的那部分时间重复付费,省下的都是真金白银。最后再分享一个小技巧:把常用的 tmux 启动命令和训练启动命令写成一个start.sh脚本,每次新任务直接bash start.sh,省得每次手敲一堆命令还容易敲错。