模型训练稳定性:从训练链路到排障手册

模型训练稳定性:从训练链路到排障手册

我写这篇手册,不是为了追求一条毫无波动的 loss 曲线,而是为了回答一个更实际的问题:怎样让模型在训练预算内持续、有效、可控地学习,并在出现异常时尽快找到真正原因?

导读:先把目标说清楚------什么才算训练稳定

0.1 loss 有波动,不等于训练不稳定

大模型训练通常使用小批量数据。不同 batch 的内容和难度不同,单步训练损失(training loss)自然会波动。因此,我不会因为某一步 loss 上升,就立刻判断训练不稳定;也不会因为 loss 曲线平滑,就断定训练一定正确。

我把稳定训练定义为:

在给定的数据、模型、计算精度和训练预算下,所有关键计算保持可执行,梯度信号能够正常传播,参数获得受控且有效的更新,训练目标呈总体改善趋势,并最终产生符合预期的验证结果。

这个定义包含五层要求:

  1. 数值可计算 :loss、激活值、梯度、参数和优化器状态不能出现 NaNinf
  2. 梯度信号健康:误差信号能够传到应该学习的参数,既没有系统性消失,也没有持续爆炸。
  3. 参数更新受控:单步更新不能大到破坏已有模型状态。
  4. 参数更新有效:应该学习的参数确实发生了足够的变化,而不是看似运行、实际没有学习。
  5. 训练方向正确:数据、标签、mask、loss 和分布式归约都符合训练目标,验证结果能够改善。

其中,前四项保证训练链路能够工作,第五项决定模型是不是在学正确的东西。

0.2 目标、指标、现象、原因和处理手段不能混在一起

训练排障最容易出现的错误,是把观察到的现象直接当成原因。例如:"loss 很震荡,所以学习率一定太大",或者"grad norm 很高,所以一定发生了梯度爆炸"。这两种推断都缺少证据。

我会严格区分下面五类信息:

类型 例子 我用它回答什么问题
训练目标 validation loss、perplexity、任务指标 模型最终是否学到了有用能力
过程指标 grad norm、相对更新量、AMP scale 训练链路内部正在发生什么
异常现象 loss spike、长期不降、出现 NaN 哪个时间点可能出现了问题
候选原因 异常 batch、学习率过大、mask 错误 哪条链路值得优先验证
处理手段 梯度裁剪、降低学习率、修复数据 根因确认后应该怎样处理

**一个指标只能提供证据,不能自动给出根因。**我会先寻找第一个偏离健康状态的信号,再用重放或单变量对照实验确认原因。

0.3 全文主线:一次训练更新怎样发生

一轮参数更新可以压缩成下面这条链路:

数据与训练目标 → 前向计算 → loss → 反向传播 → gradient → optimizer → parameter update → 新的模型状态 → 后续训练与验证结果

后文的四道保障都附着在这条链路上:

  • 数值保障覆盖前向、反向和优化器计算;
  • 梯度保障覆盖误差信号从 loss 返回各层参数的过程;
  • 更新保障覆盖优化器把梯度转换为参数变化的过程;
  • 正确性保障覆盖数据、目标函数、mask、归约和验证链路。

Part 1 一次训练更新究竟发生了什么

1.1 从一个 batch 到 loss

一个 batch 是一次前向和反向计算使用的数据集合。在大语言模型训练中,batch 通常包含 token 序列以及与训练目标有关的辅助信息:

  • token:分词器把文本转换得到的离散编号;
  • label:模型需要预测的目标 token;
  • attention mask:规定每个位置可以关注哪些位置;
  • loss mask:规定哪些位置参与 loss;
  • 有效 token 数:真正参与 loss 计算的位置数量。

模型通过前向传播产生 logits。logits 是 softmax 之前的未归一化分数。以自回归语言模型为例,模型用当前位置之前的上下文预测下一个 token,交叉熵损失衡量预测分布和正确 token 之间的差异。

如果第 \(i\) 个有效 token 的损失为 \(\ell_i\),一个 batch 的 token 平均损失通常写为:

\L=\\frac{\\sum_{i=1}\^{N}\\ell_i}{N} \\

这里的 \(N\) 必须是有效 token 数,而不是张量总长度。如果 padding、被屏蔽位置或只作为上下文的位置错误地进入分母,loss 的数值和梯度尺度都会改变。

1.2 参数、激活值和梯度分别是什么

这三个概念都以张量形式存在,但含义完全不同。

对象 英文 产生时间 作用
参数 parameter / weight 初始化或加载 checkpoint 时产生 保存模型已经学到的状态
激活值 activation 每次前向传播时产生 把输入逐层转换为 logits
梯度 gradient 反向传播时产生 描述 loss 对参数的局部变化率

例如:

python 复制代码
h = self.proj(x)
  • self.proj.weight 是参数;
  • h 是这一层的激活值;
  • 反向传播后的 self.proj.weight.grad 是梯度。

参数异常会影响后续所有 batch;激活异常首先发生在当前前向传播;梯度异常首先影响当前参数更新。排障时,我会沿着时间顺序寻找第一个异常对象。

1.3 梯度表示什么

设模型参数为 \(\theta\),当前 loss 为 \(L(\theta)\),梯度定义为:

\g_t=\\nabla_\\theta L(\\theta_t) \\

梯度告诉我:在当前参数附近,如果轻微改变某个参数,loss 会怎样变化。它提供局部方向和敏感程度,但它不是最终参数更新量

最基本的梯度下降会直接使用:

\\\theta_{t+1}=\\theta_t-\\eta g_t \\

其中,\(\eta\) 是学习率(learning rate)。但大模型训练通常使用 AdamW。AdamW 会用历史梯度计算一阶矩和二阶矩,再结合学习率、数值稳定项和权重衰减产生实际更新。

1.4 参数更新量表示什么

我把第 \(t\) 步的参数更新量定义为:

\\\Delta\\theta_t=\\theta_{t+1}-\\theta_t \\

因此:

梯度 \(g_t\) 是优化器的输入,参数更新量 \(\Delta\theta_t\) 是优化器处理后的结果。

即使梯度正常,参数也可能没有有效更新。例如:

  • 当前学习率已经降为零;
  • 参数被冻结;
  • 参数没有加入 optimizer;
  • FP16 检测到非有限梯度并跳过了 optimizer step;
  • AdamW 的二阶矩把当前梯度大幅缩小;
  • scheduler 在错误的时间点推进。

这就是为什么我会分别监控 grad norm 和相对更新量,而不是把二者视为同一个指标。


Part 2 第一道保障:所有数值必须可计算

2.1 什么是数值稳定

计算机不能表示任意大的实数,也不能无限精确地表示任意小数。训练中的张量必须落在当前浮点格式能够表示的范围和精度内。

我重点关注三类异常:

  • 上溢(overflow) :数值太大,超出表示范围,可能变成 inf
  • 下溢(underflow):数值太小,无法保留,可能被舍入为零;
  • 非法运算 :例如 0/0、对负数开平方,结果可能变成 NaN

NaN 是 Not a Number,表示结果不是合法数值;inf 是 infinity,表示正无穷或负无穷。一旦非有限值进入后续矩阵乘法、归一化或优化器状态,它通常会迅速污染更多参数。

2.2 浮点数的范围和精度来自哪里

浮点数可以直观理解为科学计数法:

