01

train() 总览:训练主循环

training.py L3104-3697 核心循环

初始化篇 中,pretrain() 完成了模型、优化器和数据迭代器的构建。接下来,它调用 train() 进入训练主循环。

train() 的核心是一个 while iteration < train_iters 循环,每次迭代调用 train_step() 执行一步训练。在循环内部还穿插着评估、检查点保存和各种回调。

1.1 进入循环前的配置

train() 在进入主循环前,做了一系列重要的配置:

training.py — train() 循环前配置L3233-3293
# 1. 将 optimizer 的 loss scaling 函数注入 config
#    PP 调度器需要用它来 scale 梯度
config.grad_scale_func = optimizer.scale_loss

# 2. 设置 no_sync_func:控制何时执行梯度 AllReduce
#    overlap_grad_reduce 开启时,DDP 自行管理通信时机
if isinstance(model[0], DDP) and args.overlap_grad_reduce:
    config.no_sync_func = [model_chunk.no_sync for model_chunk in model]
    if args.align_grad_reduce:
        config.grad_sync_func = [model_chunk.start_grad_sync ...]

# 3. 设置 param_sync_func:控制何时执行参数 AllGather
#    overlap_param_gather 开启时,参数会在前向传播之前异步收集
if args.overlap_param_gather and args.align_param_gather:
    config.param_sync_func = [model_chunk.start_param_sync ...]

# 4. 设置梯度归一化函数
config.finalize_model_grads_func = finalize_model_grads

# 5. 选择 PP 调度函数
forward_backward_func = get_forward_backward_func()
config 的角色
config (TransformerConfig) 在这里充当了训练基础设施和 PP 调度器之间的桥梁。PP 调度器(如 1F1B)不直接依赖 DDP 或 optimizer,而是通过 config 中注入的函数间接调用。这种设计让 PP 调度器可以独立于具体的并行实现。

1.2 主循环结构

while iteration < args.train_iters: ┌─ 前置处理 │ · finalize_async_save() 完成上一轮异步存盘 │ · update_num_microbatches() microbatch 数量可能随 rampup 增加 │ · CUDA Graph capture 在 warmup_steps 后捕获计算图 │ ├─ train_step() ← 核心:单步训练 │ ├── zero_grad │ ├── forward_backward_func() 前向 + 反向(PP 调度) │ ├── optimizer.step() 参数更新 │ └── loss 聚合 │ ├─ forward pre-hook 管理 │ · 首轮成功后启用 pre-hook │ · 使 overlap_param_gather 生效 │ ├─ 后置处理 │ · iteration += 1 │ · consumed_train_samples 更新 │ · training_log() 日志输出 │ ├─ 周期性操作 │ · evaluate() eval_interval 轮一次 │ · post_training_step_callbacks GC / straggler / profiler │ · checkpoint_and_decide_exit save_interval 轮一次 + 退出判断 │ └─ 退出条件:iteration / duration / signal
02

train_step():单步训练

training.py L1515-1881 核心函数

train_step() 是训练循环中每个 iteration 的执行体。它完成一次完整的"前向-反向-更新"周期。

Phase 1: 清零梯度

model.zero_grad_buffer() + optimizer.zero_grad()
清空所有梯度缓冲区,为新一轮反向传播做准备。

Phase 2: 前向 + 反向(PP 调度)

forward_backward_func()
这是整个训练步的计算核心。PP 调度器会按照 1F1B 等策略,交替执行多个 microbatch 的前向和反向传播。

Phase 3: 优化器更新

optimizer.step()
执行梯度裁剪 → loss scaling → 参数更新。返回 (update_successful, grad_norm, num_zeros)

Phase 4: Loss 聚合

在 PP 最后一个 stage 上,将各 microbatch 的 loss 求平均,并通过 AllReduce 在 DP 组内汇总。

Phase 5: 学习率更新

opt_param_scheduler.step(increment=batch_size)
如果 optimizer.step() 成功(梯度无 NaN),推进 LR 调度器。否则标记 skipped_iter=1

2.1 Rerun 状态机驱动的重试

training.py — train_step() 中的重试循环L1559-1671
rerun_state_machine = get_rerun_state_machine()

