kimi_vl.py — KimiVL (DeepSeek-V3 架构) 注意力适配¶
文件路径¶
verl/models/transformers/kimi_vl.py
文件概述¶
为 KimiVL 模型提供支持 Ulysses 序列并行的注意力替换函数。KimiVL 基于 DeepSeek-V3 架构,其注意力机制采用 MLA (Multi-Latent Attention),与标准 MHA/GQA 有很大区别,因此需要特殊处理。
关键代码讲解¶
MLA (Multi-Latent Attention) 简介¶
传统的 MHA/GQA 直接投影得到 Q、K、V,而 MLA 引入了"潜在压缩"的概念:
- Q 可选 LoRA:当
q_lora_rank不为 None 时,Q 先降维再升维 - KV 压缩:K 和 V 共享一个低秩压缩表示
compressed_kv - Q/K 拆分为 nope + pe:分别用于无位置编码和有位置编码的部分
Q 投影和 KV 压缩¶
# Q 的两种路径
if self.q_lora_rank is None:
q = self.q_proj(hidden_states)
else:
q = self.q_b_proj(self.q_a_layernorm(self.q_a_proj(hidden_states)))
q = q.view(bsz, q_len, self.num_heads, self.q_head_dim).transpose(1, 2)
# KV 共享压缩
compressed_kv = self.kv_a_proj_with_mqa(hidden_states)
compressed_kv, k_pe = torch.split(compressed_kv, [self.kv_lora_rank, self.qk_rope_head_dim], dim=-1)
k_pe = k_pe.view(bsz, q_len, 1, self.qk_rope_head_dim).transpose(1, 2)
kv = self.kv_b_proj(self.kv_a_layernorm(compressed_kv))
.view(bsz, q_len, self.num_heads, self.qk_nope_head_dim + self.v_head_dim)
.transpose(1, 2)
k_nope, value_states = torch.split(kv, [self.qk_nope_head_dim, self.v_head_dim], dim=-1)
关键理解:
- k_pe 只有 1 个头(MQA 风格),用于位置编码
- k_nope 和 value_states 有 num_heads 个头,从压缩表示解压得到
- Q 被拆分为 q_nope(不加位置编码的部分)和 q_pe(加位置编码的部分)
Ulysses SP 的特殊处理¶
ulysses_sp_size = get_ulysses_sequence_parallel_world_size()
if ulysses_sp_size > 1:
validate_ulysses_config(self.num_heads, ulysses_sp_size)
num_key_value_groups = self.config.num_attention_heads // self.config.num_key_value_heads
k_pe = repeat_kv(k_pe, ulysses_sp_size) # k_pe 只有1个头,需要重复到 SP 大小
k_nope = repeat_kv(k_nope, num_key_value_groups)
value_states = repeat_kv(value_states, num_key_value_groups)
q = gather_seq_scatter_heads(q, seq_dim=2, head_dim=1)
k_pe = gather_seq_scatter_heads(k_pe, seq_dim=2, head_dim=1)
k_nope = gather_seq_scatter_heads(k_nope, seq_dim=2, head_dim=1)
value_states = gather_seq_scatter_heads(value_states, seq_dim=2, head_dim=1)
MLA 的特殊之处:k_pe 是 MQA(只有 1 个头),需要先 repeat 到 ulysses_sp_size 倍,AlltoAll 之后才能正确切分头。
拼接最终的 Q 和 K¶
q_nope, q_pe = torch.split(q, [self.qk_nope_head_dim, self.qk_rope_head_dim], dim=-1)
cos, sin = self.rotary_emb(value_states, seq_len=full_q_len)
q_pe, k_pe = apply_rotary_pos_emb(q_pe, k_pe, cos, sin, position_ids)
# 拼接 nope 和 pe 部分形成完整的 Q 和 K
query_states = k_pe.new_empty(bsz, self.num_heads // ulysses_sp_size, full_q_len, self.q_head_dim)
query_states[:, :, :, : self.qk_nope_head_dim] = q_nope
query_states[:, :, :, self.qk_nope_head_dim :] = q_pe
key_states = k_pe.new_empty(bsz, self.num_heads // ulysses_sp_size, full_q_len, self.q_head_dim)
key_states[:, :, :, : self.qk_nope_head_dim] = k_nope
key_states[:, :, :, self.qk_nope_head_dim :] = k_pe
V 的 padding 处理¶
# V 的 head_dim 可能与 Q 不同,需要 padding
if self.q_head_dim != self.v_head_dim:
value_states = F.pad(value_states, [0, self.q_head_dim - self.v_head_dim])
Flash Attention 要求 Q、K、V 的 head_dim 相同,所以对 V 做 zero-padding。计算完成后再截断:
核心函数列表¶
| 名称 | 作用 |
|---|---|
_ulysses_flash_attn_forward() |
KimiVL MLA 注意力的 Ulysses SP 替换函数 |
rotate_half() |
RoPE 辅助函数 |
apply_rotary_pos_emb() |
自定义的 RoPE 应用函数(适配 MLA 格式) |
repeat_kv() |
GQA/MQA KV 头重复函数 |
与其他模块的关系¶
- 被
monkey_patch.py中的apply_monkey_patch()调用 - 自定义了
apply_rotary_pos_emb(),因为 MLA 的 RoPE 格式与标准 LLaMA 不同 - 依赖
verl.utils.ulysses的 AlltoAll 原语
小结¶
KimiVL 的注意力适配是所有模型中最复杂的,因为 MLA 架构将 Q/K 拆分为有无位置编码的两部分,KV 使用低秩压缩共享表示。理解这个文件需要先了解 DeepSeek-V2/V3 的 MLA 机制。