\\\text{value}=(-1)\^{\\text{sign}}\\times\\text{significand}\\times 2\^{\\text{exponent}} \\

  • 符号位决定正负;
  • 指数位主要决定能表示多大、多小的数;
  • 尾数位主要决定数值精度。

常见训练格式如下:

格式 指数位 尾数位 核心特点 训练中的主要风险
FP32 8 23 范围和精度都较高 显存和计算开销较大
FP16 5 10 精度高于 BF16,但范围较小 大值上溢、小梯度下溢
BF16 8 7 范围接近 FP32,精度低于 FP16 小变化可能被舍入,非法运算仍会产生非有限值

这里最重要的区别是:**指数位控制动态范围,尾数位控制有效精度。**BF16 不容易因为动态范围不足而溢出,但它并不等于"数值绝对安全"。

2.3 什么是自动混合精度训练

自动混合精度(Automatic Mixed Precision,AMP)不是把所有计算统一改成 FP16 或 BF16,而是让框架根据算子特性选择精度:

  • 适合低精度的矩阵乘法使用 FP16 或 BF16;
  • 对精度敏感的归约或特定算子保留 FP32;
  • 优化器状态通常使用更高精度保存。

这样可以降低显存占用并提高吞吐,同时减少低精度计算带来的风险。

在新的 PyTorch 接口中,典型写法是:

python 复制代码
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
    logits = model(input_ids, attention_mask=attention_mask).logits
    loss = loss_fn(logits, labels)

我不会假设进入 autocast 后每个算子都采用同一种 dtype。排查数值问题时,需要定位实际发生异常的算子和张量。

2.4 FP16 为什么需要 loss scaling

损失缩放(loss scaling)主要解决 FP16 反向传播中的小梯度下溢。

设原始 loss 为 \(L\),缩放系数为 \(S\)。反向传播前先计算:

\L'=S\\cdot L \\

对应梯度也被放大:

\\\nabla_\\theta L'=S\\cdot\\nabla_\\theta L \\

这样,小梯度在 FP16 中更不容易被舍入为零。optimizer step 前再把梯度除以 \(S\),恢复真实尺度。

动态损失缩放会根据非有限梯度调整 \(S\):

  • 检测到 infNaN 时,跳过当前 optimizer step,并降低 scale;
  • 连续一段时间没有溢出时,逐步增大 scale。

需要注意:scale 下降说明框架检测到了不可用于更新的梯度,但它没有告诉我根因。根因可能是 FP16 范围不足,也可能是输入、loss 或模型计算已经出现异常。

2.5 FP16 下统计和裁剪梯度的正确顺序

scaler.scale(loss).backward() 得到的是放大后的梯度。我要先还原真实梯度,再统计 grad norm 和执行裁剪:

python 复制代码
optimizer.zero_grad(set_to_none=True)

with torch.autocast(device_type="cuda", dtype=torch.float16):
    logits = model(input_ids, attention_mask=attention_mask).logits
    loss = loss_fn(logits, labels)

scaler.scale(loss).backward()
scaler.unscale_(optimizer)

grad_norm = torch.nn.utils.clip_grad_norm_(
    model.parameters(),
    max_norm=clip_threshold,
)

scaler.step(optimizer)
scaler.update()

如果先裁剪缩放后的梯度,裁剪阈值就失去了原本含义。如果使用梯度累积,我会等完整有效 batch 的梯度全部累积完成,再执行一次 unscale_、统计、裁剪和 optimizer step。

2.6 常见数值危险点

危险点 典型错误 更可靠的处理
softmax 对过大 logits 直接手写 exp 使用框架提供的 softmax、cross entropy 或 logsumexp
对数 对零概率计算 log(0) 使用数值稳定实现,并明确概率下界的业务含义
除法 有效 token 数或归一化分母为零 在运算前显式检查分母
attention mask 一整行都被设为负无穷后执行 softmax 保证每个查询至少存在合法键,或显式处理空行
loss reduction 空 mask 后继续求平均 把有效数量为零作为数据或实现错误处理
自定义算子 低精度中执行大范围累加 对敏感计算指定 FP32 或使用稳定 kernel

我不会用 torch.clamp(loss) 作为通用止损手段。直接截断总 loss 会改变训练目标,还可能把真正的实现错误隐藏起来。

2.7 数值稳定的直接监控项

我只保留能直接触发行动的数值监控:

指标 统计位置 异常后的直接行动
loss 是否有限 backward 前 保存 batch 和前向现场,停止当前更新
gradient 是否有限 backward 后、optimizer step 前 跳过更新并定位第一个异常梯度
parameter 是否有限 optimizer step 后 停止训练并回滚到最后一个健康 checkpoint
optimizer step 是否执行 每个计划更新点 检查 AMP 跳步、参数冻结和控制流
AMP scale 仅 FP16 训练 连续下降时检查溢出位置和模型是否适合 FP16

activation 和 logits 的详细分层统计不作为常驻核心指标。只有当 loss 首先出现异常、而输入仍然有限时,我才临时开启分层检查,定位第一个产生非有限值或异常放大的模块。


Part 3 第二道保障:梯度信号必须健康

3.1 梯度为什么需要沿网络传播

大模型包含许多层。靠近输出的 loss 必须通过反向传播,把误差信号传回每一层参数。对于从 \(h_0\) 到 \(h_N\) 的深层网络,可以写成:

\\\frac{\\partial L}{\\partial h_0} = \\frac{\\partial L}{\\partial h_N} \\prod_{l=1}\^{N} \\frac{\\partial h_l}{\\partial h_{l-1}} \\

这里的每一项都是局部变换对输入的敏感程度。它们连续相乘后,某些方向的信号可能逐层缩小,也可能逐层放大。

梯度信号健康不是要求所有层的梯度一样大,而是要求应该学习的层在当前阶段持续收到可用信号。

3.2 梯度消失是什么

梯度消失(gradient vanishing)是指反向信号在传播过程中被系统性削弱,使部分层长期几乎得不到有效梯度。

我会同时检查三个证据:

  1. 训练仍处于应该明显学习的阶段,但 loss 长期停在较高位置;
  2. 靠近输入的若干层,其 grad norm 长期远小于靠近输出的层;
  3. 对应层的相对更新量长期接近测量下限,验证结果也没有改善。

仅仅看到训练中后期全局 grad norm 变小,不足以证明梯度消失。正常收敛时,模型接近当前可达到的低损失区域,梯度变小是合理现象。

3.3 梯度爆炸是什么

梯度爆炸(gradient explosion)是指反向信号在传播过程中被持续放大,最终使梯度尺度异常增大,并可能造成过大的参数更新。

我会区分三种现象:

现象 更可能说明什么 不能直接断言什么
某个 batch 的 grad norm 单步尖峰,随后恢复 当前 batch 更难或存在异常样本 网络结构必然发生梯度爆炸
深层到浅层的 grad norm 系统性放大 反向传播链路可能放大信号 根因一定是学习率过大
grad norm 持续抬升并伴随 loss 恶化 训练状态正在偏离健康区域 仅靠一次裁剪就能修复根因

学习率不会改变当前 backward 已经算出的原始梯度,但它会改变这一步参数更新和下一步模型所在的位置。因此,过大学习率可能先造成参数移动过远,再让后续步骤进入高梯度区域。

3.4 grad norm 是什么

梯度范数(gradient norm,grad norm)通常指所有目标参数梯度组成的整体 L2 范数:

\\\lVert g\\rVert_2=\\sqrt{\\sum_i g_i\^2} \\

