微调训练中断断点续训,训练任务容错处理
微调训练中断断点续训,本质是让训练框架在中断前保存完整的模型权重、优化器状态和训练步数,重启后从最近一次 checkpoint 继续训练,而不是从头再来。
这篇文章会从 checkpoint 保存策略讲起,给出可直接执行的恢复命令、多卡训练常见中断原因和验证恢复是否成功的方法,适合正在跑大模型微调、希望减少算力浪费的开发者照做。
先搞清楚 checkpoint 里该存什么
断点续训能否恢复正常,取决于保存内容是否完整。
只存模型权重,只能恢复推理,不能恢复训练。
要让训练真正接得上,checkpoint 里至少要包含:
- 模型权重:包括主干和所有可训练参数。
- 优化器状态:如 Adam 的动量与方差,否则恢复后学习率调度会异常。
- 训练步数 step:决定剩余训练轮次和日志对齐。
- 学习率调度器状态:避免重启后学习率跳变。
- 随机种子状态:对追求可复现的实验很重要。
以 Hugging Face 的 Trainer 为例,训练代码里开启 checkpoint 保存的最小配置如下:
from transformers import Trainer, TrainingArguments
training_args = TrainingArguments(
output_dir="./output/checkpoints",
save_strategy="steps",
save_steps=500,
save_total_limit=3,
load_best_model_at_end=True,
metric_for_best_model="eval_loss",
logging_steps=50,
)
trainer = Trainer(
model=model,
args=training_args,
train_dataset=train_dataset,
eval_dataset=eval_dataset,
)
其中 save_total_limit=3 表示磁盘只保留最近 3 份 checkpoint,
避免占满存储;save_steps=500 是每 500 步存一次,
间隔可以根据训练集大小调整。
恢复训练:从 checkpoint 拉起来接着跑
训练中断后,不用改原训练脚本,只要在启动命令或 TrainingArguments 中指定断点位置即可。
方式一:Trainer 自动恢复
如果使用的是 Trainer,在训练脚本里加一个参数:
trainer.train(resume_from_checkpoint=True)
传入 True 时,框架会自动读取 output_dir 下最新的 checkpoint。
也可以手动指定:
trainer.train(resume_from_checkpoint="./output/checkpoints/checkpoint-5000")
方式二:命令行传入 checkpoint 目录
python train.py --model_name /data/base_model \
--resume_from_checkpoint ./output/checkpoints/checkpoint-5000
训练脚本里对应的解析逻辑通常是:
parser.add_argument("--resume_from_checkpoint", type=str, default=None)
...
if args.resume_from_checkpoint:
trainer.train(resume_from_checkpoint=args.resume_from_checkpoint)
重启后,
日志里会出现 Continuing training from epoch 或 Loading model from checkpoint-5000 的提示,
说明已经读取到断点。
多卡与平台场景下的容错设置
多卡训练时,任意一张卡掉线都会导致 NCCL 通信失败,任务直接退出。
这类中断单靠重启命令往往不够,还需要配合容错机制。
1. NCCL 超时与自动恢复
设置环境变量,让通信失败时更快暴露并重启:
export NCCL_DEBUG=INFO
export NCCL_IB_TIMEOUT=22
export NCCL_IB_RETRY_CNT=7
2. 借助训练框架的容错功能
PyTorch 2.x 支持 torch.distributed.elastic,配合 torchrun 可以在 worker 进程异常退出后自动拉起新进程:
torchrun --nnodes=1 --nproc_per_node=8 --max_restarts=3 \
--rdzv_backend=c10d \
train.py --resume_from_checkpoint ./output/checkpoints/checkpoint-5000
--max_restarts=3 表示最多自动重启 3 次,重启后如果脚本支持 checkpoint 恢复,就能从最近的保存点继续。
3. 云平台/任务队列中的续跑思路
如果使用云服务器或任务调度平台跑训练,建议把训练脚本改造成“启动时自动寻找最新 checkpoint”,这样节点被回收或重启后,无需人工指定目录。
常见做法是把检查逻辑写在训练入口处:
import os, glob
def find_latest_checkpoint(output_dir):
checkpoints = glob.glob(os.path.join(output_dir, "checkpoint-*"))
if not checkpoints:
return None
# 按 checkpoint 编号取最大
return max(checkpoints, key=lambda x: int(x.split("-")[-1]))
ckpt = find_latest_checkpoint("./output/checkpoints")
trainer.train(resume_from_checkpoint=ckpt if ckpt else False)
这样即使半夜训练中断,第二天重跑脚本也能自动续上。
避坑:这些细节会让断点续训失败
断点续训配置简单,但实际恢复时经常遇到以下问题,列出来帮你提前规避。
- 只保存了模型权重,没保存优化器状态,恢复后 loss 波动很大。检查 checkpoint 目录下是否包含
optimizer.pt和scheduler.pt文件。 - checkpoint 写入时进程被杀,目录不完整。启动脚本里留意保存时是否有
.tmp临时文件残留,建议先写临时目录再os.replace原子替换。 - 代码或数据顺序发生变更。断点续训默认继续跑后续数据,如果训练集顺序被打乱,数据会错位,微调效果会受影响。恢复前不要随意改动数据集列表。
- 换了显卡数量。从 8 卡改成 4 卡后,batch size、学习率、梯度累积步数都要重新计算,否则 loss 可能直接飞掉。
- 磁盘空间被
save_total_limit清理逻辑误删。确认保存目录挂载的磁盘有足够余量,建议预留模型大小 10 倍以上的空间。
如何确认续训真的成功了
训练重启后别只盯日志,按下面三步验证:
- 查看启动日志,确认出现
Resuming training from checkpoint或Continuing training from epoch X字样。 - 保存一份 loss 曲线,对比中断前最后 20 步和恢复后前 20 步的 loss,正常情况应在同一水平线附近,不会突然升高。
- 定期在 checkpoint 目录运行
ls -lh,确认生成时间与保存周期匹配。
另外可以做一个中断恢复演练:训练到 checkpoint 保存后,手动 kill 掉进程,再按本文的恢复命令重新启动。
如果 loss 曲线能平滑衔接,说明你的微调训练断点续训和容错处理已经可用。
如果你在配置恢复命令时遇到 checkpoint 加载失败或数据顺序错乱的问题,优先检查保存内容是否完整和数据集是否变更。
熟练之后,可以把查找最新 checkpoint 和自动重启逻辑写进训练脚本,进一步降低人工干预成本。