数据截至 (上游 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 | 位置 |
|---|---|---|---|
| 位置编码 | 学习式绝对位置 embedding | rotary,无位置 embedding | apply_rotary_emb nanochat/gpt.py:57-64 |
| 注意力稳定 | 无 | QK norm + q,k 各乘 1.2 | nanochat/gpt.py:101-103 |
| 头结构 | MHA | GQA(n_kv_head ≤ n_head) | nanochat/gpt.py:72-78 |
| MLP 激活 | GELU | relu² | MLP.forward nanochat/gpt.py:137-141 |
| embedding 绑定 | tied | untied | nanochat/gpt.py:174-177 |
| norm | LayerNorm(带参数) | 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(GPTConfignanochat/gpt.py:36-39)。 _compute_window_sizes(nanochat/gpt.py:287-317)把字符串翻译成 FA3 的(left, right)窗口参数;S 层窗口向上取整到 128 的倍数(FA3 tile 对齐,2048 → 768)。- 收益直接体现在 FLOPs 估计里:
estimate_flops(nanochat/gpt.py:319-339)按每层窗口分别算注意力开销——S 层越多,训练越省。
2.2 Value Embeddings(ResFormer 式)
- 交替层(且最后一层必含)多一张
value_embeds嵌入表:把 token 嵌入当作「先验 value」,以门控方式加进注意力的 v(has_venanochat/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_weightsnanochat/gpt.py:238-243)。
初始化与元设备
GPT.__init__跑在 meta device 上下文:只算形状与 dtype,不填数据;真正的初始化集中在init_weights(nanochat/gpt.py:204-284)——embedding std=0.8、lm_head std=0.001、矩阵 uniform(±√3·d^−1/2)、输出投影全零。训练脚本走build_model_meta→to_empty→init_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 里,用户只选「多大的模型」。
外推链(四步)
- 形状:
model_dim = depth × 64(aspect ratio),再向上对齐到 head_dim(128) 的倍数;头数 = model_dim / head_dim(scripts/base_train.py:129-140)。 - 训练 token 数:
target_tokens = 12 × (transformer_matrices + lm_head 参数量)(scripts/base_train.py:263-269)。注释解释了为什么用这组参数计数:「transformer matrices + lm_head 给出最干净的 scaling laws」。 - 批大小:以 d12 为锚(
D_REF个 token、B_REF = 2^19tokens),按 Power Lines 论文的B ∝ D^0.383外推,再 clamp 到最近的 2 的幂(scripts/base_train.py:271-284)。 - 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_multiplierscripts/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:9evaluate_bpb,见 §6)+ CORE 分数(scripts/base_eval.py:57evaluate_core),评测时临时关掉 FP8(见 §7)。 - MFU 靠硬编码的峰值 FLOPS 表换算(
nanochat/common.py:228-279get_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.step(nanochat/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_fused(nanochat/optim.py:111-180)在一个 torch.compile 全图融合核里装了五件事:
- Nesterov momentum:
g = grads.lerp(momentum_buffer, momentum)。 - MuonEq 行平衡:先把每行缩放到平均行范数,让进正交化的谱更规整。
- Polar Express 迭代:替代经典 Newton-Schulz 的正交化迭代(系数表
polar_express_coeffsnanochat/optim.py:102-108,出自 arXiv 2505.16932),bf16 里跑;高矩阵与宽矩阵两种迭代式(nanochat/optim.py:145-154)。 - Muon+ 归一化:把 Frobenius 范数钉到 √(min(m,n)),修正欠收敛(
nanochat/optim.py:158-161)。 - NorMuon 方差缩减 + cautious wd:按列/行做自适应缩放(arXiv 2510.05491);weight decay 只在更新与参数同号时施加(
mask = (g * stacked_params) >= 0,nanochat/optim.py:176-180)。
AdamW 侧同样是融合核(adamw_step_fused nanochat/optim.py:23-63):bf16 存储的参数(wte / value_embeds)在核内升 fp32 算完再写回;超参用 0-D CPU tensor 传入,改值不触发重编译。
参数分组
GPT.setup_optimizer(nanochat/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_bestfit(nanochat/dataloader.py:74-161),每行容量 T+1:
- 从缓冲区找能整个放下的最大文档放入;重复直到什么都放不下。
- 剩余空当取缓冲区里最短的文档,裁剪到刚好填满。
- 性质: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_batchesnanochat/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-60、scripts/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_bpbnanochat/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)换成Float8Linear(scripts/base_train.py:174-185的fp8_module_filter)。- 实现只有约 150 行:一个
autograd.Function(_Float8Matmulnanochat/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 存、前向时自定义Linearcast(nanochat/common.py:13-32、nanochat/gpt.py:45-49)。fp16 路径自动开 GradScaler(scripts/base_train.py:323-326)。 - FA3 是按硬件拉 kernel:
flash_attention.py用 HFkernels包按 SM 版本取对应 FA3(sm90 用 varunneal 版);拿不到就 SDPA 兜底,但兜底不支持滑窗,训练会刷屏警告(nanochat/flash_attention.py:23-50、scripts/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 · 后训练链。