Compressed Sparse Attention (压缩稀疏注意力)

image

压缩稀疏注意力 (Compressed Sparse Attention)

CSA(压缩稀疏注意力)的核心架构如上图所示,它首先将每 \(m\) 个 token 的 KV 缓存压缩为一个条目,然后应用 DeepSeek 稀疏注意力以进行进一步加速。

压缩键值条目 (Compressed Key-Value Entries)。\(H \in \mathbb{R}^{n \times d}\) 为输入隐藏状态序列,其中 \(n\) 是序列长度,\(d\) 是隐藏层大小。CSA 首先计算两系列 KV 条目 \(C^a, C^b \in \mathbb{R}^{n \times c}\) 及其对应的压缩权重 \(Z^a, Z^b \in \mathbb{R}^{n \times c}\),其中 \(c\) 是头维度 (head dimension):

\[C^a = H \cdot W^{aKV}, \quad C^b = H \cdot W^{bKV}, \tag{1} \]

\[Z^a = H \cdot W^{aZ}, \quad Z^b = H \cdot W^{bZ}, \tag{2} \]

其中 \(W^{aKV}, W^{bKV}, W^{aZ}, W^{bZ} \in \mathbb{R}^{d \times c}\) 是可训练参数。接下来,\(C^a\)\(C^b\) 中的每 \(m\) 个 KV 条目将根据其压缩权重和可学习的位置偏差 \(B^a, B^b \in \mathbb{R}^{m \times c}\) 被压缩成一个条目,生成 \(C^{\text{Comp}} \in \mathbb{R}^{\frac{n}{m} \times c}\)。每个压缩条目 \(C_i^{\text{Comp}} \in \mathbb{R}^c\) 通过下式计算:

\[[S^a_{mi:m(i+1)-1}; S^b_{m(i-1):mi-1}] = \text{Softmax}_{\text{row}}([Z^a_{mi:m(i+1)-1} + B^a; Z^b_{m(i-1):mi-1} + B^b]), \tag{3} \]

\[C_i^{\text{Comp}} = \sum_{j=mi}^{m(i+1)-1} S_j^a \odot C_j^a + \sum_{j=m(i-1)}^{mi-1} S_j^b \odot C_j^b, \tag{4} \]

其中 \(\odot\) 表示哈达玛积 (Hadamard product);\(\text{Softmax}_{\text{row}}(\cdot)\) 表示沿行维度的 softmax 操作,它对来自 \(Z^a\)\(Z^b\) 的总共 \(2m\) 个元素进行归一化。当 \(i=0\) 时,\(Z^b_{m(i-1):mi-1}\) 用负无穷填充,\(C^b_{m(i-1):mi-1}\) 用零填充。注意每个 \(C_i^{\text{Comp}}\) 派生自 \(2m\) 个 KV 条目,但用于 \(C_i^{\text{Comp}}\)\(C^b\) 索引与用于 \(C_{i-1}^{\text{Comp}}\)\(C^a\) 索引是重叠的。因此,CSA 实际上将序列长度压缩到了原来的 \(\frac{1}{m}\) 倍。

用于稀疏选择的闪电索引器 (Lightning Indexer for Sparse Selection)。 在获得压缩后的 KV 条目 \(C^{\text{Comp}}\) 后,CSA 应用 DSA 策略来选择前 k 个压缩 KV 条目用于核心注意力。首先,CSA 执行与获取 \(C^{\text{Comp}}\) 相同的压缩操作来获得压缩索引键 \(K^{\text{IComp}} \in \mathbb{R}^{\frac{n}{m} \times c^I}\),其中 \(c^I\) 是索引头维度。然后,对于查询 token \(t\),我们以低秩方式生成索引查询 \(\{\mathbf{q}_{t,1}^I, \mathbf{q}_{t,2}^I, ..., \mathbf{q}_{t,n_h^I}^I\}\)

\[\mathbf{c}_t^Q = \mathbf{h}_t \cdot W^{DQ}, \tag{5} \]

