train() 总览:训练主循环
在 初始化篇 中,pretrain() 完成了模型、优化器和数据迭代器的构建。接下来,它调用 train() 进入训练主循环。
train() 的核心是一个 while iteration < train_iters 循环,每次迭代调用 train_step() 执行一步训练。在循环内部还穿插着评估、检查点保存和各种回调。
1.1 进入循环前的配置
train() 在进入主循环前,做了一系列重要的配置:
# 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 (TransformerConfig) 在这里充当了训练基础设施和 PP 调度器之间的桥梁。PP 调度器(如 1F1B)不直接依赖 DDP 或 optimizer,而是通过 config 中注入的函数间接调用。这种设计让 PP 调度器可以独立于具体的并行实现。
1.2 主循环结构
train_step():单步训练
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 状态机驱动的重试
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()
2.2 Loss 聚合逻辑
forward_backward_func() 返回的 losses_reduced 是一个列表,每个元素对应一个 microbatch 的 loss 字典。只有 PP 最后一个 stage 才有有效的 loss 值,其他 stage 返回空字典。
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
[loss_sum, num_valid_tokens],最终 loss = 总 loss / 总 token 数。这比 per-microbatch 平均更准确,因为不同 microbatch 中有效 token 数量可能不同(padding、loss_mask 等导致)。
forward_backward_func:PP 调度选择
forward_backward_func 是 train_step() 中最核心的调用。它不是一个固定函数,而是根据并行配置动态选择的 PP 调度策略。
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。 |
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 调度机制。
forward pre-hook 机制
当 --overlap-param-gather 开启时,DistributedOptimizer 将参数分片到各 DP rank。前向传播前,需要通过 AllGather 收集完整参数。为了让这个 AllGather 与前一层的计算重叠,Megatron 使用了 PyTorch 的 forward_pre_hook 机制。
但是,train() 对 pre-hook 有一个延迟启用策略:
# ========== 循环开始前:禁用 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
forward_pre_hook 触发的 AllGather 会将错误参数扩散到所有 DP rank。禁用 pre-hook 意味着首轮每个 rank 使用自己本地的参数。如果首轮训练成功(没有 NaN、没有 skip),说明所有 rank 的参数是正确的,此时再启用 pre-hook 就安全了。
评估与检查点
5.1 周期性评估
# 每 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()
每个 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,速度更快但不跨节点持久
--rampup-batch-size 导致 microbatch 数量增加时,train() 会在变化点自动保存 checkpoint(L3400-3421)。这是因为 microbatch 数量变化会改变训练行为,保存一个 checkpoint 确保可以从变化点精确恢复。
辅助机制
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 住,从而触发自动重启。
数据流全景图
从 pretrain_gpt.py 到 optimizer.step(),数据经过了以下完整路径:
- pretrain_gpt.py 源码精读 — 入口脚本、4 个回调函数
- Rank 与并行组 — GPU 映射与通信组创建
- 集合通信操作详解 — AllReduce / AllGather / P2P 等
- training.py 初始化篇 — pretrain() → initialize → setup
- training.py 训练循环篇(本文)— train() → train_step() → 评估 → 检查点