01

全局架构:这个文件在做什么?

pretrain_gpt.py 是 Megatron-LM 训练 GPT 模型的入口脚本。它本身不实现底层的并行策略或训练循环,而是作为一个粘合层 (glue layer),通过回调函数把各个子系统组装在一起。

它的核心设计模式是依赖注入training.py 中的 pretrain() 定义了通用的训练框架,而本文件提供 GPT 特定的实现——如何取数据、如何建模型、如何做前向计算。

pretrain_gpt.py 提供 4 个回调,注入到通用训练框架
pretrain_gpt.py
① 数据来源 train_valid_test_datasets_provider
② 模型构建 model_provider (gpt_builder)
③ 前向计算 forward_step
④ 模型类型标识 ModelType.encoder_or_decoder
training.py
pretrain()
· 分布式初始化
· 训练循环
· Checkpoint
· Logging
BERT → pretrain_bert.py  ·  T5 → pretrain_t5.py  ·  只需提供不同回调即可复用框架
设计哲学
这种 "框架 + 回调" 的设计,使得 Megatron 的核心训练逻辑(并行策略、梯度累积、checkpoint 等)可以在不同模型间复用,而模型特定的代码被隔离在各自的 pretrain_*.py 中。
02

入口点 __main__:一切从这里开始

L301-318 入口
pretrain_gpt.py — __main__ L301-318
if __name__ == "__main__":

    # Temporary for transition to core datasets
    train_valid_test_datasets_provider.is_distributed = True

    # Optionally enable inprocess restart on pretrain
    pretrain, store = inprocess_restart.maybe_wrap_for_inprocess_restart(pretrain)

    pretrain(
        train_valid_test_datasets_provider,          # 回调①:构建数据集
        partial(model_provider, gpt_builder),         # 回调②:构建模型
        ModelType.encoder_or_decoder,                 # 模型类型枚举
        forward_step,                                 # 回调③:单步前向
        args_defaults={'tokenizer_type': 'GPT2BPETokenizer'},
        extra_args_provider=add_modelopt_args if has_nvidia_modelopt else None,
        store=store,
        get_embedding_ranks=get_embedding_ranks,
    )

逐行解读:

  • is_distributed = True — 告诉框架数据集构建需要在所有 rank 上协同进行(分布式构建),而不是每个 rank 独立构建
  • inprocess_restart.maybe_wrap_for_inprocess_restart() — 支持进程内重启(fault tolerance),训练出错时不需要杀掉重启进程
  • partial(model_provider, gpt_builder) — 将 gpt_builder 绑定为 model_provider 的第一个参数。这样框架调用 model_provider(pre_process, post_process) 时,实际执行的是 model_provider(gpt_builder, pre_process, post_process)
  • ModelType.encoder_or_decoder — 告诉 Pipeline Parallelism 调度器这是一个 decoder-only 模型(不需要 encoder-decoder 交叉注意力)
pretrain() 被调用后做了什么?
training.py::pretrain() 内部按顺序执行:
initialize_megatron() → 初始化分布式、解析参数
setup_model_and_optimizer() → 建模型 + DDP 包装 + 创建优化器
build_train_valid_test_data_iterators() → 建数据集
train() → 进入主训练循环
03

get_batch():数据如何流入模型

L42-100 数据流 并行切分

这是数据预处理的核心函数,负责从 data_iterator 取出一个 micro-batch,并根据各种并行策略做切分。

pretrain_gpt.py — get_batch() 第一部分:PP stage 过滤 L42-49
def get_batch(data_iterator, vp_stage: Optional[int] = None):
    """Generate a batch."""
    args = get_args()
    config = core_transformer_config_from_args(args)
    # 如果不在 PP 首尾 stage,且不是 MTP rank → 不需要数据
    if not is_first_or_last_pipeline_stage(vp_stage) and (
    (not mtp_on_this_rank(config, ignore_virtual=False, vp_stage=vp_stage))):
        return None, None, None, None, None, None, None
