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 线性层内部。