对话越长,注意力越慢——因为每生成一个字,都要和前面所有字算一遍。

DSA(DeepSeek Sparse Attention)的办法是先粗筛,再精算

  1. 用一个很便宜的打分器扫一遍全部历史,给每个历史 token 打一个分;

  2. 只留分数最高的 2048 个

  3. 昂贵的主注意力只算这 2048 个,其余全部跳过。

这样主注意力的开销就和对话长度无关了,永远是 2048。

打分器为什么开销小:



主 MLA attention

indexer 打分器

倍数

头数 x 每头维度

128 x 576

64 x 128

计算量小很多

精度

bf16

FP8

一半

每 token 每层 cache

576 维 bf16 = 1152 字节

132 字节

小 8.7 倍

关键点:打分器只需要 k,不需要 v(它只用来排序,不参与加权求和),所以它的 cache 只存一个向量。代码注释写得很直白——deepseek_v2.py:635:"Only has one vector instead of K + V"。

一层里的数据是怎么流的

1. 每层有两套 KV cache,互不相干

  • 主 MLA 的:每 token 576 维 bf16

  • indexer 的:每 token 132 字节 uint8(DeepseekV32IndexerCachedeepseek_v2.py:616

2. 打分器和主注意力之间只通过一块共享 buffer 通信

# deepseek_v2.py:1382,全模型只分配一次
topk_indices_buffer = torch.empty(
    vllm_config.scheduler_config.max_num_batched_tokens,   # 行:每个 query token 一行
    topk_tokens,                                            # 列:2048
    dtype=torch.int32,
)

打分器没有返回值,它把"该看哪 2048 个"直接写进这张表;主注意力再从表里读。-1 表示"这个位置是空的"。

调用顺序在 vllm/model_executor/layers/mla.py:206

if self.indexer and self.is_sparse and not self.skip_topk:
    self.indexer(hidden_states, q_c, positions, self.indexer_rope_emb)   # 只写 buffer
attn_out = self.mla_attn(q, kv_c_normed, k_pe, ...)                      # 读 buffer

注意打分器复用了主 MLA 已经算好的 q_c(1536 维压缩向量),不重新从 hidden_states 算,省一次大 GEMM。

3. 主注意力有三条路,不是永远走稀疏

路线

什么时候走

为什么

A. dense MHA

序列长度 ≤ 2048

反正 2048 个全都会被选中,稀疏毫无意义,直接走普通 MLA

B. masked MHA

序列 > 2048 的 prefill

用索引表造一张稀疏 mask 跑 MHA

C. sparse MQA

decode

flash_mla_sparse_fwd 只读被选中的 2048 个 slot

判断代码在 sparse_mla_attention.py:801

if force_dense or (prefill_max_seq_len <= topk_tokens and not force_masked):
    return super().forward_mha(...)      # 退化成普通 dense MLA

sparse 场景下 prefill 也可能走"矩阵吸收"的 MQA 路径(mla_attention.py:757-779),因为 top-k 把有效的 KV 长度截断到 2048 了,"解压整条 cache"的开销消失,吸收反而更划算。

Indexer内部到底做了什么

六步,对应 deepseek_v2.py:726-819Indexer.forward

① 两个投影

self.wq_b = ReplicatedLinear(1536, 128 * 64)              # q:64 个头,每头 128 维
self.wk_weights_proj = MergedColumnParallelLinear(
    7168, [128, 64])                                       # 一次 GEMM 出 k 和 weights

k 只有 128 维,64 个头共用同一个 k——所以打分是 MQA 而不是 MHA。weights 是每个头一个标量权重。

这两个投影都不做张量并行ReplicatedLinear / disable_tp=True,注释:"no tensor parallel, just replicated"),每个 GPU 都算完整的打分器。省了通信,代价是 TP 越大冗余越多。

② k 过 LayerNormk_norm

③ RoPE 只加在前 64 维上,q 和 k 一起做。后 64 维不加位置编码。

④ q 量化成 FP8,每 128 个数一个 scale。然后有个小技巧——把所有系数一次性乘进 weights

weights = weights * q_scale * self.softmax_scale * self.n_head_scale

这样后面的打分内核里就不用再乘任何系数了。CUDA 上 ③④ 还会被融合成一个 triton kernel(fused_indexer_q_rope_quant)。

⑤ k 的量化和写 cache 融合成一个算子

ops.indexer_k_quant_and_cache(k, kv_cache, slot_mapping, quant_block_size, scale_fmt)

代码注释:we only quant q here since k quant is fused with cache insertion

⑥ 清 buffer,然后算分数 + 选 top-k

分数的定义是把 64 个头的点积加权求和塌成一个数:

logits[t, s] = weights[t,0] * (q[t,0] · k[s]) + ... + weights[t,63] * (q[t,63] · k[s])

t 是当前的 query token,s 是某个历史 token,结果就是"t 对 s 的关注度"。

prefill 和 decode 用完全不同的内核:



prefill

decode

取历史 k

cp_gather_indexer_k_quant_cache 分块 gather 到固定 workspace

不 gather,直接读 paged cache

算分数

fp8_fp4_mqa_logits

fp8_fp4_paged_mqa_logits

选 top-k

top_k_per_row_prefill

cooperative_topk / persistent_topk / top_k_per_row_decode(按 SM 版本和行数三选一)

还有一层省算:不是每层都打分

deepseek_v2.py:1095-1119

if _index_topk_pattern is None:
    _skip_topk = (max(layer_id - _index_skip_topk_offset + 1, 0) % _index_topk_freq != 0)
elif 0 <= layer_id < len(_index_topk_pattern):
    _skip_topk = _index_topk_pattern[layer_id] == "S"

被标记 skip 的层Indexer 对象都不创建self.indexer = None),直接复用前面某层算出来的索引表。相邻层关注的 token 本来就高度重合,所以可以共享。

MTP / nextn 层是例外:永远创建完整打分器,在 draft 第 0 步算索引,之后靠 set_skip_topk 运行时切换。代码注释里有个警告——MTP 层绝不能一开始就 skip_topk=True,否则 draft 会去读一块从没被写过的 buffer。

几个容易踩的点

  • 两级用的是不同的 RoPE 实例:主 MLA 用 self.rotary_emb,打分器用单独的 self.indexer_rope_emb,后者的 is_neox_styleindexer_rope_interleave 单独控制,可以和主 RoPE 不一致。

  • 有个早退优化sparse_attn_indexer.py:407-424):如果主 MLA 这一批决定走 dense MHA(根本不会消费索引表),打分器直接返回,连清 buffer 都省掉。原因是打分器和主 MLA 用各自独立的 decode 阈值,可能对同一批 token 分类不同,只有主 MLA 那边知道索引会不会被用。

  • decode 时索引要做一次转换:打分器输出的是请求内的相对位置,kernel 需要的是 paged cache 里的物理 slot,中间有一步 triton_convert_req_index_to_global_indexflashmla_sparse.py:597)。

  • DCP 下 top-k 要跨 rank 归并_merge_dcp_topk_global):只交换各 rank 的本地 top-k 候选,不交换整个分数矩阵。正确性论证在 docstring 里——全局 top-k 里的 token 必然也在它所属 rank 的本地 top-k 里。这条路径只有 CuteDSL 实现,没有 PyTorch fallback。