跳转至

modeling_qwen2_megatron.py — Qwen2 Megatron 主模型

文件路径

verl/models/qwen2/megatron/modeling_qwen2_megatron.py

文件概述

Qwen2 模型在 Megatron 框架下的完整实现。结构与 LLaMA 版本高度相似,主要区别在于: 1. 使用 Qwen2Config 替代 LlamaConfig 2. PP 版本支持 share_embeddings_and_output_weights(权重共享) 3. 注意力层的 bias 设置不同

关键代码讲解

与 LLaMA 版本的差异

1. 权重共享(tie_word_embeddings)

Qwen2 的 PP 版本支持 embedding 和 lm_head 的权重共享:

class ParallelQwen2ForCausalLMRmPadPP(nn.Module):
    def __init__(self, config, megatron_config, pre_process, post_process,
                 share_embeddings_and_output_weights):
        if post_process:
            self.lm_head = tensor_parallel.ColumnParallelLinear(
                ...,
                # 如果共享权重,跳过权重分配
                skip_weight_param_allocation=self.pre_process and self.share_embeddings_and_output_weights,
            )
        if pre_process or post_process:
            self.setup_embeddings_and_output_layer()

    def setup_embeddings_and_output_layer(self):
        """设置 embedding 和 output 层的权重共享"""
        if self.pre_process:
            self.model.embed_tokens.weight.is_embedding_or_output_parameter = True
        if self.post_process and self.lm_head.weight is not None:
            self.lm_head.weight.is_embedding_or_output_parameter = True

        if not self.share_embeddings_and_output_weights:
            return

        if parallel_state.get_pipeline_model_parallel_world_size() == 1:
            # 单 PP:同一 stage 上的两层,避免梯度重复累加
            self.shared_embedding_or_output_weight().zero_out_wgrad = True
        else:
            # 多 PP:通过 all_reduce 同步 embedding 权重
            if parallel_state.is_pipeline_first_stage() and self.pre_process:
                self.shared_embedding_or_output_weight().shared_embedding = True
            if self.post_process and not self.pre_process:
                self.lm_head.weight.data.fill_(0)
                self.lm_head.weight.shared = True

    def _forward_head(self, hidden_states):
        output_weight = None
        if self.share_embeddings_and_output_weights:
            output_weight = self.shared_embedding_or_output_weight()
        logits = self.lm_head(hidden_states, weight=output_weight)[0]

2. 模型类结构

与 LLaMA 版本完全对称的六个类:

类名 对应 LLaMA 类 额外特性
ParallelQwen2Model ParallelLlamaModel -
ParallelQwen2ForCausalLM ParallelLlamaForCausalLM -
ParallelQwen2ModelRmPad ParallelLlamaModelRmPad -
ParallelQwen2ForCausalLMRmPad ParallelLlamaForCausalLMRmPad -
ParallelQwen2ForValueRmPad ParallelLlamaForValueRmPad -
ParallelQwen2ModelRmPadPP ParallelLlamaModelRmPadPP -
ParallelQwen2ForCausalLMRmPadPP ParallelLlamaForCausalLMRmPadPP 支持 share_embeddings_and_output_weights
ParallelQwen2ForValueRmPadPP ParallelLlamaForValueRmPadPP -

与其他模块的关系

  • 使用 layers/ 中的 Qwen2 并行组件
  • 被 verl 训练流程中的 actor/critic 使用
  • 与 checkpoint_utils/ 配合进行权重加载/保存

小结

Qwen2 Megatron 模型与 LLaMA 版本结构高度一致,核心差异是 Qwen2 支持 tie_word_embeddings(共享 embedding 和 lm_head 权重),这在 PP 场景下需要通过 setup_embeddings_and_output_layer 方法在不同 PP stage 间同步权重。