# 注意:这不是简单的 "执行一次"
# rerun 状态机可能要求重新执行前向反向(容错重试)
while rerun_state_machine.should_run_forward_backward(data_iterator):
    # 清零梯度
    for model_chunk in model:
        model_chunk.zero_grad_buffer()
    optimizer.zero_grad()

    # 执行前向 + 反向
    losses_reduced = forward_backward_func(
        forward_step_func=forward_step_func,
        data_iterator=data_iterator,
        model=model,
        num_microbatches=get_num_microbatches(),
        seq_length=args.seq_length,
        micro_batch_size=args.micro_batch_size,
        forward_only=False,                        # 训练模式
    )

# 重试结束后,检查是否需要存盘退出
should_checkpoint, should_exit, exit_code = \
    rerun_state_machine.should_checkpoint_and_exit()
为什么前向反向在 while 循环里?
回忆 pretrain_gpt.html — loss_func() 中的三级诊断:当检测到 spiky loss 时,rerun 状态机会要求在相同 GPU不同 GPU 上重新执行前向反向,以区分是随机波动还是硬件故障。这个 while 循环就是重试的执行入口。

2.2 Loss 聚合逻辑

training.py L1826-1870

forward_backward_func() 返回的 losses_reduced 是一个列表,每个元素对应一个 microbatch 的 loss 字典。只有 PP 最后一个 stage 才有有效的 loss 值,其他 stage 返回空字典。

training.py — Loss 聚合(简化)L1826-1870
if mpu.is_pipeline_last_stage(ignore_virtual=True):
    loss_reduced = {}
    for key in losses_reduced[0].keys():
        val = [x[key].view(-1) for x in losses_reduced]

        if val[0].numel() == 2:
            # 新模式:per-token 归一化
            # 每个 microbatch 返回 [loss_sum, num_tokens]
            # 先对所有 microbatch 求和,再 AllReduce 跨 DP
            val = torch.vstack(val).sum(dim=0)   # [total_loss, total_tokens]
            torch.distributed.all_reduce(
                val,
                group=mpu.get_data_parallel_group(with_context_parallel=True)
            )
            loss_reduced[key] = val[0] / val[1]  # total_loss / total_tokens

        elif val[0].numel() == 1:
            # 旧模式:per-microbatch 平均
            val = torch.cat(val).mean()
            loss_reduced[key] = val
losses_reduced 的数据结构详解:key 长什么样?

losses_reduced 中每个字典的 key 来自 loss_func() 返回的 report 字典。不同训练模式下 key 不同:

pretrain_gpt.py — 标准预训练的 loss_funcL107-166
# 标准预训练:只有 1 个 key
report = {
    'lm loss': torch.cat([loss.clone().detach().view(1),
                          num_tokens.view(1)])
}
# value 是 2 元素 tensor: [loss_sum, num_valid_tokens]
post_training/loss_func.py — 知识蒸馏模式L39-70
# 知识蒸馏:最多 4 个 key
report = {
    'lm loss':                        tensor([loss_lm, num_tokens]),
    'total loss':                     tensor([kd_loss, num_tokens]),
    'logits distillation loss':       tensor([logits_loss, num_tokens]),
    'intermediate distillation loss': tensor([intermediate_loss, num_tokens]),
}

因此 losses_reduced 的完整结构如下(以标准预训练、3 个 microbatch 为例):

losses_reduced = [ {'lm loss': tensor([12.5, 100])}, ← microbatch 0: loss_sum=12.5, 100 个有效 token {'lm loss': tensor([ 9.8, 80])}, ← microbatch 1: loss_sum=9.8, 80 个有效 token {'lm loss': tensor([15.3, 120])}, ← microbatch 2: loss_sum=15.3, 120 个有效 token ] for key in losses_reduced[0].keys(): ↑ 标准预训练只循环 1 次: key = 'lm loss' ↑ 知识蒸馏循环 4 次: key = 'lm loss', 'total loss', ... 每个 key 对应的 value 都是 [loss_sum, num_tokens] (numel==2) → 统一走 per-token 归一化分支

注意:非 PP last stage 的 GPU 返回的是空字典 {},不会进入这段聚合逻辑。聚合出的 loss 仅用于日志打印和监控,不影响梯度——梯度在各 microbatch 的 backward() 时已经就地计算完毕。

Per-token vs Per-microbatch 归一化
新版本 Megatron 使用 per-token 归一化:loss_func 返回 [loss_sum, num_valid_tokens],最终 loss = 总 loss / 总 token 数。这比 per-microbatch 平均更准确,因为不同 microbatch 中有效 token 数量可能不同(padding、loss_mask 等导致)。
03

