跳到主要内容

数据截至 (上游 commit 3adf61e154c3)

02 · 训练循环:300 行里的完整预训练

这一章讲什么: train.py 的 337 行如何覆盖「从零训一个 GPT-2」的全部环节——配置、分布式、混合精度、梯度累积、学习率调度、评估存档。它是一张「完整训练循环零件清单」的最小参考实现。


1. 它要解决的小问题

「训练 = 前向、反传、更新」这句话漏掉了真实训练里的全部杂事:

  • 超参有 40 个,怎么管理才不烦?
  • 8 张卡怎么一起训?
  • 显存装不下大 batch 怎么办?
  • 混合精度训练怎么不丢精度?
  • 学习率什么时候开始降、降到多少?
  • 训到一半断了怎么办?
  • 怎么知道 GPU 有没有被喂饱?

train.py 每问给一个 5~20 行的答案。本章按执行顺序走一遍。


2. 执行顺序总览

脚本从上到下就是初始化流水线,最后进入一个 while True 主循环(train.py:255):

① 默认超参写成全局变量 (train.py:32-74)
② exec configurator.py → 配置文件/命令行覆写 (train.py:77)
③ DDP 初始化: rank/world_size/设备 (train.py:82-100)
④ TF32 + autocast 上下文 (train.py:107-112)
⑤ get_batch 定义 (train.py:116-131)
⑥ 建模型: scratch / resume / gpt2 三选一 (train.py:149-188)
⑦ GradScaler + AdamW + torch.compile + DDP 包装 (train.py:196-212)
⑧ while True: 设 LR → 定期评估存档 → N 个 micro-step → clip → step

3. 机制一:穷人配置系统(configurator.py)

小问题: 不想引入 Hydra/argparse 这类配置框架,又要支持「配置文件 + 命令行覆盖」。

nanoGPT 的野路子: 超参就是脚本顶部的全局变量;train.py:77exec(open('configurator.py').read()) 把覆盖逻辑直接执行在当前全局命名空间里。

configurator.py 对命令行参数分两类处理(configurator.py:20-47):

  • 不含 = 的参数 → 当成配置文件路径,整个文件 exec 进来configurator.py:28)。所以 config/train_gpt2.py 就是一个普通 Python 文件,里面直接写 batch_size = 12
  • --key=valueliteral_eval 尝试解析成 Python 字面量,再断言类型与原全局变量一致后覆写(configurator.py:31-45)。

坑: 类型断言 type(attempt) == type(globals()[key])configurator.py:42)意味着 --dropout=0 会被拒(int ≠ float,默认是 0.0)——命令行要写 --dropout=0.0。作者自己在文件头承认这是「Probably a terrible idea」,但省了所有配置样板代码(configurator.py:1-15)。


4. 机制二:DDP——8 张卡跑同一份代码

小问题: 多卡训练的经典方案是 PyTorch DDP:每张卡一个进程、各持完整模型副本、各自前向反传,反传时梯度跨进程求平均。

nanoGPT 里的最小接入,只有三处:

  1. 检测与初始化train.py:82-95):靠环境变量 RANK 判断是否 DDP 运行(torchrun 会注入),拿到 ddp_rank / ddp_local_rank / ddp_world_size,并把 gradient_accumulation_steps 按 world_size 整除摊到各进程。
  2. 包装模型train.py:211-212):DDP(model, device_ids=[ddp_local_rank])
  3. 主循环里控制梯度同步:见下一节。

只有 master_process(rank 0)负责评估打印和写 checkpoint(train.py:90train.py:263)。

细节: seed_offset = ddp_ranktrain.py:91)配合 torch.manual_seed(1337 + seed_offset)train.py:106),让每个进程采到不同的随机 batch——不然 8 张卡全喂一样的数据就白并行了。


5. 机制三:梯度累积——小显存模拟大 batch

小问题: 复现 GPT-2 要 ~0.5M token 的 batch(config/train_gpt2.py 注释:12×1024×5×8 = 491,520),但一张 40GB 卡一次只放得下 12×1024。办法是把一步拆成 gradient_accumulation_steps 个 micro-step,梯度累加后再更新。

真实实现(主循环核心,train.py:292-305):

for micro_step in range(gradient_accumulation_steps):
if ddp:
model.require_backward_grad_sync = (micro_step == gradient_accumulation_steps - 1)
with ctx:
logits, loss = model(X, Y)
loss = loss / gradient_accumulation_steps
X, Y = get_batch('train') # 前向时异步预取下一批
scaler.scale(loss).backward()

三个点拆开看:

  • require_backward_grad_synctrain.py:298):DDP 默认每次 backward 都跨卡同步梯度,累积时这是纯浪费——前 N-1 个 micro-step 直接关掉同步,只在最后一次开。注释里作者明说:官方做法是 model.no_sync() 上下文管理器,但他嫌它逼你写两遍前向代码,翻源码发现它只是切换这个布尔值,于是直接手切(train.py:293-297)。这是全库「剥抽象」精神最典型的例子。
  • loss 先除以累积步数train.py:301):因为 loss 本身是 micro-batch 上的平均,除一下让累加后的梯度等价于大 batch 平均。
  • 预取藏时延train.py:303):get_batch 放在前向之后、backward 之前调用,CPU 取数 + pin_memory().to(device, non_blocking=True) 的传输和 GPU 上的前向计算重叠。

micro-step 循环结束后才是更新三连:可选梯度裁剪(train.py:307-309)、scaler.step(optimizer) + scaler.update()train.py:311-312)、optimizer.zero_grad(set_to_none=True)train.py:314)。


