跳转至

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 之间的模型格式转换。