forward_backward_func:PP 调度选择

schedules.py L46-146 Pipeline Parallelism

forward_backward_functrain_step() 中最核心的调用。它不是一个固定函数,而是根据并行配置动态选择的 PP 调度策略。

schedules.py — get_forward_backward_func()L46-146
def get_forward_backward_func():
    pp_size = parallel_state.get_pipeline_model_parallel_world_size()
    vp_size = parallel_state.get_virtual_pipeline_model_parallel_world_size()

    if pp_size > 1:
        if vp_size is not None:
            # VPP 模式:交错 1F1B
            return forward_backward_pipelining_with_interleaving
        else:
            # 标准 PP:1F1B
            return forward_backward_pipelining_without_interleaving
    else:
        # 无 PP:简单遍历所有 microbatch
        return forward_backward_no_pipelining
条件调度函数行为
pp_size == 1 forward_backward_no_pipelining 顺序遍历所有 microbatch,每个做完前向立刻反向。最简单。
pp_size > 1,无 VPP forward_backward_pipelining_without_interleaving 标准 1F1B:warmup 阶段连续前向填满 pipeline,稳态阶段交替 1 前向 1 反向,cooldown 阶段连续反向排空。
pp_size > 1,有 VPP forward_backward_pipelining_with_interleaving 交错 1F1B:每个 GPU 持有多个不连续层块,更细粒度的交错减少 pipeline bubble。
forward_backward_func 的统一接口
无论选择哪种调度策略,它们的接口完全相同:
losses = forward_backward_func(
    forward_step_func,   # pretrain_gpt.py 提供的回调
    data_iterator,       # 数据迭代器
    model,               # 模型(可能是列表)
    num_microbatches,    # 本步的 microbatch 数量
    forward_only,        # True=评估模式,False=训练模式
    ...
)
PP 调度器内部负责:microbatch 切分、P2P 通信、前向/反向交替、梯度累积。

调度器内部会多次调用 forward_step_func()(即 pretrain_gpt.py 中的 forward_step()),每次处理一个 microbatch。在 PP 最后一个 stage,forward_step_func 返回的 loss 函数会被立即调用以计算损失。

调度时序图与三阶段详解

1F1B · 1F1B-I (VPP) 调度图及 warmup / 稳态 / cooldown 三阶段逐步解读, 见 Pipeline Parallelism 调度机制

04

forward pre-hook 机制

training.py L3340-3349, L3491-3517 overlap_param_gather

--overlap-param-gather 开启时,DistributedOptimizer 将参数分片到各 DP rank。前向传播前,需要通过 AllGather 收集完整参数。为了让这个 AllGather 与前一层的计算重叠,Megatron 使用了 PyTorch 的 forward_pre_hook 机制。

overlap_param_gather 工作原理 背景:DistributedOptimizer 将参数分片到各 DP rank,每个 rank 只持有 1/N 参数。 前向传播前必须 AllGather 收集完整参数。 不重叠(禁用 pre-hook): AllGather 所有层的参数,完成后再开始 Forward — 通信和计算完全串行 NIC: [====== AllGather 全部参数 ======] GPU: [= Fwd L0 =][= Fwd L1 =][= Fwd L2 =] ├────── 通信阻塞,GPU 空闲 ──────┤├──────── GPU 计算 ────────────────┤ 重叠(启用 pre-hook): 每层注册 forward_pre_hook,hook 内调用 finish_param_sync() 做两件事: 1. Wait:等待当前层的 AG 完成(拿到本层完整参数) 2. Prefetch:异步发起下一层的 AG(预取下一层参数) NIC: [AG L0][AG L1] [AG L2] [AG L3] GPU: [= Fwd L0 =][= Fwd L1 =][= Fwd L2 =][= Fwd L3 =] ├─ 重叠 ──┤├─ 重叠 ──┤├─ 重叠 ──┤ AG(L1)在 AG(L2)在 AG(L3)在 NIC上异步 NIC上异步 NIC上异步 与Fwd(L0) 与Fwd(L1) 与Fwd(L2) GPU计算并行 GPU计算并行 GPU计算并行 逐步执行时序: ┌──────────────────────────────────────────────────────────────────────┐ │ ① pre_hook(L0) → finish_param_sync(L0) │ │ ├─ L0 的 AG 尚未发起 → 同步发起 AG(L0) 并等待完成 │ │ └─ Prefetch: 异步发起 AG(L1),不等待 │ │ │ │ ② Fwd(L0) 在 GPU 上执行 │ │ └─ 与此同时,AG(L1) 在 NIC 上并行传输 ← 重叠发生在这里! │ │ │ │ ③ pre_hook(L1) → finish_param_sync(L1) │ │ ├─ 等待 AG(L1) 完成(通常已经传完,无需等待) │ │ └─ Prefetch: 异步发起 AG(L2) │ │ │ │ ④ Fwd(L1) 在 GPU 上执行,AG(L2) 在 NIC 上并行传输 │ │ │ │ ⑤ ... 重复直到最后一层(最后一层无 prefetch) │ └──────────────────────────────────────────────────────────────────────┘ 关键:重叠之所以可行,是因为 AllGather 由 NIC (网卡 RDMA) 驱动, Forward 由 GPU SM 驱动 — 两者是不同硬件,可以真正并行。

