跳到主要内容

数据截至 (上游 commit 92d63d4e8bb4)

02 · GPT 模型与预训练

这一章讲什么: 全库最重的一章。四件事:模型长什么样(nanochat/gpt.py)、--depth 一个旋钮怎么推出全部超参(scripts/base_train.py)、优化器为什么是 Muon + AdamW 且不用 DDP(nanochat/optim.py)、数据怎么装(nanochat/dataloader.py)。最后附带 FP8 与常见坑。


1. 它要解决的小问题

预训练这一章要同时回答三个问题:

  • 配方:2026 年的「标准 GPT」相比 GPT-2 时代改了哪些零件,每处为什么。
  • 超参自动化:宽度 / 批大小 / lr / 步数 / weight decay 怎么不手调——用户只给 --depth
  • 分布式:8 卡训练怎么不包 DDP 也能同步梯度,还顺手把优化器显存砍到 1/N。

2. 模型:modern-recipe GPT

直觉

GPT-2 之后的几年里,社区把 Transformer 的每个零件都换了更顺手的版本。nanochat 把这些「已成共识的小改进」全收了——nanochat/gpt.py:1-13 的文件头注释就是清单。

配方一览

零件GPT-2 时代nanochat位置
位置编码学习式绝对位置 embeddingrotary,无位置 embeddingapply_rotary_emb nanochat/gpt.py:57-64
注意力稳定QK norm + q,k 各乘 1.2nanochat/gpt.py:101-103
头结构MHAGQA(n_kv_headn_headnanochat/gpt.py:72-78
MLP 激活GELUrelu²MLP.forward nanochat/gpt.py:137-141
embedding 绑定tieduntiednanochat/gpt.py:174-177
normLayerNorm(带参数)RMSNorm、无可学参数norm nanochat/gpt.py:42-43
bias全无各处 Linear(..., bias=False)
logit 处理softcap 15(tanh 软顶)nanochat/gpt.py:511-515

表格外还有三个「modded-nanogpt 系」的新零件,是 speedrun 快的重要原因,单独讲。

2.1 滑窗注意力(SWA)

  • 每层要么 L(全上下文)要么 S(约 1/4 上下文),window_pattern="SSSL" 平铺到各层,最后一层恒为 L(GPTConfig nanochat/gpt.py:36-39)。
  • _compute_window_sizesnanochat/gpt.py:287-317)把字符串翻译成 FA3 的 (left, right) 窗口参数;S 层窗口向上取整到 128 的倍数(FA3 tile 对齐,2048 → 768)。
  • 收益直接体现在 FLOPs 估计里:estimate_flopsnanochat/gpt.py:319-339)按每层窗口分别算注意力开销——S 层越多,训练越省。

2.2 Value Embeddings(ResFormer 式)

  • 交替层(且最后一层必含)多一张 value_embeds 嵌入表:把 token 嵌入当作「先验 value」,以门控方式加进注意力的 v(has_ve nanochat/gpt.py:53-54;定义 nanochat/gpt.py:189-192)。
  • 门控是输入相关、每头一个:gate = 3 * sigmoid(ve_gate(x[..., :12])),范围 (0, 3)(nanochat/gpt.py:92-97)。

2.3 smear / backout / resid+x0 lambdas

  • smear:把前一个 token 的嵌入按门控混入当前位置——白送的 bigram 信息(nanochat/gpt.py:478-492)。推理时靠 KVCache 上的 prev_embedding 传递(第四章)。
  • backout:在中点层缓存残差,投影 logits 前减掉 0.2 倍——先去掉低层特征再出词(nanochat/gpt.py:505-507)。
  • resid_lambdas / x0_lambdas:每层两个标量,x = λ_resid·x + λ_x0·x0——残差流缩放 + 初始嵌入回混(nanochat/gpt.py:498-500)。初始化有讲究:resid 从 1.15 随深度线性降到 1.05,x0 从 0.20 降到 0.05(init_weights nanochat/gpt.py:238-243)。

