观文听傑

返回

上一篇用 Token-based Batching(按 token 预算组批)让长样本少占几行、短样本多占几行。代价是每个 optimizer step 真正看见的有效 token 不再恒定:两个实验都跑了 10,000 步,可能已经消费了完全不同的数据量。

本篇只解决一个核心问题:如何把学习率写成“已学习多少 token”的函数,并正确安排 Warmup(预热)与 Cosine Decay(余弦衰减)。

01 为什么 epoch 和 step 都可能是错时钟?#

Epoch(轮次)假设数据集边界稳定;流式语料、按权重重复采样和持续去重会让“一轮”含义模糊。Step(优化步)比 epoch 明确,但在可变 batch 下,第 ss 步的有效 token 数 nsn_s 会变化。

真正的数据进度是累计监督 token:

qs=i=1sniq_s=\sum_{i=1}^{s} n_i

其中 qsq_s 是完成第 ss 次参数更新后累计消费的有效 token。若训练预算为 QQ,日程进度就是 ps=min(qs/Q,1)p_s=\min(q_s/Q,1)

flowchart LR
  A[micro-batches] --> B[统计本窗口有效 token]
  B --> C[反向传播与梯度归一化]
  C --> D[optimizer.step]
  D --> E[累计 global_tokens]
  E --> F[计算下一个学习率]
  F --> G[写入 param_groups]
  G --> A
mermaid

注意数据流顺序:当前窗口用当前学习率更新;更新成功后推进 token 时钟,再为下一窗口设置学习率。

02 Warmup 在保护什么?#

训练刚开始时,参数、激活尺度和 Adam 的一二阶矩估计都还没有稳定。直接使用峰值学习率 ηmax\eta_{\max},一次噪声较大的更新就可能破坏表示。线性 warmup 在前 WW 个 token 将学习率从较小值升到峰值:

η(q)=ηmaxqW,0q<W\eta(q)=\eta_{\max}\frac{q}{W},\qquad 0\le q<W

Warmup 不是“先不学习”,而是逐渐放大更新。它也不是修复错误归一化、异常梯度或过大峰值学习率的万能补丁。

03 Cosine Decay 怎样把更新慢慢收紧?#

Warmup 后,令

r=clip(qWQW,0,1)r=\operatorname{clip}\left(\frac{q-W}{Q-W},0,1\right)

从峰值平滑衰减到最低比例 αηmax\alpha\eta_{\max}

η(q)=ηmax[α+(1α)1+cos(πr)2]\eta(q)=\eta_{\max}\left[\alpha+(1-\alpha)\frac{1+\cos(\pi r)}{2}\right]

QQ 是总 token 预算,WW 是 warmup token,α[0,1]\alpha\in[0,1]。余弦的端点斜率为 0,切换平滑;但它并不自动知道最优训练长度,QQ 仍是实验设计。

04 用 100 个 token 手算完整日程#

Q=100Q=100W=20W=20ηmax=103\eta_{\max}=10^{-3}α=0.1\alpha=0.1

累计 token qq阶段比例 η/ηmax\eta/\eta_{\max}学习率
0warmup00
10warmup0.55.0e45.0e-4
20峰值11.0e31.0e-3
60decay0.1+0.9×0.5=0.550.1+0.9\times0.5=0.555.5e45.5e-4
100末端0.11.0e41.0e-4

若每步 token 是 [8, 12, 30, 10],更新后的时钟依次为 [8,20,50,60],而不是 [1,2,3,4]。第三步跨过多个“虚拟刻度”没有问题:日程是连续函数,不要求每个 token 都调用一次 scheduler。

05 Step 时钟什么时候仍然够用?#

若世界大小、梯度累积次数和每个 micro-batch 的有效 token 都固定,则 qs=snq_s=s\cdot n,按 step 与按 token 只是横轴单位不同。只要以下任一项改变,固定 step 日程就会漂移:

  • 长度分桶导致有效 token 波动;
  • OOM 后减小 micro-batch、增加累积次数;
  • 扩容改变 data-parallel world size;
  • 部分窗口因 FP16 overflow 被跳过;
  • loss mask 改变有效监督位置数。

06 一个与框架无关的日程函数#

先把数学写成纯函数,边界条件才容易单测。

import math

def warmup_cosine_multiplier(tokens, warmup_tokens, total_tokens, min_ratio=0.1):
    if not 0 <= warmup_tokens < total_tokens:
        raise ValueError("need 0 <= warmup_tokens < total_tokens")
    if not 0 <= min_ratio <= 1:
        raise ValueError("min_ratio must be in [0, 1]")
    q = min(max(int(tokens), 0), total_tokens)
    if warmup_tokens and q < warmup_tokens:
        return q / warmup_tokens
    progress = (q - warmup_tokens) / (total_tokens - warmup_tokens)
    cosine = 0.5 * (1.0 + math.cos(math.pi * progress))
    return min_ratio + (1.0 - min_ratio) * cosine
python

输入和输出都是标量:tokens: int -> multiplier: float。至少断言 m(0)=0m(0)=0m(W)=1m(W)=1m(Q)=αm(Q)=\alpha,并密集采样验证 warmup 单调递增、decay 单调递减且没有负数。

07 用当前 PyTorch API 落地#

PyTorch 2.14 的 LambdaLR 接受整数计数并返回相对初始学习率的乘数;官方要求 scheduler.step()optimizer.step() 之后调用。但它的参数名仍叫 epoch,不会替你统计 token。因此可直接让训练循环维护 token 时钟,并把纯函数结果写入参数组,语义更清楚:

