把 Megatron 的
pretrain()读完,你会撞上一个反直觉的事实:这份负责「预训练」的文件,对模型一无所知。模型怎么构建(model_provider)、一步前向怎么算(forward_step_func)、数据长什么样(dataset_provider),全都是参数传进来的函数。它不生产任何一个 FLOP——它是纯粹的控制面。
有读者问过我一个好问题:拿到一份 Megatron 的 training.py(就是定义 pretrain() 和 train() 的那个文件),能不能靠它彻底搞懂预训练?
诚实的回答是:搞懂一半,但恰好是大多数资料不讲的那一半。预训练可以粗暴地切成数据面和控制面。数据面是「一个 batch 进去、一个梯度出来」的数学与通信——前向反向、张量并行、流水线调度、优化器更新,这些都不在这个文件里。控制面是「几千张卡连续跑几十天不出错、出了错能恢复、每一秒都被记账」的工程——这个文件就是控制面的全部。教科书和论文讲数据面,而控制面几乎只存在于源码里。
这是本站 Megatron 源码系列的第二篇。上一篇《模型配置到底在配置什么》拆的是模型如何被造出来(Config/Spec/Builder 三层);这一篇拆造出来之后如何被训练。如果你读过 nanoGPT 的训练配方,可以把 nanoGPT 那个几百行的训练脚本当作本文的「理想气体」参照物:两者的骨架完全同构,差值就是工业级预训练的全部工程成本清单。
先交代素材:本文基于一份 Megatron-LM 训练侧的真实 training.py(NVIDIA 2025 版权头,路径 megatron/training/training.py),这个版本带若干实验性特性(线性注意力变体、RL 步、GTP 重物化等),与你手上的公开版本可能不完全一致,但主循环骨架是稳定的。文中所有代码片段来自这份源码或标注为伪代码;所有手算数字都先用一次性脚本验算过(脚本在文末)。
名词速查
| 术语 | 一句话解释 |
|---|---|
| DP / TP / PP / CP | 四种并行轴:数据并行(复制模型分数据)、张量并行(切开单个矩阵乘)、流水线并行(按层切)、上下文并行(按序列长度切) |
| micro-batch / global batch | 一次前反向吃的小批 vs 一次参数更新对应的总样本数;global = micro × 梯度累积次数 × DP 数 |
| grad buffer / bucket | DDP 把所有梯度放进一块连续大缓冲区,再按「桶」为单位分批做通信,让通信和反向计算重叠 |
| distributed optimizer | 把优化器状态(动量等)按 DP 维度分片,每张卡只存 1/N,思路与 ZeRO 第一阶段同源 |
| loss scaling | fp16 训练时把 loss 乘一个大数再反传,防梯度下溢;溢出时该步作废、缩小系数重来 |
| forward pre-hook | 挂在前向计算之前的钩子函数,Megatron 用它把「取回分片参数的 all-gather」藏进前向 |
| CUDA Graph | 把一串 GPU kernel 的启动序列录下来整体重放,消除 CPU 逐个发射 kernel 的开销 |
| BSHD / THD | 两种 batch 布局:定长矩形(batch×seq)vs 把多条变长序列首尾相接打包成一条(sequence packing) |
| SDC | Silent Data Corruption,硬件静默出错——不崩溃、不报错,只是算出来的数悄悄不对 |
| GRPO | 一种 RL 后训练算法(PPO 家族),见《PPO 训练循环解剖》 |
主线:二十行骨架,四本账
先给这篇文章的一句话主线(这是我的提炼,不是源码里的官方说法):这份文件 = 一个二十行就能写完的理想训练循环,外面包了四本工程账。 读它的正确姿势不是逐行读,而是先把骨架认领出来,再把剩下的每一段代码归入四本账之一:
- 记账本(observability):每一秒、每一个 FLOP、每一焦耳都要有出处——吞吐、进度日志、能耗、计时器。
- 容错本(fault tolerance):几千张卡跑几十天,硬件故障不是异常而是常态——checkpoint、重跑机、哈希对账、故障注入。
- 性能本(performance):凡是能藏起来的开销都要藏——通信与计算重叠、CUDA Graph、跨 rank 对齐的 GC。
- 胶水本(parallelism plumbing):让同一个循环在任意并行拓扑下都成立——进程组集合、DDP 包装、跨模型并行组的决策对齐。
骨架本身长这样:
# 伪代码:train() 的理想骨架(对应 megatron/training/training.py 的 train() 主干)
# 省略:容错、性能开关、RL 支线、日志细节——即本文其余全部内容
while iteration < train_iters:
update_num_microbatches(consumed_samples) # 批大小调度(可爬坡)
loss = train_step(data_iterator, model, optimizer) # 前向+反向+参数更新
iteration += 1
consumed_samples += global_batch_size # 进度以样本数记
training_log(loss, ...) # 记账
if iteration % eval_interval == 0:
evaluate(valid_iterator) # 验证
if should_save(iteration):
save_checkpoint(...) # 存档
if should_exit(iteration):
break # 退出(先存档)
这个骨架和 nanoGPT、和任何 PyTorch 教程里的训练循环没有本质区别。接下来整篇文章回答一个问题:从这二十行到这份数千行的文件,中间发生了什么?
开场:一个在 import 之前就开始的程序
文件的第一行可执行代码不是函数定义,是抓时间戳:
import time
_TRAIN_START_TIME = time.time() # The earliest we can measure the start time.
# ……然后才是几百行 import
为什么偏执到把计时器放在 import 之前?因为在这个体量的代码库里,import 本身是要花几十秒的(torch、Transformer Engine、各种 CUDA 扩展的加载),而「作业从调度器拉起到真正开始训练花了多久」是集群利用率的一部分,必须被记账。这是记账本的第一笔:程序还没开始运行,账已经开始记了。
更进一步,pretrain() 启动后会把每个 rank 各自的启动时间戳做一次全局 all-reduce 取最小值:
program_start_global = torch.tensor([_STARTUP_TIMESTAMPS['program_start']], ...)
torch.distributed.all_reduce(program_start_global, op=torch.distributed.ReduceOp.MIN)
几千个进程不会同时被拉起,「训练开始时间」取决于最早的那个——这才是调度器视角的真实成本。随后一组 startup-* 计时器把启动过程切成七八段(库加载、进程内设置、megatron 初始化、JIT 融合配置……)分别记账。一个隐含的判断是:在万卡集群上,启动慢十分钟 = 万卡小时级别的浪费,值得为它写几十行监控代码。
pretrain() 的整体生命周期是线性的五步,没有惊喜:初始化 Megatron(进程组、随机种子)→ 构建模型+优化器+学习率调度器 → 构建数据迭代器 → train() 主循环 → 收尾(补存 checkpoint、跑验证/测试集)。真正的内容都藏在每一步的细节里。
train():主循环的一次迭代到底在干什么
下图是主循环单次迭代的控制流(简化,省略 RL 支线和 CUDA Graph 捕获):
flowchart TD
A["迭代开始"] --> B["校准微批数<br>update_num_microbatches"]
B --> C{"微批数变了?"}
C -- "是" --> D["先存 checkpoint"]
D --> E{"在跳过名单里?"}
C -- "否" --> E
E -- "是" --> F["dummy_train_step<br>只消费数据不训练"]
F --> A
E -- "否" --> G["train_step<br>前向 + 反向 + optimizer.step"]
G --> H{"首个迭代且成功?"}
H -- "是" --> I["启用 forward pre-hook<br>打开通信重叠"]
H -- "否" --> J["FLOPs / 样本记账 + 日志"]
I --> J
J --> K{"到 eval 间隔?"}
K -- "是" --> L["evaluate<br>先关重叠再评估"]
K -- "否" --> M["post-step 回调<br>心跳 / GC / 哈希对账"]
L --> M
M --> N{"存档或退出?"}
N -- "退出" --> O["先存档再退出"]
N -- "继续" --> A
几个分支值得单独讲。
批大小是变量,不是常量
教程里 batch size 是个常量;生产里它是个调度出来的函数。Megatron 支持 batch size 爬坡(rampup):训练早期用小 batch,随训练推进逐步放大。所以每次迭代的第一件事是 update_num_microbatches(consumed_train_samples)——根据已消费的样本数算出当前应该用多少个 micro-batch。
由此带来两个连锁设计。第一,微批数一变,先存 checkpoint:源码在检测到 get_num_microbatches() != num_microbatches 时立刻保存。我的理解是(推断,源码未写明动机):恢复训练时批大小调度必须严格复现,在调度边界存档能保证断点前后的 global batch size 一致,学习率曲线不会错位。第二,源码 assert 微批数只增不减——爬坡是单向的。
同样的逻辑延伸到学习率:opt_param_scheduler.step(increment=samples_this_iteration)——学习率调度器走的是样本数轴,不是迭代数轴。batch size 会变,「第 1000 步」没有稳定含义,「第 200 万个样本」才有。
跳过一个迭代 ≠ 什么都不做
args.iterations_to_skip 允许跳过指定迭代(比如某段数据有问题)。但看 dummy_train_step() 的实现会发现它一点也不 dummy:它照常从数据迭代器取数、照常走 TP 广播、照常做 CP 切分——只是不进模型。
为什么?因为数据迭代器的位置是全局状态的一部分。几千个 rank 各自持有迭代器,跳过迭代如果不消费数据,恢复训练或对比实验时数据顺序就会错位一个 batch,且没有任何报错。「跳过」必须精确地只跳过计算,保留一切副作用。这个函数是理解「分布式训练里没有无害操作」的绝佳教材。
第一个迭代的三重戒备
主循环对第一个迭代有特殊处理,注释写得很直白:
# Disable forward pre-hook to start training to ensure that errors in checkpoint
# loading or random initialization don't propagate to all ranks in first
# all-gather (which is a no-op if things work correctly).
展开说:开启通信重叠后,参数 all-gather 挂在 forward pre-hook 里异步进行。如果 checkpoint 加载有问题(某个 rank 的参数是坏的),异步 all-gather 会把坏参数静默扩散到所有 DP 副本,错误现场被彻底污染。所以第一个迭代同步跑、关掉重叠,确认无恙后再开钩子。
另外两重戒备:fp16 训练下前若干个迭代可能因 loss scale 未稳定而溢出跳步,钩子的启用推迟到第一个成功的迭代之后;开启参数哈希对账(后述)时,训练前先做一次全量校验。三件事共享同一个思想:分布式系统的第一步要用最保守的方式走,先建立「当前状态是好的」这个基线,再开始加速。
四条退出路径,全部先存档
checkpoint_and_decide_exit() 汇聚了四种退出条件:收到 SIGTERM 信号、运行时长到限(exit_duration_in_mins)、到达指定迭代(exit_interval)、到达训练阶段切换点(phase_transition_iterations)。四条路径无一例外先存 checkpoint 再退出。
其中时长判断有个容易忽略的细节:
done_cuda = torch.tensor([train_time > args.exit_duration_in_mins], ...)
torch.distributed.all_reduce(done_cuda, op=torch.distributed.ReduceOp.MAX)
「要不要退出」这个决定要做一次 all-reduce(MAX:任何一个 rank 说到点了就全体退出)。因为各 rank 的时钟有偏差,如果各自判断,就会出现一部分 rank 退了、另一部分还在等集合通信的死锁。在分布式训练里,连「看表」都必须是集体行为。这个模式(决策 all-reduce)在文件里反复出现,后面还会遇到。
train_step():一步更新的事务边界
单步的主干值得用伪代码钉住:
# 伪代码:train_step() 主干(对应源码,省略:激活/梯度 dump、蒸馏、vision 分支)
while rerun_state_machine.should_run_forward_backward(data): # 可重跑事务
for chunk in model:
chunk.zero_grad_buffer()
optimizer.zero_grad()
losses = forward_backward_func(forward_step_func, data, model,
num_microbatches, ...) # 流水线调度器
update_ok, grad_norm, num_zeros = optimizer.step()
update_ok = logical_and_across_mp_group(update_ok) # 跨 MP 组对齐
if update_ok:
opt_param_scheduler.step(increment=samples_this_iter) # LR 按样本走
skipped = 0
else:
skipped = 1 # 本步作废
loss = per_token_average(losses, group=dp_cp) # 见下文
三个点展开。
前向反向被包在一个「可重跑事务」里
注意最外层不是顺序执行,是 while rerun_state_machine.should_run_forward_backward(...)。这个循环大多数时候只跑一遍,但它的存在改变了单步的语义:一次前向反向不是一段代码,而是一个可以被重放的事务。数据迭代器被 RerunDataIterator 包装(可以回带重放同一批数据),配合确定性计算,数值可疑时(比如 loss 出现 NaN)可以在同一批数据上重跑一遍:两次结果不同,说明是瞬态硬件错误(SDC);两次一样,才是真的软件 bug 或数据问题。
需要说明:重跑机的内部实现在 megatron/core/rerun_state_machine.py,本文没有展开读,以上是从接口用法和公开文档得出的理解。但接口本身已经传递了核心信息:在足够大的集群上,「算错了」和「写错了」是两种需要区分的故障,而区分它们的唯一方法是重放。
「这步成功了吗」也要投票
optimizer.step() 返回的 update_successful 不能直接用——要先跨模型并行组做一次逻辑与:
update_successful = logical_and_across_model_parallel_group(update_successful, group=mp_group)
grad_norm = reduce_max_stat_across_model_parallel_group(grad_norm, group=mp_group)
原因藏在注释里:模型被 TP/PP 切开后,有的 rank 可能没有可训练参数(比如冻结了部分子模型),它的「成功」是空洞真;而 fp16 溢出可能只发生在某一个分片上。参数更新必须全有或全无——一部分 rank 更新了、另一部分没更新,模型就永久性地内部不一致了。这是「看表要集体」之后的第二个决策 all-reduce。
loss 平均数的正确打开方式
汇报 loss 时源码里有个不起眼但值得学的分支:每个 micro-batch 返回的 loss 如果是二元组 (loss_sum, token_count),就走新路径——把所有 micro-batch 的二元组相加、跨 DP 组 all-reduce、最后一除:
val = torch.vstack(val).sum(dim=0)
torch.distributed.all_reduce(val, group=dp_cp_group)
loss_reduced[key] = val[0] / val[1] # 全局 per-token 平均
而 legacy 路径是「先对每个 micro-batch 求平均、再对平均值求平均」。两者在定长 batch 下等价,但一旦序列变长(SFT、sequence packing),每个 micro-batch 的有效 token 数不同,「平均的平均」就会给 token 少的 batch 更高的权重——这是个真实存在过的统计偏差。修法就是小学数学:分子分母分开传,最后再除。如果你在看训练曲线时遇到过换数据管线后 loss 突跳,这里可能就是原因之一。
记账本:FLOPs 是怎么算出来的
num_floating_point_operations() 是整个文件里最数学的部分,它回答仪表盘上那个 TFLOP/s/GPU 是怎么来的。这个数字重要,因为它是 MFU(Model FLOPs Utilization,实际吞吐占硬件峰值的比例)的分子——训练优化的一切努力最终都汇报到这个数上。
从矩阵乘数到 6ND
公式的骨架是三个系数的乘积,源码注释写得很清楚:
- ×2:一次
m×n · n×k矩阵乘是 2mnk 次浮点运算(乘加各算一次,FMA); - ×3:每个矩阵乘要做三遍——前向一遍,反向两遍(对权重的梯度 wgrad、对输入的梯度 dgrad);
- 剩下的就是逐模块数参数:QKV 投影、输出投影、MLP、logits 头。
2×3 = 6,这就是著名的「训练成本 ≈ 6ND」(N 个参数、D 个 token)估算的出处——Kaplan 等人的 Scaling Laws 论文(arXiv:2001.08361)用的就是这个近似。Megatron 的公式可以看作 6ND 的精确版:把「N 个参数」展开成逐层逐模块的精确计数,并补上 6ND 忽略的注意力矩阵项。
补多少?我用一个 LLaMA-7B 风格的稠密配置(hidden 4096、32 层、FFN 11008 + SwiGLU、词表 32000、序列长 4096、无 GQA)把源码公式复算了一遍(脚本见文末):
- 公式给出每 token 约 4.29 × 10¹⁰ FLOPs;
- 6N(N ≈ 6.61B 含 embedding)给出约 3.96 × 10¹⁰;
- 比值 1.081——精确公式比 6ND 高约 8%,其中注意力的 L² 项(QKᵀ 和加权求和 V)在序列长 4096 时占总量约 7.5%。
结论可以带走:在 4K 上下文的稠密模型上,6ND 低估不到一成;上下文越长,L² 项占比越大,6ND 越不准。顺带一提,源码在多头潜在注意力(MLA)分支的注释里直接引用了两篇 arXiv 论文(2305.10403 与 2205.05198)作为 FLOPs 推导依据——工程代码里带论文引用,这个习惯本身值得学。
sequence packing 的记账难题
真正精彩的是这个函数如何处理 THD 布局(sequence packing:把多条变长序列打包进一条定长序列,用 cu_seqlens 数组记录边界)。打包后有两笔账变了:
- token 线性项(投影、MLP、logits):应该按真实 token 数
sum(Lᵢ)记,padding 不算工作量; - 注意力 L² 项:打包后每条子序列只在自己内部做注意力(分块因果掩码),FLOPs 是
sum(Lᵢ²)而不是(总长)²。
这两个量在数学上互相独立(知道和推不出平方和),所以必须分别追踪。差距有多大?手算一个例子:把 [2048, 1024, 1024] 三条序列打包成一条 4096 的序列——
注意力的真实计算量只有按整条序列算的 37.5%。如果记账时不区分,MFU 会被虚报——你以为算了 16.7M 单位的注意力,实际只算了 6.3M。
实现上的克制同样值得学。这两个统计量放在一个 GPU 上的 2 元素 fp64 张量里逐 micro-batch 累加(无 host 同步),每个迭代结束时只做一次 2 元素的 all-reduce + 一次 host 同步取回。由于 cu_seqlens 在 TP/CP/PP 组内是复制的,全局 all-reduce 会多算 TP×CP×PP 倍,最后除掉。而不打包的 BSHD 路径从头到尾不碰这套机制——一个 _seqlen_stats_active 布尔门控保证它零成本。给记账代码本身记账(它的开销是多少),是记账本的最高修养。
不止 FLOPs:进度日志与能耗
记账本还有两页。一页是 progress.txt 进度日志:每次存 checkpoint 追加一行累计 FLOPs 和吞吐,重启后从日志里找到「同样 world size 的最早一次启动」来计算跨作业的累计吞吐——注意异步 checkpoint 要区分「Saving」和「Saved」两条记录,只有落盘确认的才算数(一个朴素的两阶段提交)。另一页是能耗:log_energy 打开后按 J/iter/GPU 和 W/GPU 记账。另外,checkpoint 保存期间计时器和能耗监控都会暂停——存档时间不算进训练吞吐,保证 MFU 数字诚实。
容错本:把硬件故障当业务逻辑
这本账的世界观可以用一句话概括:在几千张 GPU 上跑几十天,硬件故障不是异常(exception),是业务逻辑(business as usual)。文件里为此准备了一整套器官:
| 机制 | 对付什么 | 关键设计 |
|---|---|---|
| 异步 checkpoint | 存档时间挤占训练 | 每迭代非阻塞地推进落盘(maybe_finalize_async_save),退出时才阻塞收尾 |
| 本地非持久 checkpoint | 分布式文件系统太慢 | 存到本地盘,用 CliqueReplicationStrategy 在邻居节点间互备——本节点挂了,副本还在别人那里 |
| rerun state machine | SDC(静默算错) | 同批数据重放对比,区分瞬态硬件错误与软件 bug(见上文) |
| DP 副本参数哈希对账 | 副本漂移 | 定期对所有 DP 副本的参数求哈希并交叉比对——理论上永远相同,不同就是出事了 |
| GPU sniff test | 慢性硬件退化 | 训练前和定期跑一轮 GPU 自检(实现在别的文件,本文未展开) |
| 故障注入 | 容错代码本身没被测过 | FaultInjectorConfig 在指定迭代主动制造故障——训练系统的混沌工程 |
| 落野检测 | 个别慢卡拖垮全体 | StragglerDetector 用 FLOPs 记账数据找出持续偏慢的 rank |
两点评注。第一,参数哈希对账是「决策 all-reduce」思想的极致:DP 副本的参数在数学上必须逐位相同,任何分歧都意味着某张卡算错了或某次通信丢了数据——这是用冗余换来的免费校验和。第二,故障注入的存在说明容错代码被当作一等公民测试——如果恢复路径只在真故障时才第一次执行,那它大概率是坏的。这一整本账配合集群网络层的容错设计看会更完整:网络层处理链路故障,这一层处理计算故障。
性能本:凡是能藏的开销都要藏
通信藏进计算:pre-hook 与梯度桶
数据并行的两大通信——梯度归约(reduce-scatter)和参数收集(all-gather,distributed optimizer 分片后需要)——都不该让 GPU 干等。Megatron 的做法:
- 梯度侧:所有梯度放进连续的 grad buffer,切成桶;反向传播每算完一个桶的梯度,立刻异步启动这个桶的 reduce-scatter,通信和后续层的反向计算重叠。
- 参数侧:把 all-gather 挂在 forward pre-hook 上——前向即将用到某块参数时,钩子确保它已经(或正在)被收集,理想情况下上一层计算时下一层参数已在路上。
桶的大小有讲究。默认值是:
bucket_size = max(40000000, 1000000 * get_pg_size(dp_cp_group))
即 4000 万参数起步,DP 组每大一号加 100 万。手算三个点:DP=8 时 4000 万(下限兜底)、DP=64 时 6400 万、DP=512 时 5.12 亿。方向不难理解——桶按 DP 大小均分后每张卡的分片要足够大,否则集合通信退化为延迟主导;但这个具体系数是经验值,源码没有推导,我也没有实测过不同桶大小的吞吐差异(标注:未验证)。分片式数据并行的通信模式可以参考 ZeRO 论文(arXiv:1910.02054),Megatron 的 distributed optimizer 与 ZeRO-1 思路同源。
三种粒度的 CUDA Graph
文件里出现了三个独立的 CUDA Graph 开关:整迭代图(full_iteration,把一整个 train_step 的 kernel 序列录下来重放)、逐层图(transformer_engine,对每个 transformer 层单独录制)、优化器图(optimizer_cuda_graph)。它们对付同一个敌人——CPU 发射 kernel 的速度跟不上 GPU 消化的速度(launch overhead)。粒度越大收益越高,但约束也越苛刻:整迭代图要求迭代内的一切(包括 DDP 初始化用的 stream、FSDP 的参数收集时机)都是录制安全的,源码里好几处「因为 full_iteration 图所以必须……」的特判都是这个约束的涟漪。
我最喜欢的细节:跨 rank 对齐的垃圾回收
if args.manual_gc:
# Disable the default garbage collector and perform the collection manually.
# This is to align the timing of garbage collection across ranks.
gc.disable()
gc.collect()
Python 的 GC 会在任意时刻暂停进程几十毫秒。单机上无所谓;但在同步训练里,任何一个 rank 暂停,所有 rank 都在集合通信处等它。几千个 rank 各自随机 GC,等待就随机地叠加。解法朴素得动人:关掉自动 GC,所有 rank 在同一个迭代号一起 GC——把随机的停顿变成同步的、可记账的停顿。
这是我认为整份文件最有「分布式体感」的三行代码:它说明在同步 SPMD 系统里,连语言运行时的内务都必须是集体行为。看表要集体、判断成功要集体、连垃圾回收也要集体。
胶水本与 RL 支线:两笔简账
胶水本的核心是 ProcessGroupCollection——一个装着所有并行进程组(dp、tp、pp、cp、ep 及各种组合组)的容器,作为参数在整条调用链上传递。主循环从不假设自己知道全局拓扑,需要归约时问模型要它的 pg_collection。这样同一个 train() 能服务单机 8 卡、也能服务多模态的异构网格(不同子模型用不同并行度)。虚拟流水线(VPP)下连数据迭代器都是每个虚拟 stage 一份的列表——胶水的代价无处不在。DDP 包装、分片布局计算(wrap_model_chunks_with_ddp)也归这本账,上一篇模型配置三问已详细拆过,不重复。
RL 支线是这份 2025 版源码里最显眼的新器官:perform_rl_step 打开后,主循环每次迭代先用当前模型生成 rollout(GRPO 流程),再用生成的数据做若干步梯度更新。两个细节透露了工程量:其一,可以构建一个并行度完全不同的独立推理模型(推理不支持 CP、偏好低 PP),每轮把训练权重「refit」过去;其二,参考模型(reference model)的权重靠一段两次 load_checkpoint 的舞蹈拿到——先加载预训练 checkpoint 抓一份 CPU 状态字典,再把 RL checkpoint 加载回来。训练与推理两套模型在同一进程内共存、互相搬运权重——PPO 循环解剖里画过的那套多模型编排,正在被吸收进预训练框架本体。预训练与后训练的边界,在源码层面已经开始消融(这是观察,不是判断——也可能只是这个分支的特例)。
这份文件学不到什么
回到开头的问题:能靠它彻底搞懂预训练吗?现在可以给出精确的答案——这份文件是控制面的全部,但以下四样东西它一行都没有:
- 前向的数学:
forward_step_func是传进来的参数。模型结构在megatron/core/models/,损失怎么算在各pretrain_*.py入口脚本里。 - 流水线调度:
forward_backward_func是从megatron/core/pipeline_parallel/取来的现成函数。1F1B、virtual pipeline 的排布逻辑全在那边——这是数据面最精彩的部分。 - 优化器数学:
optimizer.step()是黑盒。分片、主权重副本、混合精度的细节在megatron/core/optimizer/。 - 数据管线:数据集怎么构建、怎么混合、怎么打包,在
dataset_provider的另一端。
反过来说,这个「什么都没有」正是它最深刻的设计:训练循环与它训练的东西彻底解耦。pretrain() 的函数签名就是证据——模型、数据、前向逻辑全是注入的依赖。同一个控制面既跑 GPT 也跑 Mamba 混合模型也跑多模态。上一篇讲的是「配置与构建分离」,这一篇的对应物是「编排与计算分离」——两篇合起来,Megatron 的分层就完整了。
带得走的东西:读任何训练框架主循环的六个问题
| # | 问题 | 在 Megatron 里的答案(供对照) |
|---|---|---|
| 1 | 二十行骨架在哪?先把它从包装里剥出来 | train() 的 while 循环 + train_step() |
| 2 | 一步更新的事务边界是什么?什么算「这步成功」? | rerun 事务包裹前反向;update_successful 跨 MP 组投票,全有或全无 |
| 3 | 进度用什么轴记账:迭代、样本还是 token? | 样本数为主轴(LR 调度、恢复定位),FLOPs 记吞吐,THD 下按真实 token 修正 |
| 4 | 每条退出/失败路径都过 checkpoint 吗? | 四条退出路径全部先存档;批大小变化也触发存档 |
| 5 | 哪些开销被藏进了哪个阶段? | 梯度通信藏进反向,参数收集藏进前向,kernel 发射藏进 CUDA Graph,GC 藏进对齐的迭代点 |
| 6 | 哪些决定必须跨 rank 集体做? | 退出判断、更新成败、参数哈希、GC 时机——单方面决定 = 死锁或分裂 |
这六个问题拿去读 DeepSpeed、TorchTitan 或任何内部训练框架的主循环,应该都能在半小时内建立骨架级的理解——框架间的差异几乎全部落在第 5、6 问的答案上。
诚实的提醒
- 素材边界:本文基于一份带实验特性的 Megatron-LM 训练侧源码(2025 版权头)。骨架(
pretrain/train/train_step/FLOPs 记账)与公开版本一致性高,但 RL 支线、GTP 重物化等属于较新的分支特性,你的版本可能没有。 - 亲手验算的:FLOPs 公式 vs 6ND 的比值 1.081、注意力 L² 占比 7.5%、packing 例子的 37.5%、桶大小三个取值——均用下面的脚本验算。未亲手验证的:桶大小对吞吐的实际影响、rerun state machine 的内部实现、GPU sniff test 的内容、「批大小变化即存档」的动机(已标注为推断)。
- 成本最低的验证实验:不需要 GPU。把下面的公式复算脚本跑一遍,换成你关心的任意模型配置(改 hidden/层数/FFN/词表/序列长),看 6ND 近似在你的场景里差多少——长上下文下这个差值会大到影响 MFU 结论:
h,L,ffn,vocab,s = 4096,32,11008,32000,4096
lin = 3*2*(h*3*h + h*h) + 3*2*h*ffn*3 # 注意力线性项 + SwiGLU MLP,每层每 token
core = 3*2*h # 注意力 L² 项系数,每层每单位 L²
tot = s*(L*lin + 3*2*h*vocab) + s*s*core*L # 一条序列的总 FLOPs
N = L*(4*h*h + 3*h*ffn) + vocab*h
print(f"formula/6ND = {tot/(6*N*s):.4f}, attn L^2 share = {s*s*core*L/tot:.1%}")
参考来源
工程实践
- Megatron-LM 源码仓库:https://github.com/NVIDIA/Megatron-LM (本文对应
megatron/training/training.py) - 本站前篇:《模型配置到底在配置什么:从 Megatron 源码读懂 Config、Spec、Builder 三层》
arXiv 论文(编号均已核实)
- Shoeybi et al., Megatron-LM: Training Multi-Billion Parameter Language Models Using Model Parallelism — arXiv:1909.08053
- Narayanan et al., Efficient Large-Scale Language Model Training on GPU Clusters Using Megatron-LM — arXiv:2104.04473(PTD-P 并行组合与吞吐分析)
- Kaplan et al., Scaling Laws for Neural Language Models — arXiv:2001.08361(6ND 成本近似的出处)
- Korthikanti et al., Reducing Activation Recomputation in Large Transformer Models — arXiv:2205.05198(源码 FLOPs 注释引用)
- Rajbhandari et al., ZeRO: Memory Optimizations Toward Training Trillion Parameter Models — arXiv:1910.02054(distributed optimizer 的思想源头)