为什么中间 PP stage 不需要数据?
在 Pipeline Parallelism 中:
· 第一个 stage:需要 tokens 做 embedding
· 最后一个 stage:需要 labelsloss_mask 算 loss
· 中间 stage:输入来自上一个 stage 的激活值,完全不需要原始数据
pretrain_gpt.py — get_batch() 第二部分:TP 广播 L51-55
    # 只在 TP rank 0 上真正读数据,其他 TP rank 通过 broadcast 接收
    batch = get_batch_on_this_tp_rank(
        data_iterator,
        mtp_on_this_rank=mtp_on_this_rank(config, ignore_virtual=False, vp_stage=vp_stage)
    )

同一个 TP (Tensor Parallel) 组内的所有 rank 处理的是相同的数据(因为张量并行是把模型参数切分,不是把数据切分)。所以只需要 TP rank 0 从 dataloader 读取,然后广播给组内其他 rank。

pretrain_gpt.py — get_batch() 第三部分:MTP 预处理 L66-80
    # Multi-Token Prediction:在 CP 切分之前,预先准备多个预测头的数据
    mtp_data = None
    if config.mtp_num_layers and mtp_on_this_rank(config, ignore_virtual=False, vp_stage=vp_stage):
        mtp_data = prepare_mtp_data(
            batch['tokens'], batch['position_ids'],
            batch['labels'], batch['loss_mask'],
            config.mtp_num_layers, cu_seqlens=cu_seqlens,
        )

    # 如果同时有 MTP 和 packed/hybrid CP,需要把 MTP 数据混入 batch
    # 这样后续 CP 切分会一并处理它们
    mtp_needs_batch_split = mtp_data is not None and (cu_seqlens is not None or local_cp_size is not None)
    if mtp_needs_batch_split:
        for k in range(config.mtp_num_layers):
            batch[f'mtp_input_ids_{k}'] = mtp_data['mtp_input_ids'][k]
            batch[f'mtp_position_ids_{k}'] = mtp_data['mtp_position_ids'][k]
            # ...
pretrain_gpt.py — get_batch() 第四部分:Context Parallelism 切分 L82-100
    if cu_seqlens is None and local_cp_size is None:
        # 普通 CP:沿序列维度均匀切分
        batch = get_batch_on_this_cp_rank(batch)
        packed_seq_params = None
    elif local_cp_size is None:  # Packed THD format
        # 变长序列打包格式(多条不等长序列拼在一起)
        batch, packed_seq_params = get_thd_batch_on_this_cp_rank(
            batch, cu_seqlens, cu_seqlens_padded, max_seqlen
        )
    else: # Hybrid CP format
        batch, packed_seq_params = get_batch_on_this_hybrid_cp_rank(batch, local_cp_size)

    # 最终返回 7 个元素
    return (*batch.values(), packed_seq_params, mtp_data)
flowchart TD DI([data_iterator]) --> PP PP["① PP 过滤\n中间 stage → 返回 7 个 None"] --> TP TP["② TP 广播\nrank 0 读取,broadcast 给组内所有 rank"] --> MTP MTP["③ MTP 准备\n启用时 roll 出多头 token"] --> CP CP["④ CP 切分\n沿序列维度"] --> N1 & N2 & N3 N1["普通模式\n均匀切分"] N2["THD packed\n变长序列打包"] N3["Hybrid CP\n混合切分"] N1 & N2 & N3 --> OUT(["tokens / labels / loss_mask\nattention_mask / position_ids\npacked_seq_params / mtp_data"])
04

loss_func():损失计算与异常检测

L107-166 损失函数 容错
pretrain_gpt.py — loss_func() 核心计算 L107-133
def loss_func(
    loss_mask: torch.Tensor, output_tensor: torch.Tensor, model: Optional[GPTModel] = None
):
    """Loss function.

    Args:
        loss_mask: 用于屏蔽 padding token 的 loss
        output_tensor: 模型输出的 per-token loss
        model: GPT 模型实例(用于 ModelOpt 蒸馏场景)
    """
    args = get_args()

    # 标准 loss 计算
    losses = output_tensor.view(-1).float()     # 展平为一维
    loss_mask = loss_mask.view(-1).float()       # 展平 mask
    loss = torch.sum(losses * loss_mask)         # 只对非 padding token 求和

    num_tokens = loss_mask.sum().clone().detach().to(torch.int)
    report = {'lm loss': torch.cat([loss.clone().detach().view(1), num_tokens.view(1)])}

