跳转至

parallel_attention.py — LLaMA 并行注意力

文件路径

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

文件概述

实现 LLaMA 的张量并行(TP)注意力机制,包括标准版本和去 padding 版本。同时包含多种 RoPE 实现(标准、线性缩放、动态 NTK、Llama3)。

关键代码讲解

1. RoPE 旋转位置编码

class LlamaRotaryEmbedding(nn.Module):
    def __init__(self, dim, max_position_embeddings=2048, base=10000):
        # 计算逆频率: 1/(base^(2i/dim))
        inv_freq = 1.0 / (self.base ** (torch.arange(0, self.dim, 2).float() / self.dim))
        self.register_buffer("inv_freq", inv_freq, persistent=False)
        # 预计算 cos/sin 缓存
        self._set_cos_sin_cache(seq_len=max_position_embeddings, ...)

    def _set_cos_sin_cache(self, seq_len, device, dtype):
        t = torch.arange(seq_len, device=device)
        freqs = torch.einsum("i,j->ij", t, self.inv_freq)  # 外积
        emb = torch.cat((freqs, freqs), dim=-1)  # 重复一次
        self.register_buffer("cos_cached", emb.cos())
        self.register_buffer("sin_cached", emb.sin())

还提供了三种扩展版本: - LlamaLinearScalingRotaryEmbedding:线性缩放位置(t = t / scaling_factor) - LlamaDynamicNTKScalingRotaryEmbedding:动态调整 base 频率 - LlamaLlama3ScalingRotaryEmbedding:Llama3 的分段频率缩放(高频不动、低频缩放、中频平滑插值)

2. 标准并行注意力 ParallelLlamaAttention

class ParallelLlamaAttention(nn.Module):
    def __init__(self, config, megatron_config):
        tp_size = mpu.get_tensor_model_parallel_world_size()
        self.num_heads_per_tp = self.num_heads // tp_size
        self.num_key_value_heads_per_tp = self.num_key_value_heads // tp_size

        # QKV 合并为一个 ColumnParallelLinear
        self.qkv_proj = QKVParallelLinear(
            input_size=self.hidden_size,
            num_heads=self.num_heads,
            num_key_value_heads=self.num_key_value_heads,
            head_dim=self.head_dim,
            gather_output=False, ...  # 不 gather,保持 TP 分片
        )
        # O 投影使用 RowParallelLinear(输入已按 TP 切分)
        self.o_proj = tensor_parallel.RowParallelLinear(
            input_size=self.num_heads * self.head_dim,
            output_size=self.hidden_size,
            input_is_parallel=True, ...
        )

    def forward(self, hidden_states, attention_mask, position_ids):
        bsz, q_len, _ = hidden_states.size()
        # 1. QKV 投影(ColumnParallel 内部做 all-gather + 矩阵乘)
        qkv = self.qkv_proj(hidden_states)[0]
        query, key, value = qkv.split([self.q_size, self.k_size, self.v_size], dim=-1)

        # 2. 应用 RoPE
        cos, sin = self.rotary_emb(value, seq_len=kv_seq_len)
        query, key = apply_rotary_pos_emb(query, key, cos, sin, position_ids)

        # 3. GQA: 重复 KV 头以匹配 Q 头数
        key = repeat_kv(key, self.num_key_value_groups)
        value = repeat_kv(value, self.num_key_value_groups)

        # 4. 标准注意力计算
        attn_weights = torch.matmul(query, key.transpose(2, 3)) / math.sqrt(self.head_dim)
        attn_weights = nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32)
        attn_output = torch.matmul(attn_weights, value)

        # 5. O 投影(RowParallel 内部做矩阵乘 + reduce-scatter)
        attn_output = self.o_proj(attn_output)[0]

3. 去 Padding 注意力 ParallelLlamaAttentionRmPad