初始化与元设备

  • GPT.__init__ 跑在 meta device 上下文:只算形状与 dtype,不填数据;真正的初始化集中在 init_weightsnanochat/gpt.py:204-284)——embedding std=0.8、lm_head std=0.001、矩阵 uniform(±√3·d^−1/2)、输出投影全零。训练脚本走 build_model_metato_emptyinit_weights 三步(scripts/base_train.py:129-147)。
  • vocab 补到 64 的倍数(nanochat/gpt.py:168-172),DDP 与张量核都更顺;forward 里再切掉(nanochat/gpt.py:512-513)。
  • rotary 的 cos/sin 预计算 10 倍序列长、注册为 persistent=False 的 buffer(不进 checkpoint,nanochat/gpt.py:197-201)。

3. --depth 单旋钮

直觉

muP / scaling-law 的思路:在一个小参考模型(d12)上把超参调好,然后用公式外推到大模型,而不是每个尺寸重新调。nanochat 把这套外推写死在 base_train.py 里,用户只选「多大的模型」。

外推链(四步)

  1. 形状model_dim = depth × 64(aspect ratio),再向上对齐到 head_dim(128) 的倍数;头数 = model_dim / head_dim(scripts/base_train.py:129-140)。
  2. 训练 token 数target_tokens = 12 × (transformer_matrices + lm_head 参数量)scripts/base_train.py:263-269)。注释解释了为什么用这组参数计数:「transformer matrices + lm_head 给出最干净的 scaling laws」。
  3. 批大小:以 d12 为锚(D_REF 个 token、B_REF = 2^19 tokens),按 Power Lines 论文的 B ∝ D^0.383 外推,再 clamp 到最近的 2 的幂(scripts/base_train.py:271-284)。
  4. lr 与 wd:lr 按 η ∝ √(B/B_ref) 缩放(scripts/base_train.py:287-293);weight decay 按 T_epoch 恒定推出 λ = λ_ref·√(B/B_ref)·(D_ref/D)scripts/base_train.py:296-302)。注释坦白:这些论文研究的是 AdamW,对 Muon 是「盲目跟随并希望大约成立」。

结果:敲 --depth=24,宽度、头数、步数、批大小、lr、wd 全部自动。README 的 miniseries 就是这么扫出来的;GPT-2 能力大约在 d24–d26。

训练循环本体

  • 梯度累积:grad_accum_steps = total_batch_size / (device_batch × seq_len × world_size)scripts/base_train.py:405-410)。
  • LR 调度:线性 warmup → 恒定 → 线性 warmdown 到 5%(get_lr_multiplier scripts/base_train.py:360-369);Muon momentum 400 步内 0.85→0.97,warmdown 段再降回 0.90(scripts/base_train.py:372-382);weight decay 余弦到 0(scripts/base_train.py:385-386)。
  • 周期性评估:val bpb(nanochat/loss_eval.py:9 evaluate_bpb,见 §6)+ CORE 分数(scripts/base_eval.py:57 evaluate_core),评测时临时关掉 FP8(见 §7)。
  • MFU 靠硬编码的峰值 FLOPS 表换算(nanochat/common.py:228-279 get_peak_flops),日志里直接打 bf16_mfu
  • 一个工程细节:训练循环里手动 gc.freeze() + gc.disable()——Python GC 会周期性花约 500ms 扫循环引用(scripts/base_train.py:585-592)。

4. MuonAdamW:矩阵用 Muon,其余用 AdamW

直觉

  • Muon(MomentUm Orthogonalized by Newton-Schulz):对 2D 矩阵参数先做 SGD-momentum,再把更新正交化(替换成最近的正交矩阵)——modded-nanogpt speedrun 的核心武器,比 AdamW 收敛快。
  • 但 embedding、lm_head、标量不适合正交化,继续用 AdamW
  • 分布式:不用 DDP 包模型。优化器在 step() 内部自己同步梯度,且优化器状态按 rank 分片(ZeRO-2 风格)——每个 rank 只存 1/N 的动量。