它把大量梯度压缩成一个数,适合发现整体尖峰和长期漂移,但有三个边界:

  • 它不能指出是哪一层首先异常;
  • 它会受 loss reduction、有效 token 数和参数集合影响;
  • 它的绝对大小不能跨模型直接比较。

因此,我不会把 0.1--10 写成通用健康区间。我先用健康的 pilot run 建立同一模型、同一统计口径下的基线,再关注相对变化。

3.5 全局和分层 grad norm 怎样配合

常驻监控只记录裁剪前全局 grad norm,因为它成本低,而且能够及时发现整体异常。

只有当全局指标异常时,我才开启分层统计,并按 Transformer 的自然边界比较:

  • token embedding;
  • 前部 Transformer block;
  • 中部 Transformer block;
  • 后部 Transformer block;
  • 最终归一化层;
  • language modeling head。

分层统计的目的不是画出更多曲线,而是回答一个具体问题:哪一组参数最先偏离自己的健康基线?

3.6 梯度裁剪解决什么问题

全局 L2 梯度裁剪(gradient clipping)会在梯度范数超过阈值 \(c\) 时,按同一比例缩小所有梯度:

\g\\leftarrow g\\cdot\\min\\left(1,\\frac{c}{\\lVert g\\rVert_2+\\epsilon}\\right) \\

统一缩放保留整体梯度方向,同时限制当前更新风险。

我会同时记录:

  • 裁剪前 grad norm;
  • 当前 step 是否触发裁剪;
  • 实际缩放系数。

如果多数 step 都被大幅裁剪,配置中的 learning rate 已经不再代表实际更新尺度。此时我会检查裁剪阈值、loss 归一化和异常来源,而不是把长期裁剪当作正常状态。

3.7 Transformer 结构怎样帮助梯度传播

残差连接提供接近恒等映射的主路径,让梯度不必完全穿过每个非线性子层。归一化和初始化则控制各层激活及局部变换的尺度。

在许多 Transformer 配置中,Pre-LN 比原始 Post-LN 更容易稳定启动,因为残差主路径更直接。但这不是说 Post-LN 一定不能训练,也不是说使用 Pre-LN 后就不需要 warmup。

我只把以下结构因素纳入排障:

  • 残差分支输出是否随深度合理缩放;
  • LayerNorm 或 RMSNorm 的位置和 eps 是否正确;
  • attention 的 Q、K 尺度是否符合实现预期;
  • 初始化是否与既定模型结构一致。

这些是模型设计和实现检查项,不应该在发现一次 loss spike 后随意叠加修改。


Part 4 第三道保障:参数更新必须受控且有效

4.1 正常梯度为什么不保证正常更新

反向传播结束后,优化器才决定怎样改变参数。以 AdamW 为例,同一个当前梯度会因为历史一阶矩、二阶矩、学习率和权重衰减不同,产生不同的更新量。

因此我把更新健康拆成两个问题:

  • 受控:更新没有大到破坏当前模型状态;
  • 有效:更新没有小到让应该学习的参数长期不动。

训练可能非常平滑但没有效果。例如,learning rate 过小会让 loss 缓慢变化甚至几乎不变。训练也可能短期下降很快但不可控,例如大学习率在早期跨过多个高曲率区域,随后引发震荡或发散。

4.2 AdamW 怎样把梯度变成更新

AdamW 对每个参数维护两个历史状态。

一阶矩(first moment)是梯度的指数移动平均:

\m_t=\\beta_1m_{t-1}+(1-\\beta_1)g_t \\

二阶矩(second moment)是梯度平方的指数移动平均:

\v_t=\\beta_2v_{t-1}+(1-\\beta_2)g_t\^2 \\

偏差修正(bias correction)用于修正状态从零开始带来的系统偏小:

\\\hat m_t=\\frac{m_t}{1-\\beta_1\^t},\\qquad \\hat v_t=\\frac{v_t}{1-\\beta_2\^t} \\

AdamW 的一次更新可以写为:

\\\theta_{t+1} = (1-\\eta\\lambda)\\theta_t - \\eta\\frac{\\hat m_t}{\\sqrt{\\hat v_t}+\\epsilon} \\

其中:

  • \(\eta\) 是 learning rate;
  • \(\lambda\) 是 weight decay 系数;
  • \(\epsilon\) 防止分母过小;
  • \(\hat m_t\) 提供经过平滑的方向信息;
  • \(\hat v_t\) 根据历史梯度平方调整每个参数的更新尺度。

"一阶矩管方向、二阶矩管步长"适合建立直觉,但不是完整定义。一阶矩也包含幅度信息,二阶矩也会改变不同参数之间的相对更新尺度。

AdamW 的关键变化,是把权重衰减(weight decay)与自适应梯度更新分开,避免 L2 项再次被二阶矩缩放。

4.3 learning rate 为什么会导致震荡或学不动

先看一个局部二次函数:

\L(\\theta)=\\frac{1}{2}\\lambda\\theta\^2 \\

使用普通梯度下降时:

\\\theta_{t+1}=(1-\\eta\\lambda)\\theta_t \\

只有满足:

\\|1-\\eta\\lambda\|\<1 \\

参数才会逐步接近最优点。这个例子说明,同一个 learning rate 在平坦方向可能安全,在高曲率方向可能过大。

真实大模型和 AdamW 更复杂,但判断逻辑不变:我不会只看配置中的 learning rate,而会结合实际相对更新量和训练结果判断步长是否合适。

4.4 warmup 是什么

学习率预热(learning rate warmup)是在训练开始阶段,从较小 learning rate 逐步增加到峰值 learning rate 的过程。

warmup 的作用不是单纯"等待 AdamW 的二阶矩稳定"。训练初期同时存在以下风险:

  • 模型仍在初始化附近;
  • 局部曲率和激活尺度快速变化;
  • 优化器历史统计很短;
  • 前几个 batch 可能不能代表总体数据分布。

较小的初始步长可以降低模型一步进入高风险区域的概率。warmup 是否需要、需要多长,必须结合模型结构、初始化、有效 batch 和峰值 learning rate 验证。

我会区分三个单位:

  • micro-step:执行一次 forward 和 backward;
  • optimizer step:真正更新一次参数;
  • warmup token:warmup 期间实际看过的有效 token 数。

存在梯度累积时,scheduler 应该按 optimizer step 推进,而不是按每个 micro-step 推进。

4.5 update norm 是什么

更新范数(update norm)是参数更新量的 L2 范数:

\\\text{Update Norm}=\\lVert\\Delta\\theta_t\\rVert_2 \\

它回答:这一步参数一共改变了多少?

如果某层参数从 \(\theta_t\) 变成 \(\theta_{t+1}\),可以直接计算:

python 复制代码
update_norm = (new_param - old_param).float().norm(2)

update norm 是绝对量。参数规模和参数自身尺度不同,绝对量不能直接用于跨层比较。

4.6 weight norm 和相对更新量是什么

权重范数(weight norm)表示参数当前的整体尺度:

\\\text{Weight Norm}=\\lVert\\theta_t\\rVert_2 \\

我把某个参数组的相对更新量定义为:

\r= \\frac{\\lVert\\Delta\\theta_t\\rVert_2} {\\lVert\\theta_t\\rVert_2+\\epsilon} \\

它常被称为 update-to-weight ratio,可以直观理解为:这一步参数变化相对于参数自身尺度有多大。