但是,train() 对 pre-hook 有一个延迟启用策略:

training.py — 延迟启用 pre-hookL3340-3349, L3491-3507
# ========== 循环开始前:禁用 pre-hook ==========
if should_disable_forward_pre_hook(args):
    disable_forward_pre_hook(model, param_sync=False)
    param_sync_func = config.param_sync_func
    config.param_sync_func = None      # 也暂时禁用 param_sync
    pre_hook_enabled = False

# ========== 首次成功训练后:启用 pre-hook ==========
if iteration == start_iteration:
    if skipped_iter:
        # FP16 loss scaling 还没稳定,继续等待
        start_iteration = iteration + 1
    else:
        # 首次成功!启用 pre-hook
        enable_forward_pre_hook(model)
        config.param_sync_func = param_sync_func
        pre_hook_enabled = True
为什么首轮要禁用?
首轮训练时,参数可能来自 checkpoint 加载或随机初始化。如果某个 rank 加载失败或初始化异常,forward_pre_hook 触发的 AllGather 会将错误参数扩散到所有 DP rank

禁用 pre-hook 意味着首轮每个 rank 使用自己本地的参数。如果首轮训练成功(没有 NaN、没有 skip),说明所有 rank 的参数是正确的,此时再启用 pre-hook 就安全了。
05

评估与检查点

training.py L3593-3656 evaluate() checkpoint

5.1 周期性评估

training.py — 评估触发逻辑L3593-3627
# 每 eval_interval 个 iteration 做一次评估
if args.eval_interval and iteration % args.eval_interval == 0 and args.do_valid:
    # 1. 暂停训练计时器(不计入训练时间)
    timers('interval-time').stop()

    # 2. 禁用 forward pre-hook(评估不需要梯度通信)
    disable_forward_pre_hook(model)

    # 3. 执行评估
    evaluate_and_print_results(
        prefix=f'iteration {iteration}',
        forward_step_func=forward_step_func,
        data_iterator=valid_data_iterator,
        model=model,
        iteration=iteration,
        config=config,
    )

    # 4. 重新启用 pre-hook + 恢复计时器
    enable_forward_pre_hook(model)
    timers('interval-time', log_level=0).start(barrier=True)

evaluate() 函数的核心逻辑:

  • 切换到 model.eval() 模式(禁用 Dropout)
  • 禁用 rerun 状态机(评估不需要容错重试)
  • torch.no_grad() 下执行 eval_iters 次前向传播
  • 调用 forward_backward_func(forward_only=True)——与训练共用同一个 PP 调度器,但只做前向
  • 在 PP last stage 上聚合 loss,通过 AllReduce 跨 DP 汇总
  • 结束后切回 model.train() 模式

5.2 checkpoint_and_decide_exit()

training.py L2990-3101

每个 iteration 结束后,checkpoint_and_decide_exit() 检查是否需要保存 checkpoint 和/或退出训练。它处理三种退出条件:

退出条件参数行为
信号退出 --exit-signal-handler 收到 SIGTERM 时保存 checkpoint 并退出。用于 SLURM 作业超时前的优雅退出。
时间退出 --exit-duration-in-mins 训练超过指定分钟数后退出。所有 rank 通过 AllReduce 同步判断(任一 rank 超时则全部退出)。
步数退出 --exit-interval 每隔 N 个 iteration 退出一次。用于周期性重启训练。

