跳转至

parallel_decoder.py — LLaMA 并行解码层

文件路径

verl/models/llama/megatron/layers/parallel_decoder.py

文件概述

定义 LLaMA 的 Transformer 解码层,组合注意力、MLP 和 LayerNorm。提供标准版和去 padding 版两个变体。

关键代码讲解

标准解码层 ParallelLlamaDecoderLayer

class ParallelLlamaDecoderLayer(nn.Module):
    def __init__(self, config, megatron_config, layer_idx):
        self.self_attn = ParallelLlamaAttention(config=config, megatron_config=megatron_config)
        self.mlp = ParallelLlamaMLP(config, megatron_config=megatron_config)
        self.input_layernorm = ParallelLlamaRMSNorm(config, megatron_config)
        self.post_attention_layernorm = ParallelLlamaRMSNorm(config, megatron_config)

    def forward(self, hidden_states, attention_mask, position_ids):
        # Pre-norm 架构
        residual = hidden_states
        hidden_states = self.input_layernorm(hidden_states)
        hidden_states = self.self_attn(hidden_states, attention_mask, position_ids)
        hidden_states = residual + hidden_states

        residual = hidden_states
        hidden_states = self.post_attention_layernorm(hidden_states)
        hidden_states = self.mlp(hidden_states)
        hidden_states = residual + hidden_states
        return hidden_states

去 Padding 解码层 ParallelLlamaDecoderLayerRmPad

class ParallelLlamaDecoderLayerRmPad(nn.Module):
    def __init__(self, config, megatron_config, layer_idx):
        # 使用去 padding 版注意力
        self.self_attn = ParallelLlamaAttentionRmPad(config=config, megatron_config=megatron_config)
        # MLP 和 Norm 不变
        self.mlp = ParallelLlamaMLP(config, megatron_config=megatron_config)

    def forward(self, hidden_states, position_ids, sequence_length,
                indices, cu_seqlens, max_seqlen_in_batch):
        # 输入形状: (total_nnz // sp, 1, hidden_size)
        residual = hidden_states
        hidden_states = self.input_layernorm(hidden_states)
        # 注意力内部处理 SP 的 gather/scatter
        hidden_states = self.self_attn(hidden_states, position_ids, sequence_length,
                                        indices, cu_seqlens, max_seqlen_in_batch)
        hidden_states = residual + hidden_states

        residual = hidden_states
        hidden_states = self.post_attention_layernorm(hidden_states)
        hidden_states = self.mlp(hidden_states)
        hidden_states = residual + hidden_states
        return hidden_states

序列并行下的数据流

输入: (total_nnz // sp, 1, hidden_size)    -- SP 分片
    |
  [RMSNorm]                                 -- 在 SP 分片上计算
    |
  [Attention: ColumnParallel all-gather -> compute -> RowParallel reduce-scatter]
    |
  [残差连接]                                 -- SP 分片
    |
  [RMSNorm]
    |
  [MLP: ColumnParallel all-gather -> compute -> RowParallel reduce-scatter]
    |
  [残差连接]
    |
输出: (total_nnz // sp, 1, hidden_size)    -- SP 分片

与其他模块的关系

  • 被 modeling_llama_megatron.py 的各个模型类使用
  • 组合了 parallel_attention.py、parallel_mlp.py、parallel_rmsnorm.py

小结

解码层采用 Pre-norm 残差结构(先 norm 再计算)。去 padding 版本接收 cu_seqlens 等变长信息,传递给注意力层进行 flash attention 计算。序列并行的 all-gather 和 reduce-scatter 隐藏在 ColumnParallel 和 RowParallel 线性层内部。