我不会把 0.001--0.01 当作所有模型的健康区间。相对更新量会受到参数化方式、参数组、训练阶段、weight decay 和统计频率影响。更可靠的使用方式是:

  1. 为固定模型和固定统计口径建立健康 pilot run 基线;
  2. 比较同一参数组在不同时段的变化;
  3. 比较关键参数组之间是否出现持续失衡;
  4. 将异常与 loss、grad norm、learning rate 和验证结果对齐。

4.7 哪些参数组值得常驻监控

为了避免指标过载,我不记录每一层的相对更新量。对标准 decoder-only Transformer,我选择四个代表性参数组:

  1. token embedding;
  2. 前部 Transformer blocks;
  3. 后部 Transformer blocks;
  4. language modeling head。

如果输入 embedding 和输出 head 共享权重,只记录一次并明确标记 tied weight。bias 和归一化参数的 weight norm 可能很小,不与大矩阵使用同一报警规则。

常驻监控可以降低采样频率,例如每隔固定数量的 optimizer step 计算一次。真正出现异常时,再临时展开到逐层统计。

4.8 怎样判断更新受控且有效

我把判断压缩成下面这张表:

观察 更新状态 下一步
相对更新量突增,随后 loss 和 grad norm 恶化 更新可能过大 检查 learning rate、异常 batch、裁剪和 optimizer state
grad norm 正常,但相对更新量长期极小 更新可能不足 检查 learning rate、二阶矩、冻结参数和 optimizer 参数组
optimizer step 持续跳过 参数没有按计划更新 检查 FP16 溢出和 GradScaler 状态
各参数组均有稳定更新,但验证结果不改善 更新存在但方向未必正确 检查数据、目标函数和验证链路
更新稳定,训练与验证目标总体改善 更新受控且有效 继续训练并保持监控口径不变

Part 5 第四道保障:数据、目标函数和实现必须正确

5.1 loss 下降为什么仍然可能训练错误

优化器只负责降低我提供的目标函数。只要目标函数可以被优化,loss 就可能下降;它不会替我判断标签、mask 和数据划分是否符合真实目标。

下面这些错误都可能得到平滑下降的 training loss:

  • label 与输入序列错位;
  • 验证数据混入训练数据;
  • padding 或只作为上下文的位置进入 loss;
  • packing 后不同样本之间发生错误 attention;
  • 指令微调中把不应监督的 prompt token 计入 loss;
  • 验证代码使用了和训练目标不一致的预处理。

因此,训练稳定不只是"能够算下去",还包括"模型在学正确的东西"。

5.2 loss reduction 为什么会改变梯度尺度

假设第 \(k\) 个 micro-batch 有 \(n_k\) 个有效 token,第 \(i\) 个 token 的损失为 \(\ell_{k,i}\)。如果训练目标是所有有效 token 的平均损失,正确形式是:

\L= \\frac{\\sum_k\\sum_{i=1}\^{n_k}\\ell_{k,i}} {\\sum_k n_k} \\

下面这种"先对每个 micro-batch 求平均,再对 micro-batch 求平均"的做法通常不等价:

\L_{\\text{wrong}} = \\frac{1}{K} \\sum_k \\left( \\frac{1}{n_k} \\sum_{i=1}\^{n_k}\\ell_{k,i} \\right) \\

当不同 micro-batch 的有效 token 数差异很大时,错误聚合会改变样本权重和梯度尺度。它还会让不同 batch、序列长度和并行配置下的 grad norm 无法比较。

5.3 数据异常怎样传导成训练异常

数据问题通常沿下面的顺序传播:

异常输入或标签 → 异常 logits 或单样本 loss → grad norm 尖峰 → 参数更新异常 → 后续 loss 无法恢复

因此,发现单步 loss spike 时,我会对齐查看当前 batch 的四项上下文:

  1. 有效 token 数;
  2. 序列长度;
  3. 数据来源或数据桶;
  4. 是否存在非有限输入或越界 label。

这些不是泛化的"数据统计面板",而是为了回答:当前异常是否和 batch 构成同步发生?

5.4 分布式训练最容易改变哪些口径

大模型训练通常横跨多个设备。下面四项必须明确采用全局还是本地口径:

项目 正确问题 常见错误
loss 是否按全局有效 token 加权 只记录 rank 0 的局部均值
有效 token 数 是否对所有数据并行 rank 求和 每卡先平均后再平均
grad norm 是否覆盖完整参数集合 只统计当前 shard 却当作全局值
optimizer step 所有 rank 是否一致执行或跳过 某些 rank 跳步、其他 rank 更新

此外,我会确认 scheduler 按实际 optimizer step 推进。如果某一步因为 FP16 溢出而被跳过,scheduler 是否仍然推进必须符合训练框架的既定语义,并在日志中保持一致。

5.5 最小正确性验证

正式训练前,我至少执行下面四个验证:

  1. 随机抽查样本:直接查看 token、label、attention mask 和 loss mask 是否对齐。
  2. 固定 batch 重复训练:确认模型能够在同一个小 batch 上明显降低 loss。
  3. 小数据集过拟合:确认模型、loss 和 optimizer 的完整链路具有学习能力。
  4. 单卡与多卡对照:使用相同有效 token 数比较前几个 optimizer step 的 loss 和更新趋势。

固定 batch loss 能下降,只能证明当前实现具有优化能力;它不能证明数据分布、验证集和最终目标正确。每个实验都只回答一个问题。


Part 6 最小监控面板:训练时到底应该看什么

6.1 我的取舍原则

我不会因为某个指标"可能有用"就把它放进常驻面板。一个常驻指标必须同时满足三个条件:

  1. 能稳定采集,统计口径明确;
  2. 能区分至少一种重要训练状态;
  3. 异常后有明确的下一步行动。

按照这个原则,我把常驻监控压缩成九项。

6.2 九项常驻核心监控

编号 指标 记录频率 我怎样使用
1 单步 training loss 每个 optimizer step 对齐具体异常 step 和 batch
2 training loss 滑动平均 每个 optimizer step 判断总体改善、平台期和持续漂移
3 validation loss / perplexity 固定验证间隔 判断模型是否真正改善
4 learning rate 每个 optimizer step 检查 warmup、decay 和 scheduler
5 裁剪前全局 grad norm 每个 optimizer step 发现梯度尖峰和持续偏移
6 clip coefficient 或裁剪触发标记 每个 optimizer step 判断裁剪是否长期改变优化过程
7 四组参数的相对更新量 固定采样间隔 判断更新过大、过小和组间失衡
8 非有限值与 optimizer step 执行状态 每个计划更新点 发现数值故障和实际跳步
9 每个 optimizer step 的有效 token 数 每个 optimizer step 保证 loss 和梯度尺度可比较

预训练通常使用 validation loss 或 perplexity。微调任务还应记录一个与任务目标直接对应的验证指标,例如分类准确率或生成任务的既定评估分数;具体指标由任务定义决定,不能由稳定性手册代替业务选择。

6.3 FP16 的两个条件监控

只有使用 FP16 和动态 GradScaler 时,我才额外记录:

  • 当前 AMP scale;
  • skipped step 的连续次数和累计次数。

单次 skipped step 不一定意味着训练失败。连续跳步说明模型没有按计划更新,我会立即检查 loss、梯度是否有限,以及 FP16 动态范围是否适合当前模型。

BF16 或 FP32 训练没有使用 GradScaler 时,不保留这两项,避免在面板中放置没有语义的空指标。

6.4 故障发生后再开启的诊断指标

