跳转至

modeling_llama_megatron.py — LLaMA Megatron 主模型

文件路径

verl/models/llama/megatron/modeling_llama_megatron.py

文件概述

这是 LLaMA 模型在 Megatron 框架下的完整实现。它定义了 6 个模型类,覆盖了三种使用场景: 1. 带 padding 的基础版本:ParallelLlamaModel + ParallelLlamaForCausalLM 2. 去 padding (RmPad) 版本:ParallelLlamaModelRmPad + ParallelLlamaForCausalLMRmPad + ParallelLlamaForValueRmPad 3. 去 padding + 流水线并行 (PP) 版本:ParallelLlamaModelRmPadPP + ParallelLlamaForCausalLMRmPadPP + ParallelLlamaForValueRmPadPP

关键代码讲解

1. 基础模型 ParallelLlamaModel

class ParallelLlamaModel(nn.Module):
    def __init__(self, config: LlamaConfig, megatron_config: ModelParallelConfig):
        # 使用 Megatron 的 VocabParallelEmbedding(词表按 TP 切分)
        self.embed_tokens = tensor_parallel.VocabParallelEmbedding(
            num_embeddings=config.vocab_size, embedding_dim=config.hidden_size, **embedding_kwargs
        )
        # Transformer 解码层
        self.layers = nn.ModuleList(
            [ParallelLlamaDecoderLayer(config, megatron_config) for _ in range(config.num_hidden_layers)]
        )
        self.norm = ParallelLlamaRMSNorm(config, megatron_config)

这是标准的 Transformer 解码器结构,使用了 Megatron 的 TP 并行嵌入。

2. 因果语言模型 ParallelLlamaForCausalLM

class ParallelLlamaForCausalLM(nn.Module):
    def __init__(self, config, megatron_config):
        self.model = ParallelLlamaModel(config, megatron_config)
        # lm_head 使用 ColumnParallelLinear,输出按 TP 切分
        self.lm_head = tensor_parallel.ColumnParallelLinear(
            input_size=config.hidden_size, output_size=config.vocab_size,
            bias=False, gather_output=False, ...
        )

    def forward(self, input_ids, attention_mask, position_ids):
        hidden_states = self.model(input_ids, attention_mask, position_ids)
        logits = self.lm_head(hidden_states)[0]
        # 从所有 TP rank 收集完整的 logits
        logits = tensor_parallel.gather_from_tensor_model_parallel_region(logits)
        return CausalLMOutputWithPast(logits=logits.float())

3. 去 Padding 版本 (RmPad)

去 padding 是核心优化:将 batch 中的 padding token 去除,只计算有效 token。

class ParallelLlamaForCausalLMRmPad(nn.Module):
    def forward(self, input_ids, attention_mask, position_ids):
        batch_size, sequence_length = input_ids.shape
        # 1. 去除 padding,得到紧凑表示
        input_ids, indices, cu_seqlens, max_seqlen_in_batch, *_ = unpad_input(
            input_ids.unsqueeze(dim=-1), attention_mask
        )  # (total_nnz, 1)

        # 2. 如果启用序列并行,pad 到 TP 的倍数
        if self.megatron_config.sequence_parallel:
            input_ids = sp_utils.pad_to_sequence_parallel(input_ids)

        # 3. 前向传播
        input_ids = input_ids.transpose(0, 1)  # (1, total_nnz+pad)
        outputs = self.model(input_ids=input_ids, ...)

        # 4. 去掉 SP padding + 恢复原始 padding
        logits = self._forward_head(hidden_states)
        if self.megatron_config.sequence_parallel:
            logits = logits[:cu_seqlens[-1]]
        logits = pad_input(logits, indices, batch_size, seqlen=sequence_length)

数据流变化:(batch, seq) -> unpad -> (total_nnz,) -> SP pad -> (total_nnz_padded,) -> 模型计算 -> SP unpad -> pad_input -> (batch, seq)

4. Value 模型