loss_mask 的作用是屏蔽不参与训练的 token。典型场景:

  • padding token(批内对齐补零的位置)
  • SFT 场景中 prompt 部分的 token(只对 response 部分计算 loss)
  • EOD (End of Document) 后的 token

注意 report 字典里把 lossnum_tokens 拼在一起——后续框架会在 DP 组内做 all-reduce 来计算全局平均 loss。

pretrain_gpt.py — loss_func() 异常检测 L136-164
    # NaN / Inf 检查
    rerun_state_machine = get_rerun_state_machine()
    if args.check_for_nan_in_loss_and_grad:
        rerun_state_machine.validate_result(
            result=loss,
            rejection_func=torch.isnan,                    # 检测 NaN
            message="found NaN in local forward loss calculation",
            tolerance=0.0,       # 前向计算是确定性的,不允许任何容忍度
            fatal=True,          # 致命错误,必须 rerun
        )
        rerun_state_machine.validate_result(
            result=loss,
            rejection_func=torch.isinf,                    # 检测 Inf
            message="found Inf in local forward loss calculation",
            tolerance=0.0,
            fatal=True,
        )

    # Spiky loss 检查(loss 突然飙升)
    if args.check_for_spiky_loss:
        rerun_state_machine.validate_result(
            result=loss,
            rejection_func=partial(
                rerun_state_machine.is_unexpectedly_large,
                threshold=SPIKY_LOSS_FACTOR,               # 阈值 = 10x
                context="loss",
            ),
            message="Spiky loss",
            tolerance=0.0,
            fatal=False,         # 非致命,会 rerun 但不会崩溃
        )
Rerun State Machine:三阶段硬件故障诊断
这是 Megatron 的容错机制rerun_state_machine.py)。检测到异常结果后,不是直接崩溃,而是通过最多三轮执行来定位根因:
  1. 同 GPU 重跑 — 恢复 RNG 和数据迭代器状态,在同一 GPU 上重跑。如果结果不同 → 瞬态错误(如 bit flip),丢弃继续
  2. 换 GPU 重跑 — 结果可复现时,保存 checkpoint 并退出,由调度器在另一台机器的 GPU 上重启
  3. 最终判定 — 换 GPU 后结果仍一致 → 结果本身是正确的,非硬件故障
fatal 参数的含义
fatal 只在第三阶段"结果确认正确"时生效:
  • fatal=True(NaN/Inf)— 即使确认不是硬件故障,也终止训练。因为 loss 是 NaN 说明模型或数据有 bug
  • fatal=False(spiky loss)— 确认不是硬件故障就继续训练。loss 飙升可能只是数据分布问题
is_unexpectedly_large 是一个自适应阈值检测器:先观察前 100 步记录历史最大值,之后如果当前值超过历史最大值的 SPIKY_LOSS_FACTOR(10x)倍就触发 rerun。
05

forward_step():单步前向传播

L169-207 核心回调
pretrain_gpt.py — forward_step() L169-207
def forward_step(data_iterator, model: GPTModel, return_schedule_plan: bool = False):
    """Forward training step."""
    args = get_args()
    timers = get_timers()

    # ① 取数据
    timers('batch-generator', log_level=2).start()
    global stimer
    with stimer(bdata=True):
        vp_stage = get_attr_wrapped_model(model, "vp_stage")
        tokens, labels, loss_mask, attention_mask, position_ids, \
            packed_seq_params, mtp_data = get_batch(data_iterator, vp_stage)
    timers('batch-generator').stop()

    # ② 前向计算
    with stimer:
        output_tensor = model(
            tokens, position_ids, attention_mask,
            labels=labels, loss_mask=loss_mask,
            packed_seq_params=packed_seq_params, mtp_data=mtp_data,
        )

    # ③ 返回 (模型输出, 延迟执行的 loss 函数)
    return output_tensor, partial(loss_func, loss_mask, model=model)