class ParallelLlamaAttentionRmPad(ParallelLlamaAttention):
    def forward(self, hidden_states, position_ids, sequence_length,
                indices, cu_seqlens, max_seqlen_in_batch):
        total_nnz, _, _ = hidden_states.size()

        # SP: all-gather 已在 ColumnParallel 内部完成
        if self.megatron_config.sequence_parallel:
            total_nnz = total_nnz * mpu.get_tensor_model_parallel_world_size()

        qkv = self.qkv_proj(hidden_states)[0]
        query, key, value = qkv.split([self.q_size, self.k_size, self.v_size], dim=-1)

        # SP: 去除 SP padding,恢复真实 token 数
        if self.megatron_config.sequence_parallel:
            sequence_parallel_pad = total_nnz - cu_seqlens[-1]
            total_nnz = cu_seqlens[-1]
            query, key, value = query[:total_nnz], key[:total_nnz], value[:total_nnz]

        # 使用 flash_attn 的 RoPE(只需 cos/sin 的前半部分)
        cos, sin = self.rotary_emb(value, seq_len=sequence_length)
        cos, sin = cos[:, :cos.shape[1]//2], sin[:, :sin.shape[1]//2]
        query, key = apply_rotary_pos_emb_rmpad_flash(
            query, key, cos, sin, cu_seqlens=cu_seqlens, max_seqlen=max_seqlen_in_batch
        )

        # 使用 flash_attn_varlen_func(支持变长序列)
        attn_output = flash_attn_varlen_func(
            query, key, value,
            cu_seqlens_q=cu_seqlens, cu_seqlens_k=cu_seqlens,
            max_seqlen_q=max_seqlen_in_batch, max_seqlen_k=max_seqlen_in_batch,
            causal=True,
        )

        # SP: 重新 pad 回 SP 对齐长度
        if self.megatron_config.sequence_parallel:
            attn_output = F.pad(attn_output, pad=(0, 0, 0, 0, 0, sequence_parallel_pad))

        attn_output = self.o_proj(attn_output)[0]  # reduce-scatter

4. RoPE 辅助函数

def apply_rotary_pos_emb_rmpad_flash(q, k, cos, sin, cu_seqlens, max_seqlen):
    """使用 flash_attn 的高效 RoPE(支持变长序列)"""
    q_embed = apply_rotary_emb(q, cos, sin, interleaved=False, inplace=False,
                                cu_seqlens=cu_seqlens, max_seqlen=max_seqlen)
    k_embed = apply_rotary_emb(k, cos, sin, ...)
    return q_embed, k_embed

TP 并行模式图

输入 hidden_states: (total_nnz, 1, hidden_size)
         |
    [ColumnParallel: all-gather -> matmul]
         |
  QKV: (total_nnz, 1, (q+k+v)_size_per_tp)
         |
    [split Q, K, V]
         |
    [RoPE + Flash Attention]
         |
  attn_output: (total_nnz, 1, hidden_size_per_tp)
         |
    [RowParallel: matmul -> reduce-scatter]
         |
  output: (total_nnz // sp, 1, hidden_size)

核心类/函数列表

名称 作用
LlamaRotaryEmbedding 标准 RoPE
LlamaLinearScalingRotaryEmbedding 线性缩放 RoPE
LlamaDynamicNTKScalingRotaryEmbedding 动态 NTK RoPE
LlamaLlama3ScalingRotaryEmbedding Llama3 分段 RoPE
ParallelLlamaAttention 标准 TP 注意力
ParallelLlamaAttentionRmPad 去 padding TP 注意力
apply_rotary_pos_emb() 标准 RoPE 应用
apply_rotary_pos_emb_rmpad_flash() flash_attn RoPE
repeat_kv() GQA 的 KV 头重复

与其他模块的关系

  • 被 parallel_decoder.py 使用
  • 使用 parallel_linear.py 的 QKVParallelLinear
  • 使用 flash_attn 库的变长注意力

小结

这个文件实现了 LLaMA 注意力的 Megatron TP 并行版本。QKV 使用 ColumnParallelLinear 按输出维度切分,O 使用 RowParallelLinear 按输入维度切分。去 padding 版本结合 flash_attn_varlen_func 和序列并行,实现了高效的变长序列注意力计算。