Value 模型继承 CausalLM,但将 lm_head 改为输出单个标量值:

class ParallelLlamaForValueRmPad(ParallelLlamaForCausalLMRmPad):
    def _init_head(self, config):
        # 输出维度为 1(标量值)
        self.lm_head = nn.Linear(in_features=config.hidden_size, out_features=1, bias=False)
        # 标记为序列并行参数
        sp_utils.mark_parameter_as_sequence_parallel(self.lm_head.weight)

    def _forward_head(self, hidden_states):
        logits = self.lm_head(hidden_states)
        if self.megatron_config.sequence_parallel:
            logits = tensor_parallel.gather_from_sequence_parallel_region(logits)
        return logits

5. 流水线并行版本 (PP)

PP 版本的核心区别:每个 PP stage 只包含部分层。

class ParallelLlamaModelRmPadPP(nn.Module):
    def __init__(self, config, megatron_config, pre_process, post_process):
        # pre_process=True 的 stage 才有 embedding
        if pre_process:
            self.embed_tokens = tensor_parallel.VocabParallelEmbedding(...)
        # 只创建属于本 PP stage 的层
        pp_rank = mpu.get_pipeline_model_parallel_rank()
        if vpp_size is not None:
            # 虚拟 PP:交错分配层
            offset = vpp_rank * (num_layers // vpp_size) + (pp_rank * num_layer_vpp_chunk)
        else:
            offset = pp_rank * num_layer_per_pp
        for i in range(self.num_layer_this_model):
            layer = ParallelLlamaDecoderLayerRmPad(config, megatron_config, layer_idx=offset + i)
            self.layers.add_module(f"{i}", layer)
        # post_process=True 的 stage 才有 final norm
        if post_process:
            self.norm = ParallelLlamaRMSNorm(config, megatron_config)

    def set_input_tensor(self, input_tensor):
        """PP 中间 stage 通过此方法接收上一个 stage 的输出"""
        self.input_tensor = input_tensor

    def forward(self, input_ids, ...):
        if self.pre_process:
            hidden_states = self.embed_tokens(input_ids)
        else:
            hidden_states = self.input_tensor  # 来自上一个 PP stage
        for decoder_layer in self.layers:
            hidden_states = decoder_layer(hidden_states, ...)
        if self.post_process:
            hidden_states = self.norm(hidden_states)
        return hidden_states

辅助函数

def _make_causal_mask(input_ids_shape, dtype, device):
    """创建因果注意力掩码(下三角矩阵)"""

def _expand_mask(mask, dtype, tgt_len=None):
    """将 [bsz, seq_len] 的 mask 扩展为 [bsz, 1, tgt_len, src_len]"""

核心类列表

类名 父类 特点
ParallelLlamaModel nn.Module 基础 Transformer,带 padding
ParallelLlamaForCausalLM nn.Module CausalLM,带 padding
ParallelLlamaModelRmPad nn.Module 去 padding,使用 flash_attn_varlen
ParallelLlamaForCausalLMRmPad nn.Module 去 padding 的 CausalLM
ParallelLlamaForValueRmPad ...RmPad Value 模型,输出标量
ParallelLlamaModelRmPadPP nn.Module 去 padding + PP
ParallelLlamaForCausalLMRmPadPP nn.Module 去 padding + PP 的 CausalLM
ParallelLlamaForValueRmPadPP ...RmPadPP Value 模型 + PP

与其他模块的关系

  • 使用 layers/ 子目录中的并行层组件
  • 被 verl 的训练流程调用(作为 actor/critic 模型)
  • 与 checkpoint_utils/ 配合进行权重加载/保存

小结

这个文件提供了 LLaMA 在 Megatron 下的完整模型定义,按需支持三种并行模式。去 padding (RmPad) 通过 unpad_input/pad_input 减少无效计算,PP 版本通过 pre_process/post_process 标志和 set_input_tensor 方法实现流水线并行。Value 模型通过继承 CausalLM 并重写 head 来输出标量值。