共计 2688 个字符,预计需要花费 7 分钟才能阅读完成。
问题背景
最近在 Autodl 上跑深度学习模型时,经常遇到训练中断的情况。最常见的有两种场景:

- 本地电脑休眠或网络波动导致 SSH 连接断开
- 云主机因资源不足被自动释放(虽然 Autodl 有保护机制,但极端情况仍会发生)
这些中断带来的实际影响很严重:
- 模型训练进度丢失,需要从头开始
- 数据迭代器(DataLoader)状态重置,影响数据一致性
- 已计算的特征图等中间结果无法复用
技术方案对比
先看几种常见的连接保持方案:
原生 SSH 直连
最简单但最不可靠的方式,断开后:
- 所有子进程会被终止(SIGHUP)
- 需要手动重新登录并启动训练
tmux/screen 会话托管
Linux 终端复用工具,核心优势:
- 会话(session)与窗口(window)分离
- 断开连接后进程仍在后台运行
- 支持随时重新接入(reattach)
端口转发持久化
适合 Jupyter 等可视化工具:
- 通过 SSH 隧道将云主机端口映射到本地
- 配合 autossh 自动维护连接
- 浏览器直接访问 localhost:port
核心组件配置
SSH 自动保持连接
修改 ~/.ssh/config 文件(没有就新建):
Host autodl
HostName your-instance-address
User root
Port 22
ServerAliveInterval 60
ServerAliveCountMax 3
参数说明:
ServerAliveInterval 60:每 60 秒发送一次心跳包ServerAliveCountMax 3:连续 3 次无响应才断开
tmux 基础用法
启动新会话并命名:
tmux new -s model_train
临时断开会话(保持进程运行):
# 在 tmux 会话中按 Ctrl+b 然后按 d
重新接入已有会话:
tmux attach -t model_train
代码实现
tmux 初始化脚本
创建start_train.sh:
#!/bin/bash
session="model_train"
tmux has-session -t $session 2>/dev/null
if [$? != 0]; then
tmux new-session -d -s $session
tmux rename-window -t $session:1 "main"
tmux split-window -h -t $session:1
tmux send-keys -t $session:1.1 "watch -n 1 nvidia-smi" C-m
tmux send-keys -t $session:1.2 "python train.py" C-m
fi
tmux attach -t $session
Jupyter 重连技巧
- 查找运行中的 Jupyter 进程:
ps aux | grep jupyter
- 获取类似这样的输出:
root 12345 0.3 0.8 212345 67890 ? S 14:30 0:05 /usr/bin/python3 /usr/local/bin/jupyter-lab --no-browser --port=8888
- 通过端口转发重新建立连接:
ssh -NfL 8888:localhost:8888 autodl
Python 训练检查点
在训练代码中添加自动保存逻辑:
import torch
from contextlib import contextmanager
@contextmanager
def training_context(model, optimizer, ckpt_path):
try:
# 加载已有检查点
if os.path.exists(ckpt_path):
state = torch.load(ckpt_path)
model.load_state_dict(state['model'])
optimizer.load_state_dict(state['optimizer'])
start_epoch = state['epoch'] + 1
print(f"Resuming from epoch {start_epoch}")
else:
start_epoch = 0
yield start_epoch
finally:
# 确保异常时也能保存
state = {'model': model.state_dict(),
'optimizer': optimizer.state_dict(),
'epoch': start_epoch
}
torch.save(state, ckpt_path)
print(f"Checkpoint saved at {ckpt_path}")
# 使用示例
for epoch in range(epochs):
with training_context(model, optimizer, 'checkpoint.pth') as start_epoch:
train_one_epoch(model, train_loader, epoch)
生产级优化
内存管理
- 创建 swap 文件(当物理内存不足时):
sudo fallocate -l 8G /swapfile
sudo chmod 600 /swapfile
sudo mkswap /swapfile
sudo swapon /swapfile
- 添加到
/etc/fstab实现开机自动挂载:
/swapfile none swap sw 0 0
后台进程守护
使用 nohup 防止 SSH 断开导致进程终止:
nohup python train.py > train.log 2>&1 &
查看运行状态:
tail -f train.log
资源监控面板
实时查看 GPU 状态:
watch -n 1 "nvidia-smi && free -h"
避坑指南
- 认证方式:
- 务必使用 SSH 密钥登录,避免密码会话超时
-
生成密钥对:
ssh-keygen -t ed25519 -
系统更新:
- 禁用自动更新:
sudo apt-mark hold cuda* nvidia* -
手动更新前先创建系统快照
-
多卡训练:
- 明确指定可见 GPU:
import os os.environ["CUDA_VISIBLE_DEVICES"] = "0,1" # 使用前两块 GPU
延伸思考
设计断点续训框架的几个核心维度:
- 状态完整性:
- 模型参数(model.state_dict())
- 优化器状态(optimizer.state_dict())
- 随机数种子(保证数据 shuffle 可复现)
-
DataLoader 的迭代位置
-
恢复策略:
graph LR A[检测中断] --> B{有检查点?} B -->| 是 | C[加载检查点] B -->| 否 | D[从零开始] C --> E[验证数据一致性] E --> F[继续训练] -
监控体系:
- 心跳检测(定期写入时间戳)
- 资源预警(GPU 显存、磁盘空间)
- 自动报警(邮件 /Slack 通知)
建议将上述功能封装为训练框架的基类,通过继承方式实现业务逻辑与容灾机制的分离。
正文完