为什么返回 partial(loss_func) 而不是直接算 loss?
这是为了适配 Pipeline Parallelism 的设计:
  • 非最后 stageoutput_tensor 是中间激活值,传给下一个 stage,loss_func 不会被调用
  • 最后 stageoutput_tensor 是 per-token loss,PP 调度器调用 loss_func(output_tensor) 汇总得到标量 loss
loss_mask 在取数据时就确定了,所以通过 partial 提前绑定;而 output_tensor 要等前向完成后才有,由调度器传入。

stimerStragglerDetector(掉队检测器)的实例。它通过计时来发现哪些 rank 的计算特别慢(straggler),帮助定位性能瓶颈。bdata=True 标记数据加载阶段,无参数的 with stimer 标记计算阶段。

forward_step 在 Pipeline Parallelism 中的行为:

PP Stage 0 (first)PP Stage 1 (middle)PP Stage 2 (last)
get_batch()
→ 拿到 tokens

Embedding
Transformer × N

pre_process=True
get_batch()
→ 返回全 None


Transformer × N

两者都=False
get_batch()
→ 拿到 labels


Transformer × N
Output Layer
loss_func()
post_process=True
Stage 0 ──激活值→ Stage 1 ──激活值→ Stage 2
PP 调度器如何使用 forward_step
你写的 forward_step 是一个回调函数,由 PP 调度器在合适的时机反复调用。调度器通过 set_input_tensor() 把上一个 stage 的输出注入模型,使得每个 stage 虽然都调用 model(tokens, ...),但模型内部会根据自己是哪个 stage 决定是用 tokens 做 embedding 还是直接用上游传来的激活值。
1F1B 调度三阶段
Warmup
连续前向,填满流水线
recv_forward → forward → send_forward
1F1B 稳态
一次前向 + 一次反向交替
forward → send_forward_recv_backward → backward
Cooldown
清空剩余反向传播
recv_backward → backward → send_backward

完整时间线图、Interleaved VPP 调度及 P2P 通信细节 → Pipeline Parallelism 调度机制

06

数据集构建:train_valid_test_datasets_provider()

L219-283 数据集
pretrain_gpt.py — GPTDatasetConfig 构建 L219-253
def core_gpt_dataset_config_from_args(args):
    if args.legacy_tokenizer:
        tokenizer = get_tokenizer()
    else:
        tokenizer = build_tokenizer(args)

    # 从命令行参数解析数据混合比例
    blend, blend_per_split = get_blend_and_blend_per_split(args)

    return GPTDatasetConfig(
        random_seed=args.seed,
        sequence_length=args.seq_length,
        blend=blend,                          # 多数据源混合比例
        blend_per_split=blend_per_split,       # train/val/test 各自的混合
        split=args.split,                      # "98,1,1" 格式的 train/val/test 比例
        tokenizer=tokenizer,
        reset_position_ids=args.reset_position_ids,
        reset_attention_mask=args.reset_attention_mask,
        eod_mask_loss=args.eod_mask_loss,      # EOD 后的 token 是否计算 loss
        context_parallel_size=args.context_parallel_size,
        data_parallel_size=args.data_parallel_size,
        # ...
    )
pretrain_gpt.py — 数据集构建 L256-283
def train_valid_test_datasets_provider(train_val_test_num_samples, vp_stage=None):
    """Build the train test and validation datasets."""
    args = get_args()
    config = core_gpt_dataset_config_from_args(args)

    # 根据训练模式选择数据集类型
    if args.sft:
        dataset_type = SFTDataset           # SFT 指令微调
    else:
        if args.mock_data:
            dataset_type = MockGPTDataset   # 假数据(调试用)
        else:
            dataset_type = GPTDataset       # 标准预训练

    # BlendedMegatronDatasetBuilder 支持多数据集按比例混合
    train_ds, valid_ds, test_ds = BlendedMegatronDatasetBuilder(
        dataset_type,
        train_val_test_num_samples,
        partial(is_dataset_built_on_rank, vp_stage=vp_stage),
        config
    ).build()

    return train_ds, valid_ds, test_ds