常驻面板发现异常后,我才临时开启以下统计:

诊断指标 开启条件 要回答的问题
分层 activation 最大绝对值和有限值检查 loss 在 forward 后已经异常 哪一层首先产生异常激活
分层 grad norm 全局 grad norm 偏离基线 哪一组参数首先出现梯度异常
logits 最大绝对值 loss 异常但前部 activation 正常 输出分布是否异常集中或溢出
AdamW 二阶矩摘要 grad norm 正常但更新量异常 优化器状态是否大幅改变有效步长
异常 batch 四项上下文 单步 loss 或 grad norm spike 异常是否由 batch 构成触发

这些指标的价值在于缩小范围。根因确认后,我会关闭临时统计,避免长期增加存储、通信和分析负担。

6.5 怎样建立健康基线

绝对数值只有在统计口径固定时才有意义。我会在 pilot run 中固定:

  • 模型和参数分组;
  • loss reduction;
  • 有效 batch 的定义;
  • 梯度累积步数;
  • 分布式归约方式;
  • 统计发生在裁剪前还是裁剪后。

然后保留单步原始值和滑动平均。单步值用于定位具体 batch,滑动平均用于判断持续趋势。

我重点区分三种异常形态:

  • 单步尖峰:一步偏离后迅速恢复;
  • 持续漂移:多个窗口逐渐离开基线;
  • 结构性失衡:某个参数组长期与其他组行为不同。

报警规则应围绕这三种形态建立,而不是从其他模型复制一个固定阈值。


Part 7 从异常现象到根因:不要看见 loss 就猜 learning rate

7.1 第一原则:寻找第一个异常信号

如果 loss 在第 \(t\) 步变成 NaN,原因可能在更早的位置发生:输入已经异常、某层 activation 首先溢出、loss 计算包含非法操作、backward 产生非有限梯度,或者上一步参数已经被污染。

我的排查顺序始终沿计算链路前进:

  1. 输入和 label;
  2. activation;
  3. logits;
  4. 各项 loss;
  5. gradient;
  6. optimizer state;
  7. parameter;
  8. 后续验证结果。

找到第一个异常位置,才有可能区分根因和后果。

7.2 异常发生时先保存什么

我会在继续重跑前保存最小现场:

  • global step、optimizer step 和 checkpoint 标识;
  • 数据分片、样本索引和随机种子;
  • 单步 loss、各 loss 分项和有效 token 数;
  • learning rate、裁剪前 grad norm 和 clip coefficient;
  • optimizer step 是否执行;
  • FP16 下的 AMP scale 与 skipped step;
  • 各 rank 是否出现非有限值。

如果合规要求不允许保存原始文本,我保存能够重放的样本索引和脱敏统计。

7.3 第一个 step 就出现 NaN

第一个 step 尚未经历历史更新,因此我优先检查:

  1. 输入、label 和 mask 是否合法;
  2. 初始化后的参数是否有限;
  3. forward 中哪一层首先产生非有限 activation;
  4. loss 是否包含除零、空平均或非法索引;
  5. 当前精度从 BF16/FP16 改为 FP32 后是否仍能复现。

如果 FP32 仍然复现,问题通常不只是低精度动态范围。如果只有 FP16 复现,我再检查模型数值范围和 loss scaling。

7.4 loss 突然出现 spike

我先判断 spike 之后是否恢复。

单步 spike 后恢复

优先对齐当前 batch 的有效 token 数、序列长度、数据来源和 label 合法性。如果相同 batch 可以稳定重放出 spike,继续检查数据与 loss;如果无法重放,检查随机算子、并行执行和前一状态差异。

连续多个 spike

对齐 grad norm、clip coefficient、相对更新量和 learning rate。如果这些指标持续偏离,训练状态可能已经离开健康区域。

spike 后永久不恢复

我会停止让异常状态进入新的正式 checkpoint,从最后一个健康 checkpoint 重放。如果健康 checkpoint 加相同 batch 会复现,优先检查数据;如果更换 batch 仍然复现,优先检查模型和 optimizer state。

7.5 loss 长期不下降

我按下面顺序排查:

  1. 固定 batch 能否过拟合;
  2. learning rate 是否为预期值;
  3. optimizer step 是否实际执行;
  4. 关键参数组是否加入 optimizer;
  5. 裁剪前 grad norm 是否存在有效信号;
  6. 四组相对更新量是否持续接近测量下限;
  7. label、loss mask 和 loss reduction 是否正确。

如果梯度存在但更新量极小,我检查 optimizer 和 learning rate;如果梯度本身在部分层长期近零,我检查梯度传播;如果参数持续更新但验证结果不改善,我检查数据和目标。

7.6 loss 平滑,但验证结果不改善

这种情况不属于典型数值失稳。我会检查:

  • 训练与验证预处理是否一致;
  • 训练数据是否泄漏验证信息;
  • 训练目标是否覆盖目标能力;
  • 微调数据格式和 loss mask 是否正确;
  • validation loss 与任务指标是否给出一致信号。

不能通过降低 learning rate 自动解决目标错位。

7.7 只有多卡训练异常

我使用相同初始 checkpoint 和相同数据,比较单卡与最小多卡配置的前几个 optimizer step,并检查:

  • 全局有效 token 数是否一致;
  • loss 是否按有效 token 正确归约;
  • 所有 rank 是否一致执行 backward 和 optimizer step;
  • grad norm 是否使用了正确的全局统计;
  • scheduler 是否按同一个 optimizer step 推进。

如果单卡正常、多卡异常,先验证归约和控制流,不先调整模型超参数。

7.8 诊断总表

现象 首先查看 最小验证 确认后才采取的处理
第一步 NaN 输入、activation、loss 单 batch FP32 重放 修复第一个非有限运算或精度配置
单步 loss spike 后恢复 batch 四项上下文、grad norm 重放相同 batch 修复损坏数据;合法难样本不直接丢弃
spike 后不恢复 相对更新量、optimizer state 健康 checkpoint 对照 修复根因后从健康 checkpoint 恢复
loss 长期不降 learning rate、step 状态、更新量 固定 batch 过拟合 修复更新链路或数据目标
grad norm 持续抬升 分层 grad norm、activation 降低单一风险因素做对照 调整已确认的结构、数据或步长问题
训练正常但验证变差 validation loss、任务指标 检查数据划分与评估链路 修复泄漏、目标错位或过拟合
只有多卡异常 归约和 rank 控制流 单卡/双卡对照 修复分布式实现

Part 8 从开跑到恢复:一套可以直接执行的训练流程

8.1 正式训练前

我在开跑前完成下面十项检查:

  • 固定代码版本、训练配置、数据版本和随机种子;
  • 随机检查 token、label、attention mask 和 loss mask;
  • 用一个 batch 完成 forward、backward 和 optimizer step;
  • 确认初始参数、loss 和梯度都是有限值;
  • 确认所有目标参数已经加入 AdamW 参数组;
  • 确认 weight decay 参数组符合既定配置;
  • 确认梯度累积按有效 token 正确归一化;
  • 确认 FP16 下先 unscale,再统计和裁剪梯度;
  • 确认 scheduler 按 optimizer step 推进;
  • 确认 checkpoint 包含 model、optimizer、scheduler 和 scaler state。

8.2 pilot run 是什么

pilot run(小规模试运行)是在投入完整训练资源之前,用受控的较小预算验证训练能否健康启动,并筛选关键配置。

