二十行骨架与四本账:Megatron 预训练主循环源码深读

基于 Megatron 训练侧 training.py 真实源码,把预训练主循环拆成「二十行理想骨架 + 四本工程账(记账/容错/性能/并行胶水)」。含 FLOPs 公式与 6ND 的手算对照、packed sequence 的 L² 记账、跨 rank 对齐 GC 等细节,末尾给一张读任何训练框架主循环都能用的六问清单。

把 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 / bucketDDP 把所有梯度放进一块连续大缓冲区,再按「桶」为单位分批做通信,让通信和反向计算重叠
distributed optimizer把优化器状态(动量等)按 DP 维度分片,每张卡只存 1/N,思路与 ZeRO 第一阶段同源
loss scalingfp16 训练时把 loss 乘一个大数再反传,防梯度下溢;溢出时该步作废、缩小系数重来
forward pre-hook挂在前向计算之前的钩子函数,Megatron 用它把「取回分片参数的 all-gather」藏进前向
CUDA Graph把一串 GPU kernel 的启动序列录下来整体重放,消除 CPU 逐个发射 kernel 的开销
BSHD / THD两种 batch 布局:定长矩形(batch×seq)vs 把多条变长序列首尾相接打包成一条(sequence packing)
SDCSilent Data Corruption,硬件静默出错——不崩溃、不报错,只是算出来的数悄悄不对
GRPO一种 RL 后训练算法(PPO 家族),见《PPO 训练循环解剖》

主线:二十行骨架,四本账

先给这篇文章的一句话主线(这是我的提炼,不是源码里的官方说法):这份文件 = 一个二十行就能写完的理想训练循环,外面包了四本工程账。 读它的正确姿势不是逐行读,而是先把骨架认领出来,再把剩下的每一段代码归入四本账之一:

  1. 记账本(observability):每一秒、每一个 FLOP、每一焦耳都要有出处——吞吐、进度日志、能耗、计时器。
  2. 容错本(fault tolerance):几千张卡跑几十天,硬件故障不是异常而是常态——checkpoint、重跑机、哈希对账、故障注入。
  3. 性能本(performance):凡是能藏起来的开销都要藏——通信与计算重叠、CUDA Graph、跨 rank 对齐的 GC。
  4. 胶水本(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 数组记录边界)。打包后有两笔账变了:

  1. token 线性项(投影、MLP、logits):应该按真实 token 数 sum(Lᵢ) 记,padding 不算工作量;
  2. 注意力 L² 项:打包后每条子序列只在自己内部做注意力(分块因果掩码),FLOPs 是 sum(Lᵢ²) 而不是 (总长)²

这两个量在数学上互相独立(知道和推不出平方和),所以必须分别追踪。差距有多大?手算一个例子:把 [2048, 1024, 1024] 三条序列打包成一条 4096 的序列——

Li2s2=20482+10242+1024240962=629145616777216=37.5%\frac{\sum L_i^2}{s^2} = \frac{2048^2 + 1024^2 + 1024^2}{4096^2} = \frac{6291456}{16777216} = 37.5\%

注意力的真实计算量只有按整条序列算的 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 machineSDC(静默算错)同批数据重放对比,区分瞬态硬件错误与软件 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 循环解剖里画过的那套多模型编排,正在被吸收进预训练框架本体。预训练与后训练的边界,在源码层面已经开始消融(这是观察,不是判断——也可能只是这个分支的特例)。

这份文件学不到什么

回到开头的问题:能靠它彻底搞懂预训练吗?现在可以给出精确的答案——这份文件是控制面的全部,但以下四样东西它一行都没有:

  1. 前向的数学forward_step_func 是传进来的参数。模型结构在 megatron/core/models/,损失怎么算在各 pretrain_*.py 入口脚本里。
  2. 流水线调度forward_backward_func 是从 megatron/core/pipeline_parallel/ 取来的现成函数。1F1B、virtual pipeline 的排布逻辑全在那边——这是数据面最精彩的部分。
  3. 优化器数学optimizer.step() 是黑盒。分片、主权重副本、混合精度的细节在 megatron/core/optimizer/
  4. 数据管线:数据集怎么构建、怎么混合、怎么打包,在 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%}")

参考来源

工程实践

arXiv 论文(编号均已核实)

  • Shoeybi et al., Megatron-LM: Training Multi-Billion Parameter Language Models Using Model ParallelismarXiv:1909.08053
  • Narayanan et al., Efficient Large-Scale Language Model Training on GPU Clusters Using Megatron-LMarXiv:2104.04473(PTD-P 并行组合与吞吐分析)
  • Kaplan et al., Scaling Laws for Neural Language ModelsarXiv:2001.08361(6ND 成本近似的出处)
  • Korthikanti et al., Reducing Activation Recomputation in Large Transformer ModelsarXiv:2205.05198(源码 FLOPs 注释引用)
  • Rajbhandari et al., ZeRO: Memory Optimizations Toward Training Trillion Parameter ModelsarXiv:1910.02054(distributed optimizer 的思想源头)