几个关键设计点:

  • Blending--data-path 可以传多个数据文件和权重,如 0.7 wiki.bin 0.3 books.bin,框架会按比例混合采样
  • 按 rank 过滤is_dataset_built_on_rank 确保只有需要数据的 rank 构建数据集(PP 首尾 stage + TP rank 0),节省内存
  • 三种数据集GPTDataset(预训练)、SFTDataset(指令微调)、MockGPTDataset(测试用的随机数据)
07

模型构建:model_provider() + gpt_builder()

model_provider.py gpt_builders.py 模型构建

模型构建分两层:model_provider() 是通用壳,gpt_builder() 是 GPT 特定实现。

model_provider.py — 通用模型提供函数 L24-67
def model_provider(
    model_builder: Callable,
    pre_process=True,     # 是否包含 embedding(PP 第一个 stage)
    post_process=True,    # 是否包含 output head(PP 最后一个 stage)
    vp_stage=None         # Virtual Pipeline stage 编号
) -> Union[GPTModel, MambaModel]:
    args = get_args()

    # OOM 时自动保存 memory snapshot(调试用)
    if args.record_memory_history:
        torch.cuda.memory._record_memory_history(True, ...)
        torch._C._cuda_attach_out_of_memory_observer(oom_observer)

    return model_builder(args, pre_process, post_process, vp_stage)
gpt_builders.py — GPT 模型构建核心 L27-99
def gpt_builder(args, pre_process, post_process, vp_stage=None, config=None):
    # ① 构建 TransformerConfig(从命令行参数转化)
    if config is None:
        config = core_transformer_config_from_args(args)

    # ② 选择 Transformer Layer 的具体实现规格 (spec)
    if args.spec is not None:
        transformer_layer_spec = import_module(args.spec)     # 用户自定义 spec
    else:
        use_te = args.transformer_impl == "transformer_engine"
        if args.num_experts:
            # MoE 模型:使用 decoder block spec(包含 Router + Expert)
            transformer_layer_spec = get_gpt_decoder_block_spec(config, use_te, ...)
        else:
            # 标准 Dense 模型
            if use_te:
                transformer_layer_spec = get_gpt_layer_with_transformer_engine_spec(...)
            else:
                transformer_layer_spec = get_gpt_layer_local_spec(...)

    # ③ (可选) Multi-Token Prediction
    mtp_block_spec = None
    if args.mtp_num_layers is not None:
        mtp_block_spec = get_gpt_mtp_block_spec(config, ...)

    # ④ 实例化 GPTModel
    model = GPTModel(
        config=config,
        transformer_layer_spec=transformer_layer_spec,
        vocab_size=args.padded_vocab_size,
        max_sequence_length=args.max_position_embeddings,
        pre_process=pre_process,              # PP: 是否包含 embedding
        post_process=post_process,            # PP: 是否包含 output head
        share_embeddings_and_output_weights=not args.untie_embeddings_and_output_weights,
        position_embedding_type=args.position_embedding_type,
        rotary_base=args.rotary_base,
        mtp_block_spec=mtp_block_spec,
    )
    return model
什么是 Layer Spec?
Megatron 用 Spec 模式来定义 Transformer 层的组件。一个 spec 描述了"这个 Transformer 层应该用什么 Attention 实现、什么 MLP 实现、什么 LayerNorm"等。

例如 get_gpt_layer_with_transformer_engine_spec() 会返回一个使用 NVIDIA Transformer Engine 加速的 spec(支持 FP8),而 get_gpt_layer_local_spec() 返回纯 PyTorch 实现。
flowchart TD PG["pretrain_gpt.py"] -- "partial(model_provider, gpt_builder)" --> MP MP["model_provider\npre_process / post_process / vp_stage"] --> GB GB["gpt_builder ①②③④"] --> TC & LS TC["① TransformerConfig\nhidden_size / num_layers / num_attention_heads ..."] LS["② Layer Spec"] --> D1 & D2 & D3 D1["Dense: TE spec / local spec"] D2["MoE: decoder block spec\n(config 作为参数)"] D3["custom: args.spec"] TC -- "config=" --> GM LS -- "transformer_layer_spec=" --> GM GM["④ GPTModel"] --> E & T & O E["Embedding\n(pre_process only)"] T["TransformerBlock × N"] O["Output Layer\n(post_process only)"]
08