多参数组时,例如 backbone 与 head 的初始学习率分别为 1e-41e-3,两者乘同一曲线并保留 10 倍比例。

08 训练循环中究竟统计哪一刻?#

optimizer.zero_grad(set_to_none=True)
window_tokens = 0

for micro in accumulation_window:
    logits = model(micro["input_ids"])          # [B,L,V]
    loss_sum = token_loss_sum(logits, micro["labels"])
    valid = micro["labels"].ne(-100).sum()      # scalar int64
    loss_sum.backward()
    window_tokens += int(valid)

normalize_gradients(model.parameters(), window_tokens)
optimizer.step()
schedule.step(successful_tokens=window_tokens)
optimizer.zero_grad(set_to_none=True)
python

分布式训练中,window_tokens 应先做全局 sum;每个 rank 必须得到相同的 global_tokens 和学习率。若 AMP 的 scaler 检测到非有限梯度并跳过 optimizer.step(),本窗口不应推进“成功优化”的日程时钟,否则曲线走了、参数却没走。

09 Token 日程与梯度累积不要混为一谈#

梯度累积决定多少 micro-batch 合成一次参数更新;token 日程决定该更新使用多大学习率。即使两个窗口都含 32K token,一个由 8 个 micro-batch 累积、另一个由 4 个组成,只要梯度按同一全局 token 分母归一化,它们的时钟增量相同。

反过来,若固定每 4 个 micro-batch 更新,但各批有效 token 不同,不能假装每步都是 32K。应记录三项:micro-step、optimizer-step、global-effective-tokens。

10 预算改变时能否中途重画曲线?#

QQ 从 100B 临时改成 200B 会改变当前点之后的余弦位置,甚至让学习率瞬间升高。可选策略有:

  1. 训练前固定预算,最易比较与复现;
  2. 延长时从当前学习率重新定义一段连续曲线,并记录新阶段;
  3. 使用与固定终点无关的逆平方根等日程,但仍需验证最终质量。

不要静默修改 total_tokens。配置、日志和 checkpoint 必须保留每个阶段的边界,否则同名实验无法解释。

11 常见错误与最短调试路径#

症状常见原因最短检查
第一次更新学习率为峰值初始化后未应用 q=0q=0 的倍率记录首三次 update 前后的 lr
日程比预期快 world_size 倍每个 rank 各自累加全局 tokenall-reduce 后只用统一总数推进
恢复后学习率突然跳变只恢复 optimizer,没恢复 token 时钟比较 checkpoint 内 tokens
AMP overflow 时曲线偷跑更新被跳过仍推进 schedule同时记录 step 是否成功和 scale
换 batch 策略后结果漂移仍按 optimizer-step 调度画 lr 对 global token,而非 step
末端学习率变成负数没 clamp 超预算进度测试 q>Qq>Q

一次有效的 dry run 不需要模型:喂入人工 token 序列,输出 (step, delta_tokens, total_tokens, lr),再与公式表逐项比对。

12 它会在哪里失败?#

Token 不是所有任务的自然样本单位。图像分类可能按样本数,强化学习可能按环境步,生成式训练还可能区分输入 token 与产生 loss 的目标 token。关键不是迷信 token,而是选择与统计目标一致、可跨配置比较的进度单位。

Warmup + cosine 也不是唯一日程。Inverse-square-root 在 Transformer 中常见;constant-with-warmup 适合还不知道终点的持续训练;ReduceLROnPlateau 依赖验证指标,但大规模预训练的验证噪声与成本可能让反馈滞后。日程选择不能替代峰值学习率搜索。

13 今天真正需要记住什么?#

  1. 可变 batch 下,optimizer step 不等于固定学习量;累计有效 token 是更稳定的训练时钟。
  2. Warmup 逐渐放大早期更新,cosine decay 在既定预算内平滑收紧更新。
  3. 只有成功的参数更新才推进日程;分布式所有 rank 必须共享同一个 token 计数。
  4. 日程状态与 optimizer 状态同样属于 checkpoint,预算变化必须显式版本化。

14 思考题与小练习#

  1. 设每步有效 token 为 [8,12,30,10],用本文参数计算每次更新后“下一步”的学习率,并与固定每步 15 token 的近似比较。
  2. 给纯函数增加 5% 的非零起始倍率,写出所有端点和单调性测试。
  3. 设计 DDP 日志字段,使你能发现某个 rank 少消费了一批但训练没有立刻死锁的问题。

相关工作#

  1. Vaswani et al., Attention Is All You Need,用 warmup 与逆平方根衰减训练原始 Transformer。
  2. Devlin et al., BERT: Pre-training of Deep Bidirectional Transformers for Language Understanding,展示大规模预训练中的 warmup 与线性衰减配置。
  3. Loshchilov & Hutter, SGDR: Stochastic Gradient Descent with Warm Restarts,系统提出余弦退火与重启。
  4. Kaplan et al., Scaling Laws for Neural Language Models,讨论模型、数据与计算预算的标度关系。
  5. Hoffmann et al., Training Compute-Optimal Large Language Models,说明 token 预算为何是预训练设计的核心变量。

15 下一篇预告#

学习率曲线已经可解释,但矩阵乘法全用 FP32 会浪费现代加速器吞吐。下一篇将研究自动混合精度、FP16/BF16 的数值范围与动态梯度缩放。

训练步数相同为何学习进度不同?Token 学习率时钟、Warmup 与 Cosine Decay
https://zwjcode.cn/blog/token-learning-rate-warmup-cosine-schedule
作者
发布于 2026年9月13日
版权协议 CC BY-NC-SA 4.0
评论加载似乎遇到了问题,请尝试刷新页面。