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 可以在完整序列上独立计算自己负责的注意力头。