此外,还有两种 checkpoint 保存模式:

  • 持久化 checkpoint--save-interval):保存到共享存储(如 NFS/S3),跨节点可恢复
  • 非持久化 checkpoint--non-persistent-save-interval):保存到本地 SSD,速度更快但不跨节点持久
Microbatch Rampup 触发的自动存盘
--rampup-batch-size 导致 microbatch 数量增加时,train() 会在变化点自动保存 checkpoint(L3400-3421)。这是因为 microbatch 数量变化会改变训练行为,保存一个 checkpoint 确保可以从变化点精确恢复。
06

辅助机制

CUDA Graph Profiling GC

train() 循环中还嵌入了多种辅助机制,这些不是训练的核心逻辑,但对大规模训练的稳定性和可观测性至关重要。

6.1 CUDA Graph

CUDA Graph 将一系列 GPU 操作"录制"成计算图,后续直接重放,避免了每次 iteration 的 CPU kernel launch 开销。Megatron 支持两种实现:

实现参数特点
transformer_engine --cuda-graph-impl te 由 TE 的 TECudaGraphHelper 管理,在 warmup 步后捕获。
local --cuda-graph-impl local Megatron 自研 FullCudaGraphWrapper,包装整个 forward_backward_func

CUDA Graph 需要固定的 tensor 地址和形状,因此只有在 warmup 步(--cuda-graph-warmup-steps)之后才能捕获。

6.2 Profiling

train() 支持两种 profiling 方式:

  • PyTorch Profiler--use-pytorch-profiler):在指定的 step 范围内收集 trace,输出到 TensorBoard
  • Nsys NVTX(默认):在指定 step 范围内启用 CUDA profiler + NVTX 标注,配合 nsight systems 使用

两者通过 --profile-step-start--profile-step-end 控制采集范围,只在 --profile-ranks 指定的 rank 上启用。

6.3 手动垃圾回收

--manual-gc 开启时,train() 会在循环开始前禁用 Python 自动 GC,改为按 --manual-gc-interval 手动触发。

为什么? Python GC 的触发时机不确定。在分布式训练中,如果一个 rank 的 GC 恰好在 NCCL 通信期间触发,会导致该 rank 延迟,引发集体通信超时。手动 GC 确保所有 rank 在相同的 iteration 触发 GC,避免时序不一致。

6.4 Straggler 检测

StragglerDetector 在每个 log_interval 统计各 rank 的 FLOPS,找出最慢的 rank(straggler)。这有助于诊断大集群中的性能异常节点。通过 --log-straggler 启用。

6.5 容错集成(ft_integration)

ft_integration 提供了与 NVIDIA Resiliency Extension 的集成接口。它在关键时间点发送心跳信号:

  • on_training_step_start/end — 每步训练前后
  • on_checkpointing_start/end — 存盘前后
  • on_eval_step_start/end — 评估步前后

外部容错监控器通过心跳间隔判断是否有 rank hang 住,从而触发自动重启。

07

数据流全景图

pretrain_gpt.pyoptimizer.step(),数据经过了以下完整路径:

完整数据流(一次 train_step 的执行路径) pretrain_gpt.py │ 提供 4 个回调 ▼ training.py :: pretrain() │ 调用 initialize → setup → build_iterators → train() ▼ training.py :: train() │ while iteration < train_iters: ▼ training.py :: train_step() │ ├── zero_grad() │ ├── forward_backward_func() ← PP 调度器 │ │ │ │ 对每个 microbatch: │ │ │ ├── forward_step() ← pretrain_gpt.py 的回调 │ │ ├── get_batch() 从迭代器取数据 │ │ │ ├── TP broadcast TP rank 0 广播到组内 │ │ │ └── CP split 切分到 CP rank │ │ ├── model(tokens, ...) GPT 前向传播 │ │ └── return (output, loss_fn) │ │ │ ├── P2P send/recv PP stage 间传递激活值 │ │ │ ├── loss_func() ← 仅 PP last stage │ │ ├── cross_entropy + loss_mask │ │ └── NaN/spiky 检测 │ │ │ └── backward() 反向传播 + 梯度累积 │ └── P2P send/recv PP stage 间传递梯度 │ ├── finalize_model_grads() 梯度 AllReduce (DP 组内) │ ├── optimizer.step() 参数更新 │ ├── unscale_and_check_inf │ ├── clip_grad_norm │ ├── adam/muon update │ └── AllGather 更新后的参数 (分布式优化器) │ └── opt_param_scheduler.step() LR 更新
系列文档导航