全局架构:这个文件在做什么?
pretrain_gpt.py 是 Megatron-LM 训练 GPT 模型的入口脚本。它本身不实现底层的并行策略或训练循环,而是作为一个粘合层 (glue layer),通过回调函数把各个子系统组装在一起。
它的核心设计模式是依赖注入:training.py 中的 pretrain() 定义了通用的训练框架,而本文件提供 GPT 特定的实现——如何取数据、如何建模型、如何做前向计算。
pretrain_*.py 中。
入口点 __main__:一切从这里开始
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 交叉注意力)
training.py::pretrain() 内部按顺序执行:①
initialize_megatron() → 初始化分布式、解析参数②
setup_model_and_optimizer() → 建模型 + DDP 包装 + 创建优化器③
build_train_valid_test_data_iterators() → 建数据集④
train() → 进入主训练循环
get_batch():数据如何流入模型
这是数据预处理的核心函数,负责从 data_iterator 取出一个 micro-batch,并根据各种并行策略做切分。
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
· 第一个 stage:需要
tokens 做 embedding· 最后一个 stage:需要
labels 和 loss_mask 算 loss· 中间 stage:输入来自上一个 stage 的激活值,完全不需要原始数据
# 只在 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。
# 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]
# ...
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)
loss_func():损失计算与异常检测
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 字典里把 loss 和 num_tokens 拼在一起——后续框架会在 DP 组内做 all-reduce 来计算全局平均 loss。
# 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.py)。检测到异常结果后,不是直接崩溃,而是通过最多三轮执行来定位根因:
- 同 GPU 重跑 — 恢复 RNG 和数据迭代器状态,在同一 GPU 上重跑。如果结果不同 → 瞬态错误(如 bit flip),丢弃继续
- 换 GPU 重跑 — 结果可复现时,保存 checkpoint 并退出,由调度器在另一台机器的 GPU 上重启
- 最终判定 — 换 GPU 后结果仍一致 → 结果本身是正确的,非硬件故障
fatal 只在第三阶段"结果确认正确"时生效:
fatal=True(NaN/Inf)— 即使确认不是硬件故障,也终止训练。因为 loss 是 NaN 说明模型或数据有 bugfatal=False(spiky loss)— 确认不是硬件故障就继续训练。loss 飙升可能只是数据分布问题
is_unexpectedly_large 是一个自适应阈值检测器:先观察前 100 步记录历史最大值,之后如果当前值超过历史最大值的 SPIKY_LOSS_FACTOR(10x)倍就触发 rerun。
forward_step():单步前向传播
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)
- 非最后 stage:
output_tensor是中间激活值,传给下一个 stage,loss_func不会被调用 - 最后 stage:
output_tensor是 per-token loss,PP 调度器调用loss_func(output_tensor)汇总得到标量 loss
loss_mask 在取数据时就确定了,所以通过 partial 提前绑定;而 output_tensor 要等前向完成后才有,由调度器传入。
stimer 是 StragglerDetector(掉队检测器)的实例。它通过计时来发现哪些 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
|
forward_step 是一个回调函数,由 PP 调度器在合适的时机反复调用。调度器通过 set_input_tensor() 把上一个 stage 的输出注入模型,使得每个 stage 虽然都调用 model(tokens, ...),但模型内部会根据自己是哪个 stage 决定是用 tokens 做 embedding 还是直接用上游传来的激活值。
recv_forward → forward → send_forward
forward → send_forward_recv_backward → backward
recv_backward → backward → send_backward
完整时间线图、Interleaved VPP 调度及 P2P 通信细节 → Pipeline Parallelism 调度机制
数据集构建:train_valid_test_datasets_provider()
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,
# ...
)
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(测试用的随机数据)
模型构建:model_provider() + gpt_builder()
模型构建分两层:model_provider() 是通用壳,gpt_builder() 是 GPT 特定实现。
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)
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
例如
get_gpt_layer_with_transformer_engine_spec() 会返回一个使用 NVIDIA Transformer Engine 加速的 spec(支持 FP8),而 get_gpt_layer_local_spec() 返回纯 PyTorch 实现。
辅助函数
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 需要额外的预测头数据
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)))
untie_embeddings_and_output_weights=False)。在 PP 下,这意味着第一个 stage(有 embedding)和最后一个 stage(有 output layer)必须持有同一份权重,需要额外的同步。
总结:数据流全景图
分布式初始化"] 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_stepmegatron/core/rerun_state_machine.py— 三阶段硬件故障诊断的完整实现