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 和序列并行,实现了高效的变长序列注意力计算。