Skip to content

训练稳定性与 Loss Spike 处理

大模型训练中,Loss Spike(损失突增)和训练不稳定是常见问题。本文分析训练不稳定的原因和应对策略。

训练稳定性与Loss Spike处理

常见不稳定模式

1. Loss Spike

Loss 突然大幅增长后缓慢恢复:

Step 1000: Loss = 2.1
Step 1001: Loss = 15.8  ← Spike!
Step 1002: Loss = 8.3
Step 1010: Loss = 3.5   ← 缓慢恢复
Step 1050: Loss = 2.2   ← 基本恢复

2. 梯度爆炸

梯度范数突然增大:

python
# 监控梯度范数
total_norm = 0
for p in model.parameters():
    if p.grad is not None:
        total_norm += p.grad.data.norm(2).item() ** 2
total_norm = total_norm ** 0.5
print(f"Gradient norm: {total_norm}")

# 异常信号:梯度范数 > 100× 正常值

3. Loss 不收敛

Loss 振荡或平台化。

Spike 原因分析

原因频率严重程度检测方法
坏数据样本低-中数据审计
学习率过大LR 消融
梯度累积溢出监控梯度范数
注意力分数溢出检查 NaN
优化器状态损坏检查点恢复

应对策略

梯度裁剪

python
# 常用梯度裁剪
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)

# DeepSpeed 配置
ds_config = {
    "gradient_clipping": 1.0,
}

学习率调度

python
# Cosine 调度 + 预热
from torch.optim.lr_scheduler import CosineAnnealingLR

scheduler = CosineAnnealingLR(
    optimizer,
    T_max=max_steps,
    eta_min=max_lr * 0.1,  # 最小学习率为最大的 10%
)

# Warmup
warmup_steps = 2000
for step in range(warmup_steps):
    lr = max_lr * step / warmup_steps
    for pg in optimizer.param_groups:
        pg["lr"] = lr

NaN 检测与恢复

python
def check_nan(model):
    """检测模型中的 NaN 值"""
    for name, param in model.named_parameters():
        if torch.isnan(param).any() or torch.isinf(param).any():
            print(f"NaN/Inf detected in {name}")
            return True
    return False

# 训练循环中的 NaN 检测
loss = model(inputs)
if torch.isnan(loss) or torch.isinf(loss):
    print("NaN loss detected, skipping step")
    optimizer.zero_grad()
    continue  # 跳过当前步

自动回滚

python
# 检测 Spike 并自动回滚
class SpikeDetector:
    def __init__(self, window=100, threshold=3.0):
        self.losses = []
        self.window = window
        self.threshold = threshold

    def is_spike(self, current_loss):
        self.losses.append(current_loss)
        if len(self.losses) < self.window:
            return False

        recent_mean = np.mean(self.losses[-self.window:])
        return current_loss > recent_mean * self.threshold

预防措施

  1. 使用 BF16 而非 FP16(无需 Loss Scaling)
  2. 梯度裁剪 max_norm=1.0
  3. 预热学习率(2000+ 步)
  4. 数据质量过滤(去除异常长度/字符样本)
  5. 定期保存检查点(每 1000 步)

Spike 后恢复

Spike 后不要立即继续训练。检查:1) 数据是否异常;2) 梯度范数是否恢复;3) 模型参数是否含 NaN。如有 NaN,必须回滚检查点。

相关资源

最近更新