%%{init: {'theme': 'base'}}%%
graph TB
A["硬件故障<br/>GPU/HBM/NIC/线缆"]
B["软件缺陷<br/>OOM/Nan/死锁"]
C["基础设施<br/>网络/存储/电力"]
D["数据问题<br/>坏样本/分布偏移"]
E["人为操作<br/>误删/配置错误"]
A --> F["训练中断"]
B --> F
C --> F
D --> F
E --> F
训练可靠性
训练可靠性
训练一个万亿参数模型需要数周甚至数月。在第 21 天崩溃不是一个小问题——它是数百万美元的损失。
章节导入:为什么可靠性是大模型训练的核心挑战
2024 年的一项调研显示,超过 70% 的大规模训练任务至少经历过一次非计划中断。Meta 在训练 Llama 3 时报告了 419 次意外中断,平均每 3 小时一次。
训练可靠性不是”锦上添花”的运维问题。它是系统设计的核心约束,决定了你的模型能不能在预算和时间表内完成训练。这一章我们将深入故障模式、容错机制和恢复策略。
7.1 常见故障模式与根因分析
大规模训练的故障来自多个层面。理解每种故障的症状和根因,是快速恢复的前提。
故障分类金字塔
各类故障详解
| 故障类型 | 典型症状 | 发生频率 | 恢复时间 |
|---|---|---|---|
| GPU 掉卡 | NCCL timeout / CUDA error | 最频繁(每周 1-2 次) | 5-30 分钟 |
| HBM ECC 错误 | Silent data corruption | 每月 1-2 次 | 需要排查 |
| IB 网络抖动 | AllReduce 超时 | 每周数次 | 自恢复或重启 |
| 存储 I/O 阻塞 | Checkpoint 写入超时 | 每周 1 次 | 10-60 分钟 |
| OOM | CUDA out of memory | 调试期频繁 | 重启即恢复 |
| 梯度爆炸 (NaN) | Loss 变 NaN | 每天数次(初期) | 回滚学习率 |
| 数据管道阻塞 | DataLoader 超时 | 偶发 | 5-10 分钟 |
GPU 故障诊断流程
#!/bin/bash
# GPU 健康检查脚本
echo "=== GPU 状态检查 ==="
nvidia-smi --query-gpu=index,name,temperature.gpu,utilization.gpu,memory.used,memory.total,ecc.errors.corrected.volatile.total,ecc.errors.uncorrected.volatile.total --format=csv
echo -e "\n=== 详细 ECC 错误 ==="
nvidia-smi -q | grep -A5 "ECC Errors"
echo -e "\n=== NCCL 连通性测试 ==="
# 需要在每台节点上运行
python -c "
import torch.distributed as dist
dist.init_process_group(backend='nccl')
dist.barrier()
print(f'Rank {dist.get_rank()}: NCCL OK')
"
echo -e "\n=== IB 网络状态 ==="
ibstat
ibping -S # 检查 IB 端口状态
echo -e "\n=== GPU 到 GPU 带宽测试 ==="
# 节点内 NVLink 带宽
all_reduce_bandwidth # 期望 > 800 GB/s
echo -e "\n=== DCGM 诊断 ==="
dcgmi diag -r 2 # NVIDIA DCGM 完整诊断静默数据损坏(Silent Data Corruption, SDC)
这是最危险的故障类型:GPU 计算出错但不报错,结果悄悄地变成垃圾。如果没有数值检查,Loss 会莫名其妙地停止下降或者变成 NaN。
# SDC 检测:使用冗余计算验证
def check_sdc(model, input_batch, reference_output=None):
"""通过两次前向传播检测 SDC"""
with torch.no_grad():
output1 = model(input_batch)
output2 = model(input_batch)
# 两次前向结果应该完全一致
diff = (output1 - output2).abs().max()
if diff > 1e-5:
# 硬件错误!需要隔离这张 GPU
rank = dist.get_rank() if dist.is_initialized() else 0
logging.error(f"SDC detected on rank {rank}! Max diff: {diff}")
# 上报并触发节点隔离
report_faulty_gpu(rank)
raise RuntimeError(f"SDC on rank {rank}")Google 的研究(2024)显示,在大规模 GPU 集群中,SDC 的发生率约为每 1000 张 GPU 每天发生 1-2 次。对于 10000+ 卡的集群,这意味着几乎每天都有 SDC 事件。不检测 SDC 就是拿真金白银在赌。
7.2 Checkpoint 策略
Checkpoint 是训练可靠性的最后一道防线。好的 checkpoint 策略能在最小开销下实现最快的恢复。
同步 vs 异步 Checkpoint
| 方式 | 原理 | 优点 | 缺点 |
|---|---|---|---|
| 同步 | 训练暂停,等待所有 rank 写完 | 简单、一致性强 | 训练停顿 10-60 秒 |
| 异步 | 后台线程复制 + 保存 | 不阻塞训练 | 实现复杂、内存翻倍 |
| 分布式 | 每个 rank 只写自己的分片 | 快、支持异构恢复 | 恢复时需要全量收集 |
分布式 Checkpoint(推荐)
import torch.distributed.checkpoint as dcp
from torch.distributed.checkpoint import FileSystemWriter, FileSystemReader
def save_checkpoint_distributed(model, optimizer, step, path):
"""PyTorch 2.x 分布式 Checkpoint"""
# FSDP/Sharded 模式下,每个 rank 只保存自己的分片
state_dict = {
"model": model.state_dict(),
"optimizer": FSDP.optim_state_dict(model, optimizer),
"step": step,
"rng_state": torch.get_rng_state(),
"cuda_rng_state": torch.cuda.get_rng_state(),
}
# 每个 rank 并行写入自己的分片
dcp.save(
state_dict=state_dict,
checkpoint_id=path,
storage_writer=FileSystemWriter(path=path),
)
if dist.get_rank() == 0:
# 保存元信息(用于异构拓扑恢复)
metadata = {
"world_size": dist.get_world_size(),
"step": step,
"model_config": model.config,
"timestamp": time.time(),
}
with open(f"{path}/metadata.json", "w") as f:
json.dump(metadata, f)
def load_checkpoint_distributed(model, optimizer, path):
"""支持从不同 world_size 的 checkpoint 恢复"""
state_dict = {
"model": model.state_dict(),
"optimizer": {}, # 占位
"step": 0,
}
dcp.load(
state_dict=state_dict,
storage_reader=FileSystemReader(path=path),
)
model.load_state_dict(state_dict["model"])
# 恢复优化器状态(需要 FSDP 特殊处理)
flattened_osd = FSDP.optim_state_dict_to_load(
model, optimizer, state_dict["optimizer"]
)
optimizer.load_state_dict(flattened_osd)
# 恢复随机数状态(保证可复现性)
torch.set_rng_state(state_dict["rng_state"])
torch.cuda.set_rng_state(state_dict["cuda_rng_state"])
return state_dict["step"]异步 Checkpoint:内存复制 + 后台保存
import torch
import threading
import copy
class AsyncCheckpointManager:
def __init__(self, model, optimizer, save_interval=1000):
self.model = model
self.optimizer = optimizer
self.save_interval = save_interval
self.writer_thread = None
self.pending_save = False
def maybe_save(self, step):
if step % self.save_interval != 0:
return
if self.pending_save:
# 上一次保存还没完成,跳过这次
logging.warning(f"Skipping checkpoint at step {step}: previous save in progress")
return
# 第一步:将状态复制到 CPU(短暂阻塞)
cpu_state = {
"model": copy.deepcopy(self.model.state_dict()), # 可能太慢
# 更好的方式:使用 buffer 共享
"step": step,
}
# 第二步:后台线程异步写入
self.pending_save = True
self.writer_thread = threading.Thread(
target=self._write_to_storage,
args=(cpu_state, f"/checkpoints/step_{step}"),
daemon=True,
)
self.writer_thread.start()
def _write_to_storage(self, state, path):
try:
torch.save(state, path)
logging.info(f"Checkpoint saved to {path}")
except Exception as e:
logging.error(f"Checkpoint save failed: {e}")
finally:
self.pending_save = False最优 checkpoint 频率取决于: \[T^* = \sqrt{\frac{2 \times T_{save} \times T_{total}}{MTTR}}\]
- \(T_{save}\):每次 checkpoint 的保存时间
- \(T_{total}\):总训练时间
- \(MTTR\):平均故障间隔时间
经验法则:checkpoint 间隔 ≈ MTTR 的 1/3。如果平均每 3 小时故障一次,就每小时保存一次。这样每次故障平均丢失 30 分钟的计算。
增量 Checkpoint
对于超大模型(千亿参数以上),全量 checkpoint 的 I/O 开销可能超过 1 分钟。增量 checkpoint 只保存发生变化的部分:
class IncrementalCheckpoint:
def __init__(self, base_path):
self.base_path = base_path
self.previous_state = None
self.version = 0
def save(self, model, optimizer, step):
current_state = {
"model": model.state_dict(),
"optimizer": optimizer.state_dict(),
"step": step,
}
if self.previous_state is None:
# 第一次:全量保存
torch.save(current_state, f"{self.base_path}/ckpt_v{self.version}.pt")
else:
# 增量:只保存变化的部分
delta = self._compute_delta(current_state, self.previous_state)
torch.save(delta, f"{self.base_path}/ckpt_v{self.version}_delta.pt")
self.previous_state = current_state
self.version += 1
def _compute_delta(self, current, previous):
"""计算两个 checkpoint 的差异"""
delta = {}
for key in current["model"]:
diff = (current["model"][key] - previous["model"][key]).abs()
if diff.max() > 1e-8: # 有变化才保存
delta[key] = current["model"][key]
return delta7.3 故障自动检测与恢复
自动恢复架构
%%{init: {'theme': 'base'}}%%
graph TB
subgraph "自动恢复系统"
A["Watchdog 进程"] -->|监控| B["训练进程健康状态"]
A -->|监控| C["GPU 状态"]
A -->|监控| D["网络连通性"]
B --> E{异常?}
C --> E
D --> E
E -->|否| F["✅ 继续运行"]
E -->|是| G["诊断故障类型"]
G --> H{可恢复?}
H -->|是| I["从 Checkpoint 重启"]
H -->|否| J["告警 + 隔离故障节点"]
I --> K["降级运行或等待替换"]
J --> L["人工介入"]
end
Watchdog 实现
import torch.distributed as dist
import time
import logging
import signal
import sys
class TrainingWatchdog:
def __init__(self, timeout=300, check_interval=30):
self.timeout = timeout # 心跳超时(秒)
self.check_interval = check_interval
self.last_heartbeat = time.time()
self.running = True
def heartbeat(self):
"""训练进程定期调用"""
self.last_heartbeat = time.time()
def monitor(self):
"""监控线程:检测训练进程是否卡住"""
while self.running:
time.sleep(self.check_interval)
elapsed = time.time() - self.last_heartbeat
if elapsed > self.timeout:
logging.error(
f"Training appears hung! "
f"No heartbeat for {elapsed:.0f}s"
)
# 收集诊断信息
self._collect_diagnostics()
# 尝试优雅终止然后重启
self._trigger_restart()
def _collect_diagnostics(self):
"""收集故障诊断信息"""
# GPU 状态
import subprocess
result = subprocess.run(
["nvidia-smi", "--query-gpu=utilization.gpu,memory.used,temperature.gpu",
"--format=csv"],
capture_output=True, text=True
)
logging.error(f"GPU status at failure:\n{result.stdout}")
# NCCL 状态
result = subprocess.run(
["nvidia-smi", "-q", "-d", "UTILIZATION,MEMORY,ECC"],
capture_output=True, text=True
)
logging.error(f"NCCL/GPU detail:\n{result.stdout[:5000]}")
# 堆栈跟踪(py-spy)
result = subprocess.run(
["py-spy", "dump", "--pid", str(os.getpid())],
capture_output=True, text=True
)
logging.error(f"Python stack:\n{result.stdout}")
def _trigger_restart(self):
"""触发任务重启"""
# 发送信号给主进程
os.kill(os.getpid(), signal.SIGTERM)
# K8s/Volcano 会自动重启 Pod
# 使用方式
watchdog = TrainingWatchdog(timeout=300)
import threading
monitor_thread = threading.Thread(target=watchdog.monitor, daemon=True)
monitor_thread.start()
for step, batch in enumerate(loader):
# 训练代码...
watchdog.heartbeat() # 定期心跳NCCL 超时与自动处理
# 设置 NCCL 超时和中断处理
import os
# NCCL 故障检测配置
os.environ["NCCL_TIMEOUT"] = "600" # 10 分钟超时
os.environ["NCCL_ASYNC_ERROR_HANDLING"] = "1" # 异步错误检测
os.environ["NCCL_DEBUG_SUBSYS"] = "ALL" # 详细日志
# 初始化时设置超时
dist.init_process_group(
backend="nccl",
timeout=torch.distributed.DefaultTimedelta(timedelta(minutes=10)),
)
try:
# 训练循环
for batch in loader:
loss = model(batch)
loss.backward()
optimizer.step()
except torch.distributed.DistBackendError as e:
logging.error(f"NCCL error: {e}")
# 检查哪些 rank 出了问题
healthy = check_all_ranks_health()
if healthy.count(False) <= max_failed_ranks:
# 少数 rank 故障,尝试弹性恢复
restart_from_checkpoint()
else:
# 大规模故障,需要人工介入
alert_oncall_engineer(f"Catastrophic failure: {healthy}")7.4 掉队者(Straggler)识别与缓解
在一个有数百张 GPU 的训练任务中,只要有一张卡变慢,整个训练的吞吐量就会被拖到那张卡的速度。这就是掉队者问题。
掉队者的原因
正常迭代时间:2.0 秒
掉队者迭代时间:4.5 秒(GPU 7 碰到了 ECC 错误重试)
集群吞吐量下降:50%
这相当于每天浪费 11.5 小时的 GPU 时间 × N 张卡。
%%{init: {'theme': 'base'}}%%
graph TB
subgraph "AllReduce 同步等待"
G0["GPU 0: 2.0s"] --> WAIT["AllReduce Barrier"]
G1["GPU 1: 2.0s"] --> WAIT
G2["GPU 2: 2.1s"] --> WAIT
G3["GPU 3: 4.5s ⚠️"] --> WAIT
G4["GPU 4: 2.0s"] --> WAIT
G5["GPU 5: 2.0s"] --> WAIT
G6["GPU 6: 2.0s"] --> WAIT
G7["GPU 7: 2.0s"] --> WAIT
WAIT --> RESULT["实际迭代时间:4.5s<br/>7 张卡等待 2.5s = 浪费"]
end
检测方法
class StragglerDetector:
def __init__(self, threshold_ratio=1.5, window_size=50):
self.threshold_ratio = threshold_ratio
self.window_size = window_size
self.step_times = collections.defaultdict(list) # rank -> [times]
def record_step(self, rank, duration):
self.step_times[rank].append(duration)
# 保持滑动窗口
if len(self.step_times[rank]) > self.window_size:
self.step_times[rank].pop(0)
def detect_stragglers(self):
"""返回掉队者 rank 列表"""
if not self.step_times:
return []
# 计算每个 rank 的平均迭代时间
avg_times = {
rank: sum(times) / len(times)
for rank, times in self.step_times.items()
}
# 计算中位数作为基准
median_time = sorted(avg_times.values())[len(avg_times) // 2]
# 找出偏离中位数过大的 rank
stragglers = [
rank for rank, avg in avg_times.items()
if avg > median_time * self.threshold_ratio
]
return stragglers, avg_times
def report(self):
stragglers, avg_times = self.detect_stragglers()
if stragglers:
logging.warning(f"Stragglers detected: {stragglers}")
for rank in stragglers:
ratio = avg_times[rank] / sorted(avg_times.values())[len(avg_times) // 2]
logging.warning(
f" Rank {rank}: {avg_times[rank]:.2f}s "
f"({ratio:.1f}x median)"
)
# 在训练循环中使用
detector = StragglerDetector(threshold_ratio=1.3)
for step, batch in enumerate(loader):
start = time.time()
# 给每个 rank 记录迭代时间
rank = dist.get_rank()
loss = model(batch)
loss.backward()
# 同步点
dist.barrier()
duration = time.time() - start
detector.record_step(rank, duration)
if step % 50 == 0:
# AllGather 所有 rank 的迭代时间
detector.report()缓解策略
| 策略 | 做法 | 适用场景 |
|---|---|---|
| 节点隔离 | 将掉队节点从集群中移除 | 硬件降级(ECC 错误累积) |
| 弹性缩容 | 去掉掉队者,用更少的 GPU 继续 | 可以容忍 world_size 变化 |
| NCCL_BUFFSIZE 调优 | 增大通信缓冲区 | 网络抖动型掉队 |
| 动态 batch 分配 | 给慢卡分配更小的 batch | 计算能力不均 |
| 通信异步化 | 用异步 AllReduce 替代同步 | 可容忍少量梯度滞后 |
nvidia-smi -q检查 GPU 频率是否被热限制(throttling)ibstat检查 IB 端口速率是否降级(从 200G 降到 100G)dmesg | grep -i ecc检查 ECC 错误日志nvidia-smi topo -m检查 GPU 拓扑是否变化- 检查其他进程是否占用了 GPU(
fuser -v /dev/nvidia*)
7.5 数据完整性与可复现性
可复现性的层次
| 层次 | 需要保证的内容 | 难度 |
|---|---|---|
| 完全位级复现 | 同样的输入 → 逐位相同的输出 | 极难(非确定性算法) |
| 确定性复现 | 同样的种子 → 同样的 loss 曲线 | 中等(需要固定所有随机性) |
| 统计复现 | 同样的配置 → 相近的最终精度 | 可行且实用 |
固定随机性
import torch
import numpy as np
import random
def set_seed(seed=42):
"""固定所有随机性源"""
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
# 确定性算法(可能降低性能)
torch.backends.cudnn.deterministic = True
torch.backends.cudnn.benchmark = False
# PyTorch 2.x 的确定性模式
torch.use_deterministic_algorithms(True)
# 环境变量
os.environ["CUBLAS_WORKSPACE_CONFIG"] = ":4096:8" # CUDA 确定性
os.environ["NCCL_DETERMINISTIC"] = "1" # NCCL 确定性数据验证
class DataIntegrityChecker:
def __init__(self, expected_hash_file):
self.expected_hashes = self._load_hashes(expected_hash_file)
def verify_sample(self, idx, sample):
"""验证单个样本的完整性"""
# 检查内容是否为空
if sample is None:
raise ValueError(f"Sample {idx} is None")
# 检查 NaN/Inf
if isinstance(sample, torch.Tensor):
if torch.isnan(sample).any():
raise ValueError(f"Sample {idx} contains NaN")
if torch.isinf(sample).any():
raise ValueError(f"Sample {idx} contains Inf")
# 检查数值范围
if isinstance(sample, torch.Tensor) and sample.is_floating_point():
if sample.abs().max() > 1e4:
logging.warning(f"Sample {idx} has extreme values: max={sample.abs().max()}")
def verify_batch(self, batch):
"""在训练循环中验证数据"""
for key, value in batch.items():
if isinstance(value, torch.Tensor):
if torch.isnan(value).any():
raise RuntimeError(f"NaN detected in batch field '{key}'!")
return True # 数据正常梯度健康检查
class GradientHealthMonitor:
def __init__(self, model, check_interval=10):
self.model = model
self.check_interval = check_interval
self.history = []
def check(self, step):
if step % self.check_interval != 0:
return
stats = {}
for name, param in self.model.named_parameters():
if param.grad is None:
continue
grad = param.grad
stats[name] = {
'mean': grad.mean().item(),
'std': grad.std().item(),
'max': grad.abs().max().item(),
'num_zeros': (grad == 0).sum().item(),
'num_nan': torch.isnan(grad).sum().item(),
'num_inf': torch.isinf(grad).sum().item(),
}
# 检查 NaN/Inf
if torch.isnan(grad).any():
raise RuntimeError(
f"NaN gradient in {name} at step {step}! "
f"Consider reducing learning rate or checking data."
)
# 检查梯度消失/爆炸
total_norm = torch.norm(
torch.stack([p.grad.norm() for p in self.model.parameters() if p.grad is not None])
).item()
if total_norm < 1e-7:
logging.warning(f"Gradient vanishing at step {step}! Total norm: {total_norm}")
elif total_norm > 1000:
logging.warning(f"Gradient exploding at step {step}! Total norm: {total_norm}")
self.history.append({'step': step, 'total_norm': total_norm})- 记录所有配置:代码版本、参数、环境(
pip freeze)、硬件型号 - 保存 RNG 状态:在每个 checkpoint 中保存随机数生成器的状态
- 版本化数据:用 DVC 或类似的工具追踪数据集版本
- 确定性数据加载:固定 DataLoader 的 worker seed
- 记录完整的训练日志:每一步的 loss、梯度范数、学习率
完全位级复现在分布式训练中几乎不可能(NCCL 的 AllReduce 顺序不确定)。追求统计复现——同配置、同数据、相近结果——是务实的目标。
小结
训练可靠性的核心是预期故障、拥抱故障、快速恢复:
| 关注点 | 关键措施 |
|---|---|
| 故障预防 | GPU 健康检查、网络监控、SDC 检测 |
| Checkpoint | 分布式保存、异步 I/O、增量更新 |
| 自动恢复 | Watchdog、NCCL 超时处理、弹性重启 |
| 掉队者 | 迭代时间监控、根因排查、节点隔离 |
| 数据完整性 | NaN 检测、梯度监控、版本化数据 |
训练可靠性不是某一个组件的工作,而是整个系统的属性——从数据管道到网络拓扑,从框架配置到运维流程,每个环节都需要考虑故障场景。
延伸阅读
- Meta Llama 3 训练报告:Meta, “The Llama 3 Herd of Models” (2024) — 真实的大规模训练故障统计和恢复实践
- Google SDC 研究:Ho et al., “Silent Data Corruptions at Scale” (2024)
- PyTorch Distributed Checkpoint:PyTorch 官方文档
- Megatron-LM Checkpoint:NVIDIA Megatron-LM 的 recovery 机制
- MLSys 可靠性:“Managing the Training of Large Language Models on 10,000+ GPUs” (Microsoft, 2024)