DeepSeekV3.2之DSA

鱿鱼圈 Lv4

论文链接:arxiv.org/pdf/2512.02556

DeepSeek-V3.2 的 DSA(DeepSeek Sparse Attention)详解

面向小白:从「为什么需要 DSA」讲起,结合论文数学公式和 sglang 源码, 最后把 MLA + DSA 串起来,看清 DeepSeek-V3.2 完整的注意力是怎么跑的。

建议先读完同目录的MLA,本文大量复用 MLA 的结论。


img

第 0 部分:一句话先记住

  • MLA 解决的是 「KV 太占显存」:把每个 token 的 KV 压成一个 512 维的小向量 c_kv
  • DSA 解决的是 「序列太长,注意力算不动」:每个 query token 不再看全部历史 token, 而是先用一个轻量打分器(lightning indexer)挑出最相关的 top-k 个 token,只对这 k 个做注意力。

MLA 省的是显存(KV cache 大小),DSA 省的是计算量(attention 的 FLOPs)。 两者正交,DeepSeek-V3.2 把它们叠在一起用。


第 1 部分:为什么需要 DSA —— 长序列的「平方墙」

标准注意力的代价是 序列长度的平方。query 第 个 token 要和前面所有 个 token 算分:

  • 时还能忍。
  • 时, 大了约 1000 倍。decode 阶段每生成一个新 token, 都要把它和前面 12.8 万个 token全部算一遍注意力分 —— 这就是长上下文又慢又贵的根源。

核心观察:对某个 query 来说,前面 12.8 万个 token 里,真正重要的其实只有几千个, 绝大多数 token 的注意力权重接近 0。那何必全算?

DSA 的思路:先花极小的代价(一个很轻的打分器)估计「哪些 token 重要」, 只挑 top-k 个(V3.2 里 通常是 2048)做真正的注意力。 于是计算量从 降到 是常数 → 近似线性


第 2 部分:DSA 的两个部件

DSA = Lightning Indexer(闪电索引器) + Top-k 稀疏注意力

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
                 ┌─────────────────────────────────────┐
query token t ──▶│ Lightning Indexer(轻量打分器) │
│ 对每个历史 token s 算一个分数 I(t,s) │
└──────────────┬──────────────────────┘
│ 得到分数向量 [I(t,1), …, I(t,L)]

┌─────────────────────────────────────┐
│ Top-k 选择:挑出分数最高的 k 个 token │
│ → topk_indices = [s1, s2, …, sk] │
└──────────────┬──────────────────────┘


┌─────────────────────────────────────┐
│ 主注意力(MLA)只对这 k 个 token 算 │
│ → 输出 = Attention(q_t, {K,V}_topk) │
└─────────────────────────────────────┘

关键点:索引器很便宜,主注意力很贵。所以让便宜的索引器去筛选, 让昂贵的主注意力只处理被筛中的少数 token。


第 3 部分:Lightning Indexer 的数学公式

3.1 论文公式

对 query token 和历史 token ,索引分数定义为:

逐项解释(这是理解 DSA 的核心,慢慢看):

符号 含义 在 sglang 里的来源
索引器的「头数」,很小(比主注意力头数少很多) index_n_heads
query token 的第 索引头的查询向量,维度 (小,如 128) wq_b(q_lora)
历史 token 索引 key只有一个头(所有索引头共享) wk(hidden)
点积,衡量 在第 个索引头下有多相关 FP8 矩阵乘
把负相关「砍成 0」,只保留正向关联 kernel 内 relu
query token 给第 个索引头的权重(每个头多重要) weights_proj(hidden)
把所有索引头的贡献加起来 kernel 内 logits_sum

直觉:索引器是一个简化版多头注意力打分。它对每个 (t, s) 对, 在 个小头里各算一个 ReLU 后的相似度,再用 加权求和,得到一个标量分数 。 分数越高 = token 对 query 越重要。

3.2 为什么这么设计「便宜」?

对比一下主注意力(MLA)和索引器的成本:

主注意力(MLA) Lightning Indexer
头数 128(num_attention_heads ,很小(如 64)
key 头数 共享 1 份 512 维 latent 共享 1 份 维(如 128)
精度 bf16 FP8(更快、更省)
每个 token 算什么 完整 softmax 注意力 + 加权 V 只算一个标量分数

所以索引器对全部 个历史 token打分,也比让主注意力算 top-k 个便宜得多。 (论文里索引器的 FLOPs 只占主注意力的很小一部分。)

3.3 Top-k 选择

拿到分数向量 后,选出分数最大的 个下标:

然后主注意力只在 这 k 个 token 上做:

如果历史长度 (短序列),就退化成「全看」,DSA 不起作用也不出错。


第 4 部分:对应 sglang 源码

源码主目录:python/sglang/srt/layers/attention/dsa/dsa_indexer.py,核心类是 Indexer

4.1 索引器的三个投影层(__init__,对应公式里的

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
# dsa_indexer.py  __init__  (约 346–367 行)

self.wq_b = ReplicatedLinear(
self.q_lora_rank, # 输入:Q 的低秩 latent(复用 MLA 的 q_lora!)
self.n_heads * self.head_dim, # 输出:H_I 个索引头 × d_I → q^I_{t,h}
bias=False, ...)

self.wk = ReplicatedLinear(
self.hidden_size, # 输入:原始 hidden
self.head_dim, # 输出:单头索引 key → k^I_s(所有头共享一份)
bias=False, ...)

self.weights_proj = ReplicatedLinear(
self.hidden_size, # 输入:原始 hidden
self.n_heads, # 输出:每个索引头一个标量权重 → w_{t,h}
bias=False, params_dtype=torch.bfloat16, ...)

self.k_norm = LayerNorm(self.head_dim, ...) # 稳定 key
self.rotary_emb = get_rope_wrapper(rope_head_dim, ...) # 索引器也有自己的 RoPE
self.softmax_scale = self.head_dim**-0.5 # 1/√d_I 缩放

逐行对应公式 (DSA-1):

公式符号 代码
wq_b(q_lora) 然后 reshape 成 [L, H_I, d_I]
wk(hidden)k_norm → 单头
weights_proj(hidden)[L, H_I]

注意 wq_b 的输入是 q_lora_rank直接复用了 MLA 把 Q 压成的低秩 latent—— 这是 MLA 和 DSA 的第一个交汇点:索引器不另算 Q,而是蹭 MLA 已经算好的 q_lora

4.2 算 q / k 并加 RoPE(_get_q_k_bf16,约 439–529 行)

1
2
3
4
5
6
7
8
9
query, _ = self.wq_b(q_lora)                       # q^I
query = rearrange(query, "l (h d) -> l h d", d=self.head_dim)
q_rope, _ = torch.split(query, [rope_head_dim, ...], dim=-1)

key, _ = self.wk(x) # k^I(单头)
key = self.k_norm(key)
k_rope, _ = torch.split(key, [rope_head_dim, ...], dim=-1)

q_rope, k_rope = self.rotary_emb(positions, q_rope, k_rope) # 给索引器的 q/k 加位置

索引器有自己独立的一套 RoPE,和主注意力的 RoPE 分开。

4.3 算分数 + ReLU + 加权求和(核心公式 DSA-1)

这一步落在 FP8 kernel 里。看 tilelang_kernel.pyfp8_index 注释(约 244–247 行), 它字面就是公式 (DSA-1):

1
2
3
4
fp8 q @ fp8 k        -> fp32 logits        # q^I_{t,h} · k^I_s(每个头一个点积)
relu(fp32 logits) * q_s (weights) -> fp32 logits # ReLU(...) × w_{t,h}
fp32 logits -> fp32 logits_sum # Σ_h(对索引头求和)
fp32 logits_sum * k_s(e8m0) -> fp32 index_score # 反量化缩放

weights(即 )是在 forward_cuda(约 1393–1498 行)里算的:

1
2
weights = self._get_logits_head_gate(x_for_gate, q_scale)
# 内部:weights_proj(x) * n_heads**-0.5 * q_scale * softmax_scale

也就是 先经过 weights_proj,再乘上 的缩放因子。

4.4 Top-k 选择(公式 DSA-2)

最干净的参考实现是 tilelang 路径 forward_indexer(约 1207–1214 行):

1
2
index_score = fp8_index(q_fp8_partial, weights_partial, k_fp8, k_scale)  # 公式 DSA-1
topk_indices = index_score.topk(min(topk, end_pos), dim=-1)[1].squeeze(0) # 公式 DSA-2

topk(...)[1] 取的是下标(indices),不是分数本身——这正是我们要的「挑哪些 token」。 高性能路径(_get_topk_ragged / _get_topk_paged)用 deep_gemm.fp8_mqa_logits

  • metadata.topk_transform 做同样的事,只是融合得更狠、支持分块防 OOM。

4.5 一个重要优化:短序列直接跳过打分

forward_cuda(约 1362–1383 行):

1
2
if max_kv_len <= self.index_topk:
skip_logits_computation = True # 历史比 k 还短,全看就行,没必要打分

对应公式 (DSA-2) 里 的退化情况。


第 5 部分:把 MLA + DSA 串起来(DeepSeek-V3.2 完整流程)

DeepSeek-V3.2 的每个注意力层,是 MLA 当「主注意力」+ DSA 当「token 筛选器」。 在 sglang 里,索引器被 DeepseekV2AttentionMLAmodels/deepseek_v2.py)持有, 在 forward_absorb_preparedeepseek_common/attention_forward_methods/forward_mla.py)里调用。

5.1 一次 decode 的完整数据流

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
输入 hidden_states  x  (一个新 token)

