训练可靠性

训练可靠性

训练一个万亿参数模型需要数周甚至数月。在第 21 天崩溃不是一个小问题——它是数百万美元的损失。

章节导入:为什么可靠性是大模型训练的核心挑战

2024 年的一项调研显示,超过 70% 的大规模训练任务至少经历过一次非计划中断。Meta 在训练 Llama 3 时报告了 419 次意外中断,平均每 3 小时一次。

训练可靠性不是”锦上添花”的运维问题。它是系统设计的核心约束,决定了你的模型能不能在预算和时间表内完成训练。这一章我们将深入故障模式、容错机制和恢复策略。

7.1 常见故障模式与根因分析

大规模训练的故障来自多个层面。理解每种故障的症状和根因,是快速恢复的前提。

故障分类金字塔

%%{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

各类故障详解

故障类型 典型症状 发生频率 恢复时间
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}")
WarningSDC 的发生率

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
TipCheckpoint 频率公式

最优 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 delta

7.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 替代同步 可容忍少量梯度滞后
Tip掉队者根因排查清单
  1. nvidia-smi -q 检查 GPU 频率是否被热限制(throttling)
  2. ibstat 检查 IB 端口速率是否降级(从 200G 降到 100G)
  3. dmesg | grep -i ecc 检查 ECC 错误日志
  4. nvidia-smi topo -m 检查 GPU 拓扑是否变化
  5. 检查其他进程是否占用了 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})
Tip可复现性的实用建议
  1. 记录所有配置:代码版本、参数、环境(pip freeze)、硬件型号
  2. 保存 RNG 状态:在每个 checkpoint 中保存随机数生成器的状态
  3. 版本化数据:用 DVC 或类似的工具追踪数据集版本
  4. 确定性数据加载:固定 DataLoader 的 worker seed
  5. 记录完整的训练日志:每一步的 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)