\[[\mathbf{q}_{t,1}^I; \mathbf{q}_{t,2}^I; ...; \mathbf{q}_{t,n_h^I}^I] = \mathbf{c}_t^Q \cdot W^{UQ}, \tag{6} \]

其中 \(\mathbf{h}_t \in \mathbb{R}^d\) 是查询 token \(t\) 的输入隐藏状态;\(\mathbf{c}_t^Q \in \mathbb{R}^{d_c}\) 是查询的压缩潜在向量;\(d_c\) 表示查询压缩维度;\(n_h^I\) 表示索引查询头的数量;\(W^{DQ} \in \mathbb{R}^{d \times d_c}\)\(W^{UQ} \in \mathbb{R}^{d_c \times c^I n_h^I}\) 分别是索引查询的下投影和上投影矩阵。接下来,查询 token \(t\) 与前一个压缩块 \(s\) (\(s < \text{Floor}(\frac{t}{m})\)) 之间的索引分数 \(I_{t,s} \in \mathbb{R}\) 计算如下:

\[[w_{t,1}^I; w_{t,2}^I; ...; w_{t,n_h^I}^I] = \mathbf{w}_t^I = \mathbf{h}_t \cdot W^w, \tag{7} \]

\[I_{t,s} = \sum_{h=1}^{n_h^I} w_{t,h}^I \cdot \text{ReLU}\left(\mathbf{q}_{t,h}^I \cdot K_s^{\text{IComp}}\right), \tag{8} \]

其中 \(W^w \in \mathbb{R}^{d \times n_h^I}\) 是一个可学习矩阵;\(w_{t,h}^I \in \mathbb{R}\) 是第 \(h\) 个索引头的权重。对于查询 token \(t\),给定其索引分数 \(I_{t,:}\),我们采用 top-k 选择器有选择地保留一部分压缩 KV 条目 \(C_t^{\text{SprsComp}}\) 用于后续的核心注意力:

\[C_t^{\text{SprsComp}} = \left\{ C_s^{\text{Comp}} \mid I_{t,s} \in \text{Top-k}(I_{t,:}) \right\}. \tag{9} \]

共享键值多查询注意力(Shared Key-Value MQA)。 在选择了稀疏 KV 条目后,CSA 随后以多查询注意力(MQA)(Shazeer,2019)的方式执行核心注意力,其中 \(C_t^{\text{SprsComp}}\) 中的每个压缩 KV 条目同时充当注意力键和值。具体来说,对于查询 token \(t\),我们首先从压缩潜在向量 \(\mathbf{c}_t^Q\) 生成注意力查询 \(\{\mathbf{q}_{t,1}; \mathbf{q}_{t,2}; ...; \mathbf{q}_{t,n_h}\}\)

\[[\mathbf{q}_{t,1}; \mathbf{q}_{t,2}; ...; \mathbf{q}_{t,n_h}] = \mathbf{q}_t = \mathbf{c}_t^Q \cdot W^{UQ}, \tag{10} \]

其中 \(n_h\) 表示查询头的数量;\(W^{UQ} \in \mathbb{R}^{d_c \times c n_h}\) 是查询的上投影矩阵。请注意,潜在查询向量 \(\mathbf{c}_t^Q\) 与用于索引器查询的向量是共享的。接下来,我们在 \(\{\mathbf{q}_{t,i}\}\)\(C_t^{\text{SprsComp}}\) 上执行 MQA:

\[\mathbf{o}_{t,i} = \text{CoreAttn}\left(\text{query}=\mathbf{q}_{t,i}, \text{key}=C_t^{\text{SprsComp}}, \text{value}=C_t^{\text{SprsComp}}\right), \tag{11} \]

其中 \(\mathbf{o}_{t,i} \in \mathbb{R}^c\) 是第 \(t\) 个 token 处第 \(i\) 个头的核心注意力输出;\(\text{CoreAttn}(\cdot)\) 表示核心注意力操作。

posted @ 2026-07-31 17:01  幽赏未已高谈转清  阅读(0)  评论(0)    收藏  举报