辅助函数

L210-298
pretrain_gpt.py — is_dataset_built_on_rank() L210-216
def is_dataset_built_on_rank(vp_stage=None):
    args = get_args()
    config = core_transformer_config_from_args(args)
    return (
        is_first_or_last_pipeline_stage(vp_stage)
        or mtp_on_this_rank(config, ignore_virtual=False, vp_stage=vp_stage)
    ) and parallel_state.get_tensor_model_parallel_rank() == 0

判断当前 rank 是否需要构建数据集。条件:(PP 首尾 stage 或 MTP rank) 且 TP rank 0

  • TP 组内数据相同,只需 rank 0 构建然后广播
  • PP 中间 stage 不需要数据
  • MTP rank 需要额外的预测头数据
pretrain_gpt.py — get_embedding_ranks() L286-298
def get_embedding_ranks(pp_ranks: List[int]):
    """确定哪些 PP rank 需要持有 embedding 权重。"""
    embedding_ranks = [pp_ranks[0]]        # 第一个 stage 一定需要
    if len(pp_ranks) > 1:
        args = get_args()
        if not args.untie_embeddings_and_output_weights:
            embedding_ranks.append(pp_ranks[-1])   # 权重共享时,最后一个 stage 也需要
        config = core_transformer_config_from_args(args)
        mtp_ranks = get_mtp_ranks(pp_ranks, config)
        embedding_ranks.extend(mtp_ranks)          # MTP rank 也需要 embedding
    return sorted(list(set(embedding_ranks)))
Tied Embeddings(权重共享)
GPT 模型通常让 input embedding 和 output linear layer 共享权重(untie_embeddings_and_output_weights=False)。在 PP 下,这意味着第一个 stage(有 embedding)和最后一个 stage(有 output layer)必须持有同一份权重,需要额外的同步。
09

总结:数据流全景图

flowchart TD subgraph GPT["pretrain_gpt.py"] CB1["datasets_provider()"] CB2["model_provider() + gpt_builder()"] CB3["forward_step()"] end subgraph TR["training.py"] PRETRAIN["pretrain()"] INIT["initialize
分布式初始化"] SETUP["setup_model
模型+优化器"] DATA["build_data
数据集"] LOOP["train loop"] end PRETRAIN --> INIT PRETRAIN --> SETUP PRETRAIN --> DATA CB1 -.->|回调| DATA CB2 -.->|回调| SETUP INIT & SETUP & DATA --> LOOP subgraph FWD["forward_step() — 每个 micro-batch 调用一次"] GET["get_batch()
PP 过滤 → TP 广播 → MTP 准备 → CP 切分"] MODEL["model()
前向计算"] LOSS["loss_func()
mask 加权 · NaN 检测 · spiky 检测"] end LOOP -->|"每个 iteration"| FWD CB3 -.->|回调| FWD GET --> MODEL --> LOSS LOOP --> BW["backward → optimizer → checkpoint → logging"]

总结一下 pretrain_gpt.py 的核心职责:

定义数据流:get_batch()

从 data_iterator 取数据 → PP/TP/CP/MTP 多维度切分 → 返回 7 元素 tuple。每种并行都有对应的数据处理逻辑。

定义损失计算:loss_func()

per-token loss × loss_mask → 标量 loss。附加 NaN/Inf/spiky loss 检测和 rerun 容错机制。

组装前向步骤:forward_step()

调用 get_batch + model + partial(loss_func)。返回 (output, loss_fn) 供 PP 调度器使用。

注入到训练框架:__main__

把上述回调传给 training.py::pretrain(),由框架负责训练循环、checkpoint、logging 等通用逻辑。

下一步阅读建议
  • Rank 与并行组 — 理解 GPU 编号如何映射到 TP/PP/DP 多维网格
  • megatron/core/parallel_state.py — TP/PP/DP/CP 并行组如何划分
  • megatron/training/training.py :: train() — 训练主循环的具体实现
  • megatron/core/pipeline_parallel/schedules.py — 1F1B 调度如何驱动 forward_step
  • megatron/core/rerun_state_machine.py — 三阶段硬件故障诊断的完整实现