三阶段异步通信

MuonAdamW.stepnanochat/optim.py:428-459)的结构,类 docstring 写得很清楚(nanochat/optim.py:214-227):

Phase 1 对所有 param group 发异步 reduce(reduce_scatter / all_reduce)
Phase 2 逐 group:等它的 reduce → 算更新 → 发异步 all_gather
Phase 3 等所有 gather 完成 → Muon 把堆叠参数拷回原参数

读法:Phase 1 的通信和 Phase 2 的计算互相重叠——先发的 reduce 先等,边等边算,算完立刻发 gather。

两类参数的通信路径

AdamW 大参数AdamW 小参数(<1024 元素)Muon 组
梯度同步reduce_scatter(每 rank 拿 1/N)all_reduce同形状堆成 (K,m,n) 后 reduce_scatter 按 chunk 切
状态存储只存自己切片的 exp_avg/exp_avg_sq全量复制(量小)只存自己拥有的参数的动量
更新后all_gather 拼回无需all_gather 拼回
位置nanochat/optim.py:274-360同左nanochat/optim.py:295-417

Muon 组还有两个省显存的招:K 不整除时补零再忽略(nanochat/optim.py:302-312);reduce_scatter 的输入 buffer 之后直接复用为 all_gather 的输出 buffer(nanochat/optim.py:414-416)。单 rank 时所有通信路径自然退化——同一个类覆盖单机与分布式(nanochat/optim.py:430-436)。

融合核里的现代配方

muon_step_fusednanochat/optim.py:111-180)在一个 torch.compile 全图融合核里装了五件事:

  1. Nesterov momentumg = grads.lerp(momentum_buffer, momentum)
  2. MuonEq 行平衡:先把每行缩放到平均行范数,让进正交化的谱更规整。
  3. Polar Express 迭代:替代经典 Newton-Schulz 的正交化迭代(系数表 polar_express_coeffs nanochat/optim.py:102-108,出自 arXiv 2505.16932),bf16 里跑;高矩阵与宽矩阵两种迭代式(nanochat/optim.py:145-154)。
  4. Muon+ 归一化:把 Frobenius 范数钉到 √(min(m,n)),修正欠收敛(nanochat/optim.py:158-161)。
  5. NorMuon 方差缩减 + cautious wd:按列/行做自适应缩放(arXiv 2510.05491);weight decay 只在更新与参数同号时施加(mask = (g * stacked_params) >= 0nanochat/optim.py:176-180)。

AdamW 侧同样是融合核(adamw_step_fused nanochat/optim.py:23-63):bf16 存储的参数(wte / value_embeds)在核内升 fp32 算完再写回;超参用 0-D CPU tensor 传入,改值不触发重编译。

参数分组

GPT.setup_optimizernanochat/gpt.py:419-456):

  • 六组 AdamW:lm_head / wte / value_embeds / resid_lambdas / x0_lambdas / smear+backout——各组有自己的 lr、betas、weight decay。
  • Muon 按形状分组(同形状才能堆叠通信)。
  • AdamW 各组 lr 还按 (d/768)^-0.5 随宽度缩放(muP 味道,nanochat/gpt.py:433-435)。

5. 数据装载:BOS 对齐 best-fit packing

直觉

经典做法把文档首尾相接再切成定长行——简单但「混乱」:一行里可能有半个文档,模型被迫在缺失开头的情况下预测。nanochat 选另一条路:每行都从 BOS 开始,文档整个整个往里装。

算法

tokenizing_distributed_data_loader_with_state_bos_bestfitnanochat/dataloader.py:74-161),每行容量 T+1:

  1. 从缓冲区找能整个放下的最大文档放入;重复直到什么都放不下。
  2. 剩余空当取缓冲区里最短的文档,裁剪到刚好填满。
  3. 性质:100% 利用(无 padding)、每行必有完整文档开头;代价是约 35% 的 token 被裁掉(文件头注释 nanochat/dataloader.py:1-17)。

示意:

