transformer_impl.py — Megatron 训练引擎实现¶
文件概述¶
基于 Megatron-LM 的训练引擎实现(约 870 行),支持 TP + PP + EP + CP 的全套并行策略,适合超大规模模型训练。
核心类¶
MegatronEngine(BaseEngine)¶
class MegatronEngine(BaseEngine):
def __init__(self, model_config, engine_config, ...):
# 初始化 Megatron 并行状态
mpu.initialize_model_parallel(
tensor_model_parallel_size=...,
pipeline_model_parallel_size=...,
expert_model_parallel_size=...,
context_parallel_size=...,
)
初始化流程¶
def initialize(self):
self._build_tf_config() # 构建 TransformerConfig
self.module = self._build_megatron_module() # 构建 Megatron 模型
self._maybe_enable_fused_kernels() # 融合内核
self.optimizer = self._build_optimizer() # Megatron 优化器
self.lr_scheduler = self._build_lr_scheduler() # 学习率调度器
模型构建(Megatron-Bridge)¶
def _build_tf_config(self):
"""使用 Megatron-Bridge 将 HuggingFace 配置转换为 Megatron 配置"""
if self.vanilla_bridge:
bridge = AutoBridge.from_config(self.model_config.hf_config, ...)
else:
bridge = AutoBridge.from_hf_pretrained(self.model_config.local_path, ...)
provider = bridge.to_megatron_provider(load_weights=False)
前向+反向(流水线并行)¶
def forward_backward_batch(self, data, loss_function, forward_only=False):
"""使用 Megatron 的 PP 调度器"""
forward_backward_func = get_forward_backward_func()
# Megatron PP 调度器自动处理:
# - 微批次在 PP stage 之间的传递
# - 1F1B 或 interleaved 调度策略
losses_reduced = forward_backward_func(
forward_step_func=self.forward_step,
data_iterator=batch_generator,
model=self.module,
num_microbatches=n_micro_batch,
)
Router Replay 支持¶
if enable_routing_replay:
# 训练时重用推理阶段的路由决策
RouterReplay.set_global_router_replay_action(RouterReplayAction.REPLAY_FORWARD)
MegatronEngineWithLMHead¶
@EngineRegistry.register(model_type="language_model", backend="megatron")
class MegatronEngineWithLMHead(MegatronEngine):
def forward_step(self, batch_iter, model, postprocess_micro_batch_func):
"""语言模型的前向步骤"""
# 支持融合内核和非融合内核两种路径
if use_fused_kernels:
output = fused_forward_fn(model, input_ids, ...)
else:
output = forward_fn(model, input_ids, ...)
MegatronEngineWithValueHead¶
@EngineRegistry.register(model_type="value_model", backend="megatron")
class MegatronEngineWithValueHead(MegatronEngineWithLMHead):
"""Critic 价值模型的 Megatron 引擎"""
权重同步¶
def get_per_tensor_param(self, base_sync_done=False, **kwargs):
"""导出权重用于同步到推理引擎
通过 Megatron-Bridge 将 Megatron 格式权重转换为 HuggingFace 格式
"""
if self.vanilla_bridge:
per_tensor_param = self.bridge.export_weights(self.module)
else:
per_tensor_param = self.bridge.export_hf_weights(self.module)
与其他模块的关系¶
- 继承
engine/base.py的BaseEngine - 使用 Megatron Core 的并行基础设施
- 使用
engine/megatron/utils.py的随机种子工具 - 被
engine_workers.py通过EngineRegistry使用
小结¶
Megatron 引擎提供了最完整的并行支持(TP+PP+EP+CP),适合训练超大规模模型。通过 Megatron-Bridge 实现 HuggingFace 和 Megatron 之间的模型格式转换。