它不是"随便跑几步",也不是固定使用总训练步数的某个百分比。一个有用的 pilot run 至少覆盖:

  1. 完整 warmup;
  2. 一小段达到峰值 learning rate 后的训练;
  3. 具有代表性的数据来源和序列长度;
  4. 至少一次 validation;
  5. 九项常驻核心监控。

如果训练启动风险通常在特定 token 规模后出现,pilot run 应覆盖该风险窗口,而不是机械地按 step 百分比截断。

8.3 怎样比较 learning rate 候选

我先从同类模型的可靠配置选择中心值,再比较少量候选。每个候选使用相同:

  • 初始 checkpoint;
  • 数据顺序;
  • 有效 batch;
  • warmup 口径;
  • 训练 token 预算;
  • validation 流程。

比较结果时,我同时看 validation loss、裁剪前 grad norm、相对更新量和实际吞吐。early training loss 更低,不自动代表最终配置更好。

8.4 正式训练期间

我把监控分成三个时间尺度:

  • 每步检查:loss 是否有限、梯度是否有限、optimizer step 是否执行;
  • 窗口检查:training loss、grad norm、裁剪行为和相对更新量是否持续偏离;
  • 周期检查:validation loss、任务指标和 checkpoint 可恢复性。

单步异常触发现场保存,持续异常触发停止或回滚,验证异常触发数据和目标检查。三者不能使用同一报警逻辑。

8.5 发生异常后怎样恢复

我的恢复顺序是:

  1. 停止让异常状态写入新的正式 checkpoint;
  2. 保存异常 step 的最小现场;
  3. 找到最后一个所有核心指标仍健康的 checkpoint;
  4. 用异常 batch 和普通 batch 分别重放;
  5. 通过单变量实验确认根因;
  6. 修复后重新执行最小正确性验证;
  7. 从健康 checkpoint 恢复正式训练。

如果非有限梯度已经进入 optimizer state,我不会只修复模型参数后继续使用可疑 optimizer state。


Part 9 训练规模发生变化后,我具体怎么调

规模变化后,最没有帮助的建议是"重新调参"。我需要的是一个默认动作:哪些配置先不动,哪些配置必须跟着改,出现什么证据后才继续调整。

我先用下面这张表做决策,后面再解释每种情况。

发生的变化 我的默认动作 什么时候继续调整
增加同分布新数据,并从头训练 保持 AdamW、峰值 learning rate 和有效 batch 不变;增加总训练 token;按新训练终点重算 decay pilot run 出现持续裁剪、相对更新量异常或 validation 变差
从已有 checkpoint 继续训练同分布数据 不把 learning rate 拉回初始峰值;从 checkpoint 当前 learning rate 平滑继续,并重新安排剩余 decay 当前 learning rate 已接近零且模型几乎不再学习
加入新领域或新来源数据 先保持旧 learning rate;为旧分布和新分布分别保留 validation 新数据引起 loss、grad norm 或相对更新量系统性偏移
对原数据训练更多 epoch 不重新使用初始峰值;继续 decay 或保持较低 learning rate validation 仍改善但更新量已经过小时才向上试探
只改变 micro-batch,有效 batch 不变 learning rate、warmup token 和 decay 全部不变 不需要因为 micro-batch 本身调参
有效 batch 扩大为原来的 \(k\) 倍 先保留旧 learning rate,再把 \(\sqrt{k}\) 倍作为第二候选 旧值稳定但 validation 改善明显变慢
序列长度增加,有效 token/step 不变 learning rate 不变;只调整 micro-batch 和梯度累积 attention logits 或梯度行为出现新异常
序列长度增加,有效 token/step 同时增加 按"有效 batch 增大"处理 比较旧值和平方根缩放候选
模型规模增大,使用标准参数化 固定数据和有效 batch,试 0.5×、0.7×、1.0× 原 learning rate 根据 validation、持续裁剪和相对更新量选择
从预训练进入微调 使用明显低于预训练的 learning rate,并保留基础能力回归评估 当前任务学不动时向上试;基础能力下降时向下调

9.1 增加同分布新数据

如果模型、数据分布、有效 batch 和优化器都没有变化,我的默认动作是:

峰值 learning rate 先不动,AdamW 配置先不动,有效 batch 先不动;我只增加总训练 token,并把 decay 的终点移动到新的训练终点。

例如,原计划训练 300B token,现在增加同分布数据后训练 500B token。我不会因为数据多了就自动放大 learning rate。我要重新计算的是 learning rate schedule:原本在 300B token 降到最低点的 decay,现在应该和新的 500B token 训练终点对齐。

如果从头训练,并且有效 batch 不变,我先保持原来的 warmup token 数,而不是保持 warmup 占总训练量的百分比。总训练量变大,不代表启动阶段需要同比例变长。

9.2 从 checkpoint 继续训练同分布数据

如果原训练已经结束,现在从最终 checkpoint 接着训练新增数据,我不会重新执行完整 warmup,也不会把 learning rate 突然拉回原始峰值。这样做可能直接破坏已经形成的模型状态。

我的默认做法是:

  1. 从 checkpoint 当前 learning rate 开始;
  2. 用短过渡保持 learning rate 连续,不制造突然跳变;
  3. 根据新增 token 预算重新安排后续 decay;
  4. 使用同一 validation 判断模型是否继续改善。

如果 checkpoint 的 learning rate 已经接近零,模型几乎没有更新,我才增加一个候选实验:使用明显低于原始峰值的新 learning rate,并配合短 warmup。这个新值必须通过 pilot run 确认,不能直接恢复到最初峰值。

9.3 加入新分布,或者重复训练原数据

新领域数据会改变模型看到的梯度分布。我先保持旧 learning rate,把旧分布 validation 和新分布 validation 分开记录。这样我才能判断新能力是否增长,以及旧能力是否退化。

如果新数据加入后出现持续更高的 grad norm、更强的裁剪或更大的相对更新量,我比较两个候选:

  • 原 learning rate;
  • 0.5 × 原 learning rate。

数据混合比例决定模型最终学习什么,属于训练目标和业务取舍,不能由稳定性指标替我决定。

如果只是对原数据训练更多 epoch,我不会把 learning rate 重置到初始峰值。我继续现有 decay,或者保持一个较低的恒定值,并用 validation 判断是否继续训练。validation 已经恶化时,正确动作通常是停止,而不是继续调学习率强行拟合。

9.4 batch 发生变化

我先区分 micro-batch 和有效 batch:

  • micro-batch 变小,但通过梯度累积保持每个 optimizer step 的有效 token 数不变:learning rate 和 schedule 都不改。
  • 每个 optimizer step 的有效 token 数扩大为原来的 \(k\) 倍:先保留旧 learning rate,再把 \(\sqrt{k}\) 倍作为第二候选。

如果总训练 token 不变,而有效 batch 扩大 \(k\) 倍,总 optimizer step 大约会缩短为原来的 \(1/k\)。为了保持相同 warmup token 数,warmup step 也应相应缩短,而不是继续使用原来的 step 数。

对 AdamW,我不默认使用 \(k\) 倍线性放大学习率。只有旧 learning rate 下训练稳定、相对更新量偏小、validation 改善明显变慢时,我才尝试更大的候选。

9.5 序列长度增加

如果序列长度增加,但我通过减小 micro-batch 保持每个 optimizer step 的有效 token 数不变,learning rate 先不改。我只检查:

  1. attention mask 和 packing 边界;
  2. loss 是否仍按有效 token 归一化;
  3. attention logits 是否出现新的异常范围;
  4. micro-batch 与梯度累积是否仍然对齐。