# 示意,非源码
for row in batch:
while row_not_full:
doc = largest_doc_that_fits(remaining) # 第一步:能放下的最大文档
if doc is None:
doc = shortest_doc_in_buffer() # 第二步:裁最短的填满
doc = doc[:remaining]
row.append(doc)

分布式与性能细节

  • DDP 分片:rank r 读第 r, r+N, … 个 row group(_document_batches nanochat/dataloader.py:25-71);train/val 按「最后一个 parquet 文件」切(nanochat/dataloader.py:38)。
  • 状态可恢复:每次 yield 带 (pq_idx, rg_idx, epoch),存进 checkpoint 的 dataloader_state_dict,恢复时跳过已读(nanochat/dataloader.py:52-60scripts/base_train.py:340-347)。
  • 零拷贝管线:预分配 row_buffer → pinned CPU buffer → 单次 HtoD 拷进常驻 GPU buffer(nanochat/dataloader.py:111-161)。
  • 预取藏在训练循环里:算完当前 micro-batch 的 backward 前,next(train_loader) 已经开始准备下一批(scripts/base_train.py:508-515)。

6. 评测指标:bpb 与 CORE

  • bpb(bits per byte):把 loss 按「目标 token 的字节数」归一,词表大小变了也能横向比(evaluate_bpb nanochat/loss_eval.py:9-65);特殊 token 与 ignore_index 位置既不进分子也不进分母,分母表来自第一章的 token_bytes.pt
  • CORE:DCLM 论文的基座模型综合分(nanochat/core_eval.py),多选/schema/语言建模三类题型统一渲染成「前缀 + 候选续写」比 loss;speedrun 排行榜的登顶线就是 GPT-2 的 CORE = 0.256525。

7. FP8(H100+)

  • --fp8 把够大的 Linear(维度 16 的倍数且 min(din,dout) ≥ 128)换成 Float8Linearscripts/base_train.py:174-185fp8_module_filter)。
  • 实现只有约 150 行:一个 autograd.Function_Float8Matmul nanochat/fp8.py:125)内部做 tensorwise 动态量化 → torch._scaled_mm → 反量化;输入/权重用 e4m3,梯度用 e5m2(nanochat/fp8.py 文件头)。
  • 文件头还诚实对比了 torchao 方案(tensor subclass + dispatch 表)与自家方案(单个不透明 autograd.Function):GPU matmul 完全相同,差别只在编译器眼里的「胶水 op」。
  • 评测时怕精度漂移,disable_fp8 上下文管理器把 Float8Linear 临时换回普通 Linear(共享 weight、meta device 防显存峰值,scripts/base_train.py:193-237)。

8. 关键细节 / 坑

  • COMPUTE_DTYPE 替代 autocast:全局一个 dtype(SM80+ 自动 bf16;可用 NANOCHAT_DTYPE 覆盖),权重 fp32 存、前向时自定义 Linear cast(nanochat/common.py:13-32nanochat/gpt.py:45-49)。fp16 路径自动开 GradScaler(scripts/base_train.py:323-326)。
  • FA3 是按硬件拉 kernelflash_attention.py 用 HF kernels 包按 SM 版本取对应 FA3(sm90 用 varunneal 版);拿不到就 SDPA 兜底,但兜底不支持滑窗,训练会刷屏警告(nanochat/flash_attention.py:23-50scripts/base_train.py:107-117)。
  • resume 恢复到数据加载器级别--resume-from-step 连同 optimizer state 与 dataloader_state_dict 一起恢复(scripts/base_train.py:150-157)。
  • checkpoint 目录按阶段分base_checkpoints/ chatsft_checkpoints/ chatrl_checkpoints/load_model(source, ...) 一个参数切换(nanochat/checkpoint_manager.py:163-171);老 checkpoint 缺新字段时 _patch_missing_* 自动补默认值(nanochat/checkpoint_manager.py:22-39)。
  • 跑之前先看 runs/speedrun.sh——它是「参考配置」的活文档。

下一章:03 · 后训练链