跳转至

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 引入了"潜在压缩"的概念:

  1. Q 可选 LoRA:当 q_lora_rank 不为 None 时,Q 先降维再升维
  2. KV 压缩:K 和 V 共享一个低秩压缩表示 compressed_kv
  3. 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。计算完成后再截断:

if self.q_head_dim != self.v_head_dim:
    attn_output = attn_output[:, :, :, : self.v_head_dim]

核心函数列表

名称 作用
_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 机制。