如果序列长度增加的同时,每个 optimizer step 的有效 token 数也增加,我就把它视为有效 batch 增大,按照上一节比较旧 learning rate 和平方根缩放候选。

9.6 模型规模增大

在标准参数化下,我不会直接套用"参数量扩大十倍,learning rate 降低几倍"的固定公式。我的默认候选是原 learning rate 的:

\0.5\\times,\\qquad 0.7\\times,\\qquad 1.0\\times \\

三个实验保持数据顺序、有效 batch、训练 token、优化器和 validation 完全一致。我按下面的顺序选择:

  1. 排除出现非有限值、持续强裁剪或异常相对更新量的候选;
  2. 比较同等 token 预算下的 validation loss;
  3. 如果差异很小,选择更新更受控的较低 learning rate;
  4. 确定峰值后,再调整 warmup 和 decay。

如果模型使用 \(\mu\)P 或其他专门的参数化规则,应遵循对应规则重新设计候选,不能把上面的标准参数化建议直接搬过去。

9.7 从预训练进入微调

微调从已经具备能力的 checkpoint 开始,我的默认动作是使用明显低于预训练的 learning rate,同时检查当前任务和基础能力:

现象 我的下一步
当前任务 loss 不降,相对更新量很小 提高一个 learning rate 档位
当前任务改善,但基础能力明显下降 降低 learning rate 或减少训练步数
training loss 下降,validation 不改善 检查数据、loss mask 和过拟合,不先提高 learning rate
两类 validation 都改善,更新保持受控 保持配置继续训练

微调的具体 learning rate 仍应从同类模型和任务的可靠配置出发。稳定性手册负责给调整方向,不替任务实验决定唯一数值。


Part 10 边界说明:不要把 token embedding 和超大稀疏 ID 表混为一谈

10.1 本文的直接结论

在标准大模型预训练和微调中,我默认让 token embedding 跟随模型主体使用 AdamW。除非已有训练方案明确把 embedding 放进不同参数组,否则没有必要仅因为输入是离散 token,就单独换成 Adagrad。

10.2 AdamW 和 Adagrad 的关系

二者都使用历史梯度平方,按参数调整有效步长,这是它们相似的地方。但二者不是同一个算法。

Adagrad 累加从训练开始以来的梯度平方:

\G_t=G_{t-1}+g_t\^2 \\

AdamW 使用指数移动平均,较早历史会逐渐衰减:

\v_t=\\beta_2v_{t-1}+(1-\\beta_2)g_t\^2 \\

AdamW 还包含一阶矩和解耦 weight decay。因此,更准确的说法是:AdamW 延续了自适应梯度方法按历史梯度平方调整步长的思想,但它不等于把 Adagrad 原样包含进来。

对本文的大模型主线,这个算法差异不会改变操作结论:token embedding 继续使用既定 AdamW 配置。

10.3 稀疏数据真正解决不了的是什么

某个 token 或 ID 很少出现,首先意味着模型几乎没有关于它的训练数据。AdamW 和 Adagrad 都只能调整已经观察到的梯度,不能替从未出现的 ID 创造训练信号。

如果低频 token 的学习效果不足,我优先检查数据覆盖、分词方式和采样策略,而不是先更换优化器。

推荐系统的超大 ID 表还涉及参数存储、梯度布局和优化器状态成本,这是另一套工程问题,不属于本文的大模型训练稳定性主线。这里不再扩展监控指标,也不替推荐系统选择统一优化器。


附录 A:核心概念速查

中文名称 英文名称或缩写 定义 我用它判断什么
训练损失 training loss 当前训练数据上的目标函数值 模型是否在优化当前目标
验证损失 validation loss 固定验证集上的目标函数值 模型是否真正改善和泛化
困惑度 perplexity 语言模型平均 token loss 的指数形式 模型对验证 token 的不确定程度
激活值 activation 前向传播中各层产生的中间张量 哪一层首先发生数值异常
梯度 gradient loss 对参数的偏导数 当前误差信号怎样作用于参数
梯度范数 gradient norm / grad norm 目标梯度集合的 L2 范数 梯度整体是否出现尖峰或持续漂移
梯度裁剪 gradient clipping 超过阈值时按比例缩小梯度 限制当前 step 的更新风险
更新范数 update norm 参数更新量的 L2 范数 当前 step 的绝对参数变化量
权重范数 weight norm 参数本身的 L2 范数 参数当前的绝对尺度
相对更新量 update-to-weight ratio update norm 除以 weight norm 当前更新相对参数尺度有多大
学习率 learning rate 控制优化器整体更新尺度的系数 当前训练阶段的计划步长
学习率预热 learning rate warmup 从较小学习率逐步升至峰值 降低训练启动阶段的大步更新风险
自动混合精度 Automatic Mixed Precision / AMP 按算子选择高低精度计算 在性能和数值可靠性之间取得平衡
损失缩放 loss scaling 放大 loss 和反向梯度后再还原 减少 FP16 小梯度下溢
优化器更新步 optimizer step 优化器真正改变参数的一次操作 区分计算过 backward 和真正更新过参数
小规模试运行 pilot run 正式训练前的受控短训练 验证启动稳定性并筛选关键配置

困惑度与 token 平均交叉熵的关系通常为:

\\\text{Perplexity}=\\exp(\\text{Token Average Loss}) \\

不同数据集、分词器和 loss mask 下的 perplexity 不能直接横向比较。


附录 B:关键指标的统一计算口径

B.1 training loss

  • 分子:所有有效 token 的 loss 之和;
  • 分母:所有数据并行 rank 和所有 micro-batch 的有效 token 总数;
  • 单步值:对应一次 optimizer step 累积的数据;
  • 滑动平均:只用于观察趋势,不覆盖原始单步值。

B.2 grad norm

  • 统计时间:完整有效 batch backward 完成之后;
  • FP16:在 scaler.unscale_(optimizer) 之后;
  • 裁剪:记录裁剪前数值;
  • 参数范围:必须明确是完整模型、完整参数组还是当前 shard;
  • 分布式:使用训练框架能够表示完整目标参数集合的实现。

B.3 clip coefficient

若全局裁剪阈值为 \(c\),裁剪前 grad norm 为 \(G\),缩放系数可以写为:

\\\alpha=\\min\\left(1,\\frac{c}{G+\\epsilon}\\right) \\

\(\alpha=1\) 表示没有缩小梯度;\(\alpha\) 越小,当前 step 受到的裁剪越强。

B.4 update norm 和相对更新量

  • 统计对象:四个已定义的代表性参数组;
  • 统计时间:optimizer step 前后;
  • FP16:使用 FP32 转换后计算范数,减少统计本身的精度误差;
  • tied weight:共享参数只统计一次;
  • skipped step:更新量记录为未执行,而不是和正常的零更新混为一谈;
  • 采样频率:固定间隔采集,保持实验间一致。

如果我希望把 weight decay 的贡献和梯度更新贡献分开,需要从优化器内部直接取得分量。仅通过 \(\theta_{t+1}-\theta_t\) 计算的 update norm 包含该 optimizer step 的全部参数变化。

B.5 optimizer step 执行状态

这个指标必须是显式布尔状态或单调递增的真实更新计数,不能用 micro-step 或 dataloader iteration 代替。FP16 发生溢出时,即使完成 backward,参数也可能没有更新。