▼ ① MLA 的 Q 下投影(论文公式 (37))
q_lora = W^DQ · x # Q 压成 1536 维低秩 latent

├──────────────► ② DSA 索引器(复用 q_lora!)
│ q^I = wq_b(q_lora) (DSA-1 的 q)
│ k^I = wk(x) (DSA-1 的 k,写进 index-K cache)
│ w = weights_proj(x) (DSA-1 的权重)
│ I(t,s) = Σ_h w_h·ReLU(q^I_h·k^I_s) ← 对所有历史 token 打分
│ 𝒮 = Top-k(I, k=2048) (DSA-2) → topk_indices
│ │
▼ ▼
③ MLA 主注意力(只在 𝒮 这 2048 个 token 上做)
q_nope, q_pe = split(q)
q_nope ← 吸收 W^UK(论文公式,weight absorption)
q_pe, k_pe ← RoPE
┌─ 从 KV cache 取出 topk_indices 指定的那 2048 个 c_kv(512维 latent)
│ (MLA:num_kv_heads=1 的 MQA,所有头共享一份压缩 KV)
└─ attn = softmax(q · [c_kv; k_pe]ᵀ) · c_kv ← 只对 2048 个算,不是全部!

▼ ④ MLA 输出解压(吸收 W^UV)
out = attn ← 吸收 W^UV → 还原成 v_head_dim → o_proj

5.2 三个关键交汇点

  1. 共享 Q latent:索引器的 wq_b 输入是 MLA 算出来的 q_lora,不重复算 Q 下投影。 (源码:forward_mla.py 里把 q_lora 同时喂给 self.indexer(...)。)

  2. DSA 决定「看谁」,MLA 决定「怎么看」

    • DSA 输出 topk_indices(一个长度 k 的下标列表)。
    • 这个 topk_indices 被传进 MLA 的注意力 kernel(attn_mqa), 让它只读取 KV cache 里这 k 个槽位的 c_kv,跳过其余历史。
    • 于是 MLA 那个 num_kv_heads=1 的 MQA,从「读全部 L 个 c_kv」变成「读 k 个 c_kv」。
  3. 两套独立的 cache

    • MLA 存的是主注意力用的 512 维 c_kvkv_a_proj_with_mqa 产出)。
    • DSA 索引器额外存一份很小的 index-K cachewk 产出,FP8),只给打分用。
    • 两者都按 token 存,靠同一套 topk_indices 对齐。

5.3 复杂度对比(直观感受收益)

设上下文长度 ,top-k

阶段 不用 DSA(只 MLA) 用 DSA(MLA+DSA)
索引器打分 ,很便宜(FP8、单头 K)
主注意力 每步,全看 每步,只看
看 131072 个 token 看 2048 个 token,约省 64×

MLA 让每个 token 的 KV 只占 512 维(省显存); DSA 让每步只对 2048 个 token 算注意力(省计算)。 两者叠加 → DeepSeek-V3.2 在 128K 长上下文下又省显存又快。


第 6 部分:三句话总结

  1. MLA 把每个 token 的 KV 压成一个 512 维 latent c_kv, 并通过权重吸收退化成 num_kv_heads=1 的 MQA —— 解决显存瓶颈。
  2. DSA 用一个 FP8 的轻量 lightning indexer,按公式 给每个历史 token 打分, 选 top-k(2048)个,让主注意力只算这一小撮 —— 解决长序列的 计算瓶颈。
  3. DeepSeek-V3.2 = MLA + DSA:索引器复用 MLA 的 q_lora 打分得到 topk_indices, MLA 的 MQA 再只读取这 k 个槽位的压缩 KV 做注意力 —— DSA 管「看谁」,MLA 管「怎么省着看」

附:关键源码索引

内容 文件 : 行
索引器投影层 wq_b/wk/weights_proj dsa/dsa_indexer.py:346-367
算 q/k + 索引器 RoPE dsa/dsa_indexer.py:439-529
公式 (DSA-1) 的 FP8 kernel 注释 dsa/tilelang_kernel.py:244-247
Top-k 选择(参考实现) dsa/dsa_indexer.py:1207-1214
短序列跳过打分优化 dsa/dsa_indexer.py:1362-1383
MLA 里实例化 Indexer models/deepseek_v2.py:1576-1636
MLA 调用 indexer 拿 topk_indices deepseek_common/attention_forward_methods/forward_mla.py:262-292
  • 标题: DeepSeekV3.2之DSA
  • 作者: 鱿鱼圈
  • 创建于 : 2026-04-21 23:50:00
  • 更新于 : 2026-06-30 22:31:19
  • 链接: https://yuyanqi.com/2026/04/21/DeepSeekV3.2之DSA/
  • 版权声明: 本文章采用 CC BY-NC-SA 4.0 进行许可。
评论