训练步数相同为何学习进度不同?Token 学习率时钟、Warmup 与 Cosine Decay
从可变长度 batch 让每步 token 数漂移出发,手算 token 进度下的 warmup 与余弦衰减,实现可恢复的 PyTorch 学习率控制器,并解释预算、调用顺序与调试方法。
上一篇用 Token-based Batching(按 token 预算组批)让长样本少占几行、短样本多占几行。代价是每个 optimizer step 真正看见的有效 token 不再恒定:两个实验都跑了 10,000 步,可能已经消费了完全不同的数据量。
本篇只解决一个核心问题:如何把学习率写成“已学习多少 token”的函数,并正确安排 Warmup(预热)与 Cosine Decay(余弦衰减)。
01 为什么 epoch 和 step 都可能是错时钟?#
Epoch(轮次)假设数据集边界稳定;流式语料、按权重重复采样和持续去重会让“一轮”含义模糊。Step(优化步)比 epoch 明确,但在可变 batch 下,第 步的有效 token 数 会变化。
真正的数据进度是累计监督 token:
其中 是完成第 次参数更新后累计消费的有效 token。若训练预算为 ,日程进度就是 。
flowchart LR
A[micro-batches] --> B[统计本窗口有效 token]
B --> C[反向传播与梯度归一化]
C --> D[optimizer.step]
D --> E[累计 global_tokens]
E --> F[计算下一个学习率]
F --> G[写入 param_groups]
G --> Amermaid注意数据流顺序:当前窗口用当前学习率更新;更新成功后推进 token 时钟,再为下一窗口设置学习率。
02 Warmup 在保护什么?#
训练刚开始时,参数、激活尺度和 Adam 的一二阶矩估计都还没有稳定。直接使用峰值学习率 ,一次噪声较大的更新就可能破坏表示。线性 warmup 在前 个 token 将学习率从较小值升到峰值:
Warmup 不是“先不学习”,而是逐渐放大更新。它也不是修复错误归一化、异常梯度或过大峰值学习率的万能补丁。
03 Cosine Decay 怎样把更新慢慢收紧?#
Warmup 后,令
从峰值平滑衰减到最低比例 :
是总 token 预算, 是 warmup token,。余弦的端点斜率为 0,切换平滑;但它并不自动知道最优训练长度, 仍是实验设计。
04 用 100 个 token 手算完整日程#
设 、、、。
| 累计 token | 阶段 | 比例 | 学习率 |
|---|---|---|---|
| 0 | warmup | 0 | 0 |
| 10 | warmup | 0.5 | |
| 20 | 峰值 | 1 | |
| 60 | decay | ||
| 100 | 末端 | 0.1 |
若每步 token 是 [8, 12, 30, 10],更新后的时钟依次为 [8,20,50,60],而不是 [1,2,3,4]。第三步跨过多个“虚拟刻度”没有问题:日程是连续函数,不要求每个 token 都调用一次 scheduler。
05 Step 时钟什么时候仍然够用?#
若世界大小、梯度累积次数和每个 micro-batch 的有效 token 都固定,则 ,按 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) * cosinepython输入和输出都是标量:tokens: int -> multiplier: float。至少断言 、、,并密集采样验证 warmup 单调递增、decay 单调递减且没有负数。
07 用当前 PyTorch API 落地#
PyTorch 2.14 的 LambdaLR 接受整数计数并返回相对初始学习率的乘数;官方要求 scheduler.step() 在 optimizer.step() 之后调用。但它的参数名仍叫 epoch,不会替你统计 token。因此可直接让训练循环维护 token 时钟,并把纯函数结果写入参数组,语义更清楚:
class TokenLRSchedule:
def __init__(self, optimizer, warmup_tokens, total_tokens, min_ratio=0.1):
self.optimizer = optimizer
self.base_lrs = [g["lr"] for g in optimizer.param_groups]
self.warmup_tokens = warmup_tokens
self.total_tokens = total_tokens
self.min_ratio = min_ratio
self.tokens = 0
self._apply()
def _apply(self):
scale = warmup_cosine_multiplier(
self.tokens, self.warmup_tokens,
self.total_tokens, self.min_ratio,
)
for group, base_lr in zip(self.optimizer.param_groups, self.base_lrs):
group["lr"] = base_lr * scale
def step(self, successful_tokens):
if successful_tokens <= 0:
raise ValueError("successful_tokens must be positive")
self.tokens += int(successful_tokens)
self._apply()
def state_dict(self):
return {"tokens": self.tokens, "base_lrs": self.base_lrs}
def load_state_dict(self, state):
self.tokens = int(state["tokens"])
self.base_lrs = list(state["base_lrs"])
self._apply()python多参数组时,例如 backbone 与 head 的初始学习率分别为 1e-4、1e-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 预算改变时能否中途重画曲线?#
把 从 100B 临时改成 200B 会改变当前点之后的余弦位置,甚至让学习率瞬间升高。可选策略有:
- 训练前固定预算,最易比较与复现;
- 延长时从当前学习率重新定义一段连续曲线,并记录新阶段;
- 使用与固定终点无关的逆平方根等日程,但仍需验证最终质量。
不要静默修改 total_tokens。配置、日志和 checkpoint 必须保留每个阶段的边界,否则同名实验无法解释。
11 常见错误与最短调试路径#
| 症状 | 常见原因 | 最短检查 |
|---|---|---|
| 第一次更新学习率为峰值 | 初始化后未应用 的倍率 | 记录首三次 update 前后的 lr |
| 日程比预期快 world_size 倍 | 每个 rank 各自累加全局 token | all-reduce 后只用统一总数推进 |
| 恢复后学习率突然跳变 | 只恢复 optimizer,没恢复 token 时钟 | 比较 checkpoint 内 tokens |
| AMP overflow 时曲线偷跑 | 更新被跳过仍推进 schedule | 同时记录 step 是否成功和 scale |
| 换 batch 策略后结果漂移 | 仍按 optimizer-step 调度 | 画 lr 对 global token,而非 step |
| 末端学习率变成负数 | 没 clamp 超预算进度 | 测试 |
一次有效的 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 今天真正需要记住什么?#
- 可变 batch 下,optimizer step 不等于固定学习量;累计有效 token 是更稳定的训练时钟。
- Warmup 逐渐放大早期更新,cosine decay 在既定预算内平滑收紧更新。
- 只有成功的参数更新才推进日程;分布式所有 rank 必须共享同一个 token 计数。
- 日程状态与 optimizer 状态同样属于 checkpoint,预算变化必须显式版本化。
14 思考题与小练习#
- 设每步有效 token 为
[8,12,30,10],用本文参数计算每次更新后“下一步”的学习率,并与固定每步 15 token 的近似比较。 - 给纯函数增加 5% 的非零起始倍率,写出所有端点和单调性测试。
- 设计 DDP 日志字段,使你能发现某个 rank 少消费了一批但训练没有立刻死锁的问题。
相关工作#
- Vaswani et al., Attention Is All You Need ↗,用 warmup 与逆平方根衰减训练原始 Transformer。
- Devlin et al., BERT: Pre-training of Deep Bidirectional Transformers for Language Understanding ↗,展示大规模预训练中的 warmup 与线性衰减配置。
- Loshchilov & Hutter, SGDR: Stochastic Gradient Descent with Warm Restarts ↗,系统提出余弦退火与重启。
- Kaplan et al., Scaling Laws for Neural Language Models ↗,讨论模型、数据与计算预算的标度关系。
- Hoffmann et al., Training Compute-Optimal Large Language Models ↗,说明 token 预算为何是预训练设计的核心变量。
15 下一篇预告#
学习率曲线已经可解释,但矩阵乘法全用 FP32 会浪费现代加速器吞吐。下一篇将研究自动混合精度、FP16/BF16 的数值范围与动态梯度缩放。