Attention Sink、Learnable Sink 与 dsink:从输出含义到反向梯度
讨论整理:2026-09-20。 关联论文:Efficient Streaming Language Models with Attention Sinks,ICLR 2024,本文参考 v4。 范围:记录本次概念澄清和本地源码阅读;不是论文全部实验的精读,也不是 TileLang 实现或 GPU 验收报告。 相关笔记:Sparse Attention、DSA 算法流程。
1. 先分清 score、概率和 attention 输出
对当前 query 的一个 head,attention 的数据流是:
Q 与可见 K 做点积并缩放
↓
scores
↓ softmax
每个可见 token 的概率
↓ 对 V 加权求和
当前 head 的输出向量记 query 位置为
attention 输出是:
Score 是匹配分数,概率是取信息的权重,输出才是汇总得到的特征向量。 输出不是最终的词表概率。
例如,一个 head 读取两个 value:
V₁ = [2, 0]
V₂ = [0, 4]
权重 = [0.75, 0.25]
O = 0.75 × V₁ + 0.25 × V₂ = [1.5, 1]代码中 out 的形状为 [T, Hq, D]:每个 packed query、每个 query head 都产生一个长度为
2. Sink 不是“选择重要 token”的同义词
论文观察到:开头 token 即使语义不重要,也可能获得很高 attention,成为 attention sinks。StreamingLLM 主方案在训练后的推理阶段保留开头少量 KV 与最近窗口,丢弃中间历史,无需微调;这是利用已有 sink 现象的稀疏访问方式,不是对所有历史做内容重要性排序。论文 §3.1–3.2
完整历史:[开头] [大量中间历史] [最近窗口]
保留 KV:[开头] [最近窗口]它支持持续流式生成,但不会使被丢弃的历史仍然可访问。作者仓库 FAQ
下面是解释归一化影响的假设例子,不是论文测量结果:
原概率:开头 sink 0.80,A 0.15,B 0.05
删掉开头 token、保持 A/B logits 不变后:A 0.75,B 0.25移除 sink 后,剩余概率必须重新归一化。这说明“语义不重要”不等于“删除后不影响计算”。
在当前代码讨论中,两个职责应分开:
block_indices描述稀疏选择的块;窗口、因果关系和样本边界进一步约束可见 token。learnable_sink影响可见 token 的归一化和输出幅度,不执行 Top-k 或选择 block。
3. 论文的 sink token 与代码的标量 sink
论文中的真实 sink token 有 K/V,value 不被规定为零。论文 §3.3 还比较了分母加固定
当前 Mimikyu 代码采用的是显式标量:每个 query head 有一个可训练 logit
| 对象 | 在本次讨论中的含义 |
|---|---|
| 真实 sink token | 序列中吸收较多 attention 的位置,有真实 K/V |
learnable_sink[h] | 当前实现的参数 |
| 当前 query、head 分给虚拟 sink 的概率 | |
sink_s[i,h] | 当前实现中真实 token 的总概率,也就是输出缩放比例 |
dsink[h] | 损失对参数 |
这里仅建立机制和公式上的联系,不据此推断 Mimikyu 实现的设计来源。
4. Learnable sink 为什么等价于输出缩放
对固定的 query 和 head,暂时省略下标,定义:
没有显式 sink 时,输出为
真实 token 总共获得的概率为:
由于 sink value 为零:
因此,本实现可以先计算普通 sparse attention,再缩放输出;数学上等价于在 softmax 分母增加
延续第一节的例子,若
原始概率:A 0.75,B 0.25
新概率: A 0.15,B 0.05,零值 sink 0.80
原输出: [1.5, 1]
新输出: [0.3, 0.2] = 0.2 × 原输出需要记住:
- QK scores 本身没有被 sink 修改;改变的是归一化分母、真实 token 概率及输出。
- 真实 token 之间的相对概率保持不变。
按 head 共享,但 随 query 改变,因此 不是固定的 head 权重。 learnable_sink=0仍贡献;当前 Python 接口用 None表示不启用。- 固定
时,分母形式对应 Zero Sink;将其变为可训练 是数学推广,并不等于训练一个真实 sink token。
5. 为什么这不是“中和不同 head 的 score”
每个 head 在自己可见的 token 之间独立做 softmax,不是在不同 head 之间共享一个概率预算。
这个实现允许某个 head 在当前 query 上减少从真实 token 取回的信息。它不比较不同 head 的 score,不把它们拉到相同水平,也不意味着该 head 永久不重要。
“允许少输出信息”是理解零值 sink 的功能性解释。是否真的改善训练或推理质量,需要实验;不能仅从公式推导质量收益。
6. Head 输出变小,对后续有什么影响
标准多头 attention 会拼接各 head 输出,再经过输出投影。按 head 分块后可以写成:
若仅 head 1 缩放为
用常见残差结构示意,省略归一化和其他操作:
Sink 改变的是 attention 分支提供的一份更新,原表示仍通过残差路径保留。之后的 MLP、后续层和词表投影会继续处理这个变化。
因此,减弱某个 head 不等于所有词的 logits 或概率一起变小;特征与投影有正负,后续还有非线性运算。具体预测怎样变化由整个网络决定。
7. dsink 是什么,公式怎样得到
dsink 是损失 learnable_sink 的梯度:
求
记上游输出梯度为
同一
对应代码:
delta = (dout.float() * out.float()).sum(dim=-1) # [T, Hq]
dsink = -(delta * (1.0 - sink_s)).sum(dim=0) # [Hq]
dsink = dsink.to(sink_orig_dtype)沿 dout 中,不应额外随意除以 token 数。
用梯度下降解释符号:若某个 query 的
两个易错点:
- 使用
out_new时不再额外乘;若用 out_orig,公式必须包含。 - 本实现的 sink logit 没有乘 softmax scale,所以
dsink不再额外乘。
8. 为什么 dQ/dK/dV 也能正确包含 sink 的影响
真实 token 的 score 梯度仍是:
因此,现有 backward 若用包含 sink 的 dsink 是新增参数的梯度,不是给 dQ/dK/dV 再做一次缩放修正。
9. 本地源码核对记录
Mimikyu 检查版本:83c7cf172。以下路径相对本机仓库 /Users/kuangjux/codes/mimikyu-dsa,行号是本次检查时的定位,不是稳定 API。
| 源码 | 位置与确认内容 |
|---|---|
mmq/mmq/modules/block/memory_optimizer/qkv_attn/ring_sparse_varlen_attn.py | 125–143 行:logaddexp 计算新 LSE,保存 FP32 sink_s,缩放输出 |
| 同上 | 229–230 行:recompute 将重算的普通输出乘上保存的 sink_s |
| 同上 | 272–314 行:将 sink-adjusted out/LSE 传给 backward,单独计算 dsink |
| 同上 | 163–167 行:返回给 teacher 的是 raw LSE;ctx 中用于主 attention backward 的是新 LSE |
mmq/mmq/modules/attention/allgather_ring_sparse_attn.py | 394–415 行:将 out/LSE 传入 dsa_bwd |
mmq_kernels/mmq_kernels/triton_kernels/dsa/dsa_backward.py | 1651 行:底层已有 delta = preprocess(out, dout, head_dim) |
mmq/mmq/modules/attention/cp_attention.py | 147–151 行:reduce-scatter 处理 dK/dV,不包括上层单独计算的 dsink |
mmq/mmq/modules/block/submodules/mixer.py | 1717、1775 行:提取 dsink 并作为 sink 输入对应的梯度返回 |
分布式边界: 该 attention 函数只对本 rank 的 query 归约 dsink。共享参数需要合并各 rank 的梯度贡献,但本次未验证训练 reducer 的完整覆盖,不能断言全局同步已正确或遗漏,也不应未经核对就增加一次 all-reduce。
数值边界: 前向将 FP32 sink_s 转成输出 dtype 后相乘;反向把已有的 out/dout 转成 FP32 计算。上面的数学等价不意味着 BF16 实现与全 FP32 或融合实现逐位一致。
10. 下一次讨论:如何接到 TileLang
本次仅记录已经检查出的接口条件和候选方向,尚未修改或验证 kernel。
当前仓库 /Users/kuangjux/codes/welm-sparse-attention:
kernels/blackwell/sparse-attn/tilelang_sparse_fwd.py的前向输出已包含 sink,但写出的Lse是真实 token 的 LSE。backward 所需LseTotal必须通过logaddexp(lse_real, sink)得到,或修改前向接口直接提供;不能混用。tilelang_sparse_bwd.py已接收 FP32Delta和LseTotal,现有公式可表达 sink-aware dQ/dK/dV,但当前没有dsink输出。- 可复用已计算的
Delta,按 query tile 生成 FP32 sink 梯度 partial,再归约成[Hq]。这是一种候选方案,具体融合位置和接口留待下一次讨论。 - 候选概率表达式为
exp(sink - lse_total),可避免1 - sink_s在sink_s接近 1 时的相减精度损失;替换后需要检查数值差异。 - 当前 fused backward 按 KV block 的 reverse CSR 分工,同一 query 可出现在多个 block。不能在每次遇到 query 时都累加完整的
-Delta * p_sink,否则会重复计数。 - 后续需明确归约所有权、参数 dtype、workspace、分布式同步与验证用例;空支持集、极端 sink 和 padding 等边界也应纳入验证。
以上是源码检查和公式推导,没有 GPU 精度或性能实测结论。