跳转至

llama.py — LLaMA 注意力层 Ulysses SP 适配

文件路径

verl/models/transformers/llama.py

文件概述

为 LLaMA 模型的注意力层提供支持 Ulysses 序列并行 (SP) 的替换函数。根据 transformers 库版本的不同,提供了两个版本的注意力前向函数。这些函数会在 monkey_patch.py 中被用来替换 HuggingFace LLaMA 模型原始的注意力计算。

关键代码讲解

版本分支

该文件提供两个函数,分别适配不同版本的 transformers: - llama_flash_attn_forward():适用于 transformers 4.45.0 ~ 4.47.1 - llama_attn_forward():适用于 transformers 4.48.0+(API 有变化)

两者核心逻辑相同,区别在于接口签名和调用方式。

Ulysses SP 核心流程

以 llama_flash_attn_forward 为例,核心流程分为三步:

第一步:QKV 投影 + AlltoAll 前置通信

query_states = self.q_proj(hidden_states)
key_states = self.k_proj(hidden_states)
value_states = self.v_proj(hidden_states)

# reshape 为 (bsz, n_head, seq_len/n, head_dim)
query_states = query_states.view(bsz, q_len, self.num_heads, self.head_dim).transpose(1, 2)
# ...

########## AlltoAll for Ulysses ##########
ulysses_sp_size = get_ulysses_sequence_parallel_world_size()

if ulysses_sp_size > 1:
    validate_ulysses_config(self.num_heads, ulysses_sp_size)
    # (bsz, n_head, seq_len/n, head_dim) -> (bsz, n_head/n, seq_len, head_dim)
    query_states = gather_seq_scatter_heads(query_states, seq_dim=2, head_dim=1)
    key_states = gather_seq_scatter_heads(key_states, seq_dim=2, head_dim=1)
    value_states = gather_seq_scatter_heads(value_states, seq_dim=2, head_dim=1)

关键理解:每个 GPU 原本只持有 seq_len/n 长度的序列片段,但拥有所有注意力头。AlltoAll 通信后,每个 GPU 持有完整的 seq_len 序列,但只拥有 n_head/n 个注意力头。这样每个 GPU 可以独立完成自己负责的注意力头的计算。

第二步:RoPE + Flash Attention

cos, sin = position_embeddings
query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin)

attn_output = _flash_attention_forward(
    query_states, key_states, value_states,
    attention_mask, full_q_len,
    position_ids=position_ids,
    dropout=dropout_rate,
    sliding_window=getattr(self, "sliding_window", None),
    use_top_left_mask=flash_attn_supports_top_left_mask(),
    is_causal=self.is_causal,
)

注意使用的是 full_q_len(聚合后的完整序列长度),而非原始的 q_len。

第三步:AlltoAll 后置通信 + 输出投影

attn_output = attn_output.reshape(bsz, full_q_len, -1, self.head_dim).contiguous()
########## AlltoAll for Ulysses ##########
if ulysses_sp_size > 1:
    attn_output = gather_heads_scatter_seq(attn_output, seq_dim=1, head_dim=2)
attn_output = attn_output.reshape(bsz, q_len, -1).contiguous()
attn_output = self.o_proj(attn_output)

反向 AlltoAll 将注意力头聚合回来,同时将序列维度切分回去,恢复到每个 GPU 只持有 seq_len/n 的状态。

4.48+ 版本的主要区别

def llama_attn_forward(
    self,
    hidden_states: torch.Tensor,
    position_embeddings: tuple[torch.Tensor, torch.Tensor],  # 变为必选参数
    attention_mask: Optional[torch.Tensor],
    past_key_value: Optional[Cache] = None,
    cache_position: Optional[torch.LongTensor] = None,
    **kwargs,
):
    # 使用 ALL_ATTENTION_FUNCTIONS 注册表选择注意力实现
    from transformers.modeling_utils import ALL_ATTENTION_FUNCTIONS
    attention_interface = ALL_ATTENTION_FUNCTIONS[self.config._attn_implementation]

    attn_output, attn_weights = attention_interface(
        self, query_states, key_states, value_states, attention_mask,
        dropout=0.0 if not self.training else self.attention_dropout,
        scaling=self.scaling, **kwargs,
    )

4.48+ 版本将 position_embeddings 改为必选参数,并引入了 ALL_ATTENTION_FUNCTIONS 注册表机制。

核心函数列表

函数名 作用
llama_flash_attn_forward() transformers 4.45-4.47 的 LLaMA 注意力替换函数
llama_attn_forward() transformers 4.48+ 的 LLaMA 注意力替换函数

与其他模块的关系

  • 被 monkey_patch.py 中的 apply_monkey_patch() 调用,用于替换 LLaMA 的注意力层
  • 依赖 verl.utils.ulysses 提供的 AlltoAll 通信原语(gather_seq_scatter_heads, gather_heads_scatter_seq)
  • 与 qwen2.py 结构几乎相同,区别仅在于 sliding window 支持

小结

这个文件是 Ulysses 序列并行在 LLaMA 模型上的具体实现。核心思路是在注意力计算前后各做一次 AlltoAll 通信,将"按序列切分"转换为"按注意力头切分",使得每个 GPU 可以在完整序列上独立计算自己负责的注意力头。