6. 机制四:混合精度——autocast + GradScaler

小问题: 用 16 位浮点训练快一倍、省一半显存,但 float16 动态范围太小,梯度会下溢成 0。

nanoGPT 的分工(train.py:110-112train.py:196):

  • autocast 上下文ctx = torch.amp.autocast(device_type=..., dtype=ptdtype),前向时自动把合适的算子降到 16 位。CPU 上则退化为 nullcontext()
  • dtype 默认选 bfloat16(硬件支持的话,train.py:73):bf16 动态范围和 float32 一样,不需要 GradScaler
  • GradScaler 只在 float16 时启用GradScaler(enabled=(dtype == 'float16'))train.py:196)。enabled=Falsescaler.scale/step/update 全是 no-op——所以主循环里那一串 scaler.* 调用在 bf16 下是零成本直通。

train.py:73 的这一行还内嵌了一个硬件探测:torch.cuda.is_bf16_supported() 不行才退回 float16。


7. 机制五:学习率调度——warmup + 余弦衰减

小问题: 训练初期参数是随机的,大学习率会直接打飞;后期又要小学习率细调。

get_lrtrain.py:231-242)三段式:

  1. 线性 warmup:前 warmup_iters 步从 0 线性升到 learning_ratetrain.py:233-234)。
  2. 余弦衰减:之后按 0.5*(1+cos(π·ratio)) 从峰值降到 min_lrtrain.py:239-242)。
  3. 过点保底:超过 lr_decay_iters 后恒定 min_lrtrain.py:236-237)。

每步开头把 get_lr(iter_num) 写进所有 param_group(train.py:258-260)。默认超参直接按 Chinchilla 口径给:lr_decay_iters ≈ max_itersmin_lr ≈ learning_rate/10train.py:67-68 的注释)。


8. 机制六:优化器分组与 fused AdamW

小问题: weight decay 不该一视同仁——LayerNorm 的 gain 和 bias 做衰减会伤表现,矩阵和嵌入才该衰减。

configure_optimizersmodel.py:263-287)的一刀切规则:按参数维度分。 p.dim() >= 2 进衰减组,否则进零衰减组(model.py:270-271):

decay_params = [p for n, p in param_dict.items() if p.dim() >= 2]
nodecay_params = [p for n, p in param_dict.items() if p.dim() < 2]

不用记任何参数名,bias/LayerNorm 天然是 1D、矩阵天然是 2D,规则自动落对。

另外两处讲究:先过滤 requires_gradmodel.py:267);有 CUDA 且 PyTorch 支持时用 fused=True 的 AdamW(model.py:281-284),把逐参数更新融成一个 kernel。


9. 机制七:评估、checkpoint、断点续训

评估(estimate_losstrain.py:215-228): 对 train/val 各抽 eval_iters 个 batch 算平均 loss,model.eval() 进、model.train() 出。每 eval_interval 步由 master 进程调用(train.py:263-265)。

存档(train.py:274-286): val loss 改进时(或 always_save_checkpoint 时)把五样东西打进 ckpt.pt

内容
modelraw_model.state_dict()(注意是解开 DDP/compile 后的裸模型,train.py:253
optimizer优化器状态(动量等,续训必需)
model_args建模用的六个结构参数
iter_num / best_val_loss训练进度
config全部超参快照

续训(train.py:158-180): init_from='resume' 时从 ckpt.pt 恢复:六个结构参数(n_layer/n_head/n_embd/block_size/bias/vocab_size)强制以 checkpoint 为准(train.py:166-167),优化器状态也读回(train.py:200-201),从原 iter_num 接着跑。

坑: torch.compile 会给 state_dict 的 key 加 _orig_mod. 前缀,读档时要先剥掉(train.py:174-177;同样逻辑在 sample.py:42-45)。作者在原注释里诚实写道「honestly no idea how checkpoints sometimes get this prefix」。


10. 机制八:MFU——GPU 有没有被喂饱

MFU(model FLOPs utilization,模型算力利用率) 回答「我实测的 FLOPS 占硬件峰值的百分之几」。estimate_mfumodel.py:289-303)的做法:

  • 每 token 的 FLOPs 用 PaLM 论文附录 B 的口径:6N + 12·L·H·Q·Tmodel.py:296)——6N 是矩阵乘的前向+反传,12LHQT 是注意力部分。
  • 乘以每步 token 数、除以每步耗时,得实测 FLOPS,再除以 A100 bf16 峰值 312e12model.py:301)。

主循环用滑动平均滚动报告:running_mfu = 0.9*running + 0.1*mfutrain.py:325-326)。这套估计换成别的卡要改 312e12 这个硬编码峰值——它只是一个量级仪表,不是精确测量。


11. 本章代码地图

主题文件符号
默认超参train.py:32-74顶层全局变量
配置覆盖configurator.py:20-47顶层 for arg in sys.argv[1:]
DDP 初始化train.py:82-100init_process_groupmaster_process
混合精度上下文train.py:110-112ctxtorch.amp.autocast
取 batchtrain.py:116-131get_batch(下一章细讲)
三种初始化train.py:149-188init_from 分支
_orig_mod. 前缀train.py:174-177unwanted_prefix 循环
GradScalertrain.py:196torch.cuda.amp.GradScaler
优化器model.py:263-287configure_optimizers
学习率train.py:231-242get_lr
评估train.py:215-228estimate_loss
主循环train.py:255-333顶层 while True
梯度同步开关train.py:298model.require_backward_grad_sync
更新三连train.py:307-314clip_grad_norm_scaler.stepzero_grad
MFUmodel.py:289-303estimate_mfu