附录 C:最小 PyTorch 训练循环

下面的示例强调执行顺序。实际分布式训练应使用框架提供的全局 grad norm 和参数分片接口,不能直接假设 model.parameters() 表示所有全局参数。

python 复制代码
import torch


def tensors_are_finite(tensors):
    return all(
        tensor is None or torch.isfinite(tensor).all().item()
        for tensor in tensors
    )


optimizer.zero_grad(set_to_none=True)

with torch.autocast(device_type="cuda", dtype=torch.float16):
    outputs = model(
        input_ids=batch["input_ids"],
        attention_mask=batch["attention_mask"],
    )
    loss_sum = token_loss_sum(outputs.logits, batch["labels"])
    valid_tokens = batch["loss_mask"].sum()

if valid_tokens.item() == 0:
    raise ValueError("当前 optimizer step 没有有效 token")

loss = loss_sum / valid_tokens

if not torch.isfinite(loss):
    save_debug_context(batch=batch, loss=loss)
    raise FloatingPointError("loss 出现非有限值")

scaler.scale(loss).backward()
scaler.unscale_(optimizer)

if not tensors_are_finite(parameter.grad for parameter in model.parameters()):
    save_debug_context(batch=batch, loss=loss)
    raise FloatingPointError("gradient 出现非有限值")

grad_norm = torch.nn.utils.clip_grad_norm_(
    model.parameters(),
    max_norm=clip_threshold,
)

# 如果需要计算相对更新量,只在固定采样 step 保存代表性参数组。
before_update = snapshot_monitored_parameter_groups(model)

previous_scale = scaler.get_scale()
scaler.step(optimizer)
scaler.update()
current_scale = scaler.get_scale()

# 具体框架应直接暴露 step 是否执行;这里仅展示 FP16 GradScaler 的判断思路。
optimizer_step_executed = current_scale >= previous_scale

if optimizer_step_executed:
    scheduler.step()
    relative_updates = compute_relative_updates(
        before_update=before_update,
        model=model,
    )
else:
    relative_updates = None

optimizer.zero_grad(set_to_none=True)

这段代码只展示控制顺序,不定义分布式 loss 聚合、参数分片和监控函数。生产训练代码应复用现有训练框架的接口,并为这些接口写最小验证。

如果使用 BF16,通常不需要 GradScaler。执行链可以简化为:

forward → loss 有限值检查 → backward → gradient 有限值检查 → grad norm → clip → optimizer step → scheduler step → 相对更新量


附录 D:训练稳定性检查清单

D.1 开跑前

  • 数据、代码、配置和随机种子可以追溯;
  • token、label、attention mask 和 loss mask 已抽查;
  • loss 按全局有效 token 正确归一化;
  • 固定 batch 可以明显降低 loss;
  • 参数组、weight decay 和冻结规则符合预期;
  • AMP、梯度累积、裁剪和 scheduler 顺序正确;
  • 九项常驻指标均有明确统计口径;
  • checkpoint 可以完整恢复训练状态。

D.2 pilot run

  • 覆盖完整 warmup;
  • 覆盖一小段峰值 learning rate 阶段;
  • 覆盖代表性数据来源和序列长度;
  • 完成至少一次 validation;
  • 没有持续非有限值或 optimizer step 跳过;
  • grad norm 没有持续偏离健康趋势;
  • 四组参数均有受控更新;
  • validation loss 或任务指标朝预期方向变化。

D.3 正式训练

  • 单步数据和滑动平均同时保留;
  • 每次异常都能定位到 step 和数据分片;
  • skipped step、连续裁剪和持续漂移有报警;
  • validation 使用固定数据和固定口径;
  • checkpoint 定期执行可恢复性验证。

D.4 异常恢复

  • 已保存最小现场;
  • 已确认第一个异常信号;
  • 已分别重放异常 batch 和普通 batch;
  • 已用单变量对照实验确认根因;
  • 已确认最后一个健康 checkpoint;
  • 修复后重新完成最小正确性验证;
  • 没有继续使用被污染的 optimizer state。

参考资料

下面的论文支撑原理与实验现象,官方文档支撑具体工程操作。论文中的配置和阈值只代表相应实验设置,不自动成为通用规则。

优化器与参数更新

  1. Kingma, Ba. Adam: A Method for Stochastic Optimization。用于理解一阶矩、二阶矩和偏差修正。
  2. Loshchilov, Hutter. Decoupled Weight Decay Regularization。用于理解 AdamW 和解耦权重衰减。
  3. Duchi, Hazan, Singer. Adaptive Subgradient Methods for Online Learning and Stochastic Optimization。用于核对 Adagrad 累积历史梯度平方的定义。
  4. PyTorch. AdamW documentation。用于核对 AdamW 的当前接口、公式和参数含义。

混合精度与数值稳定

  1. Micikevicius et al. Mixed Precision Training。用于理解 FP16 混合精度和 loss scaling。
  2. PyTorch. Automatic Mixed Precision package。用于核对 autocast、GradScaler 和梯度缩放行为。
  3. PyTorch. Automatic Mixed Precision examples。用于核对梯度累积、unscale 和 clipping 的正确顺序。
  4. NVIDIA. Mixed-Precision Training of Deep Neural Networks。用于理解低精度训练的工程背景。

Transformer 梯度与训练尖峰

  1. Xiong et al. On Layer Normalization in the Transformer Architecture。用于理解 Pre-LN、Post-LN 和训练启动行为的适用边界。
  2. Kalra, Barkeshli. Why Warmup the Learning Rate? Underlying Mechanisms and Improvements。用于理解 warmup 与初始化、曲率和优化动态的关系。
  3. Takase et al. Spike No More: Stabilizing the Pre-training of Large Language Models。用于理解 Transformer 子层尺度、梯度范数和 loss spike 之间的关系。
  4. Fan et al. HLAT: High-quality Large Language Model Pre-trained on AWS Trainium。用于参考长周期大模型训练中的 loss、grad norm 和 parameter norm 监控案例。

扩容与适用边界

  1. Brown et al. Language Models are Few-Shot Learners。用于参考 GPT-3 各模型规模的具体 batch token 和 learning rate 配置,不用于推导通用缩放公式。
  2. Hoffmann et al. Training Compute-Optimal Large Language Models。用于理解固定计算预算下模型规模和训练 token 的共同选择,不作为训练稳定性阈值。
  3. Yang et al. Tensor Programs V: Tuning Large Neural Networks via Zero-Shot Hyperparameter Transfer。用于理解参数化方式与超参数迁移的关系。

数据、loss 和梯度累积

  1. Hugging Face. Fixing Gradient Accumulation。用于理解变长序列下按有效 token 聚合 loss 的工程问题。

写在最后

训练稳定不是某一个指标落在固定区间,也不是 loss 每一步都下降。它是一条完整训练链路持续保持健康的结果。

如果只保留一套判断顺序,我会记住下面五句话:

  1. 先确认数值能不能算。
  2. 再确认梯度能不能传。
  3. 再确认参数有没有以合理幅度更新。
  4. 再确认模型是不是在优化正确目标。
  5. 最后用验证结果判断训练是否真正有效。

看到异常时,我不先猜 learning rate,也不先堆更多监控。我会保存现场,沿训练链寻找第一个异常信号,再用最小重放和单变量实验确认根因。

稳定训练的核心,不是消灭所有波动,而是让每一次异常都能够被观察、解释、验证和处理。