parallel_decoder.py — Qwen2 并行解码层¶
文件路径¶
verl/models/qwen2/megatron/layers/parallel_decoder.py
文件概述¶
Qwen2 的 Transformer 解码层,结构与 LLaMA 完全相同。提供标准版和去 padding 版两个变体。
关键代码讲解¶
class ParallelQwen2DecoderLayer(nn.Module):
def __init__(self, config, megatron_config, layer_idx):
self.self_attn = ParallelQwen2Attention(config=config, megatron_config=megatron_config)
self.mlp = ParallelQwen2MLP(config, megatron_config=megatron_config)
self.input_layernorm = ParallelQwen2RMSNorm(config, megatron_config)
self.post_attention_layernorm = ParallelQwen2RMSNorm(config, megatron_config)
class ParallelQwen2DecoderLayerRmPad(nn.Module):
def __init__(self, config, megatron_config, layer_idx):
self.self_attn = ParallelQwen2AttentionRmPad(config=config, megatron_config=megatron_config)
# 其余与标准版相同
前向传播逻辑与 LLaMA 完全一致:Pre-norm -> Attention -> 残差 -> Pre-norm -> MLP -> 残差。
与其他模块的关系¶
- 被
modeling_qwen2_megatron.py使用 - 组合了 Qwen2 版本的注意力、MLP、RMSNorm
小结¶
Qwen2 解码层与 LLaMA 的 ParallelLlamaDecoderLayer 结构完全对称,仅使用了 Qwen2 版本的子组件。