跳转至

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 版本的子组件。