跳转至

parallel_attention.py — Qwen2 并行注意力

文件路径

verl/models/qwen2/megatron/layers/parallel_attention.py

文件概述

Qwen2 的张量并行注意力实现。与 LLaMA 版本结构基本一致,主要区别是 Qwen2 的 QKV 投影带有 bias,且 RoPE 只使用标准版本。

关键代码讲解

与 LLaMA 注意力的区别

class ParallelQwen2Attention(nn.Module):
    def __init__(self, config, megatron_config):
        # Qwen2 的 QKV 带 bias(LLaMA 默认不带)
        self.qkv_proj = QKVParallelLinear(
            ..., bias=True, ...  # Qwen2 特有
        )
        # Qwen2 的 o_proj 不带 bias
        self.o_proj = tensor_parallel.RowParallelLinear(
            ..., bias=False, ...
        )

    def _init_rope(self):
        # Qwen2 只使用标准 RoPE(不像 LLaMA 支持多种缩放方式)
        self.rotary_emb = Qwen2RotaryEmbedding(
            self.head_dim, max_position_embeddings=self.max_position_embeddings,
            base=self.rope_theta,
        )

RoPE 实现

Qwen2 提供了三种 RoPE(与 LLaMA 类似但没有 Llama3 的分段缩放): - Qwen2RotaryEmbedding:标准 RoPE - Qwen2LinearScalingRotaryEmbedding:线性缩放 - Qwen2DynamicNTKScalingRotaryEmbedding:动态 NTK 缩放

去 Padding 版本

class ParallelQwen2AttentionRmPad(ParallelQwen2Attention):
    """与 LLaMA 的 ParallelLlamaAttentionRmPad 逻辑完全相同"""

核心类列表

类名 作用
Qwen2RotaryEmbedding 标准 RoPE
Qwen2LinearScalingRotaryEmbedding 线性缩放 RoPE
Qwen2DynamicNTKScalingRotaryEmbedding 动态 NTK RoPE
ParallelQwen2Attention 标准 TP 注意力(QKV 带 bias)
ParallelQwen2AttentionRmPad 去 padding TP 注意力

与其他模块的关系

  • 被 parallel_decoder.py 使用
  • 使用 parallel_linear.py 的 QKVParallelLinear

小结

Qwen2 注意力的 TP 并行实现与 LLaMA 结构几乎相同,关键区别是 QKV 带 bias。这影响了 checkpoint_utils 中权重加载/保存时需要额外处理 bias 的 TP 切分。