差分注意力 本节摘要:Softmax 注意力给每个不匹配 token 都撒一点概率。到 10 万 token 这噪声累积,淹没信号。Differential Transformer(Ye 等,ICLR 2025)用两个 softmax 相减修这个——减掉共享噪声底,像降噪耳机之于注意力。DIFF V2(Microsoft,2026 年 1 月)是生产栈重写:decode 延迟匹配基线 Transformer、无需定制核、FlashAttention 兼容。本节带你走 V1 到 V2 端到端,用纯 Python 跑通差分操作的玩具实现,经验证噪声消除特性——你会理解为什么 softmax 永不产零、为什么相减抵消共享噪声却保留信号。
本节摘要:Softmax 注意力给每个不匹配 token 都撒一点概率。到 10 万 token 这噪声累积,淹没信号。Differential Transformer(Ye 等,ICLR 2025)用两个 softmax 相减修这个——减掉共享噪声底,像降噪耳机之于注意力。DIFF V2(Microsoft,2026 年 1 月)是生产栈重写:decode 延迟匹配基线 Transformer、无需定制核、FlashAttention 兼容。本节带你走 V1 到 V2 端到端,用纯 Python 跑通差分操作的玩具实现,经验证噪声消除特性——你会理解为什么 softmax 永不产零、为什么相减抵消共享噪声却保留信号。
阅读完本节,你应当能够:
标准 softmax 注意力有个在大规模变成操作头痛的数学性质。对查询 q,注意力权重是 softmax(qK^T/sqrt(d))。Softmax 永不产精确零——每个不匹配 token 都得一点正质量。那残留质量是噪声,且随上下文长缩放。到 128K token,即使每个不匹配 token 只得 0.001% 概率,127,999 个合起来约贡献总量的 12%。模型得学着绕开一个随上下文增长的噪声底。
经验上这表现为注意力头干扰:长上下文 RAG 的幻觉引用、10 万 token 检索任务的 lost-in-the-middle 失败、32K 以上针在草堆基准的微妙精度退化。Differential Transformer 论文(arXiv:2410.05258,ICLR 2025)度量了差距:同尺寸下 DIFF Transformer 困惑度更低、长上下文精度更高、幻觉更少。
DIFF V1 有三个问题把它挡在前沿预训练流水线外:decode 每步要加载两次值缓存、需定制 CUDA 核破坏 FlashAttention 兼容、每头 RMSNorm 在 70B+ 规模长训练不稳。DIFF V2(Microsoft unilm 博客,2026 年 1 月 20 日)三者都修了。本节走两版,搭差分算子,在玩具查询上基准噪声消除。
对查询 q 与键 K=[k₁,...,k_N],注意力权重 w_i = exp(q·k_i/sqrt(d)) / Σⱼ exp(q·k_j/sqrt(d))。没有 w_i 永远为零。若 k_i 与 q 完全无关,分数 q·k_i 不是 0——它围绕零波动,方差 ‖q‖²/d。Softmax 归一化后,每个无关 token 仍贡献 O(1/N) 给加权和。无关 token 的总贡献是 O((N-1)/N)=O(1)——不小的量。模型想要的是硬 top-k:匹配 token 高权,其余近零。Softmax 太平滑做不到。
把每头的 Q、K 投影各分两半:Q=(Q₁,Q₂)、K=(K₁,K₂),算两张注意力图:
A₁ = softmax(Q₁ K₁^T / sqrt(d)) A₂ = softmax(Q₂ K₂^T / sqrt(d)) 差分注意力 = A₁ - λ · A₂
思路:A₁ 和 A₂ 都含相同噪声底(无关 token 的 O(1/N) 贡献),但信号不同(它们学关注不同模式)。相减抵消共享噪声底,保留差异化信号。λ 是每头可学标量,可负,参数化为 exp(lq1·lk1) - exp(lq2·lk2) + λ_init。
DIFF V1(ICLR 2025):每头维减半保参数量(两个半头用一个头的参数),需定制核(decode 每步加载两次值缓存),每头 RMSNorm 作差分后的稳定器,但在 70B+ 长训练后期不稳。DIFF V2(2026 年 1 月):三改——(1) 翻倍 Q 头保持 KV 头(而非半头维),decode 算术强度升(每 KV 加载更多查询),匹配基线 decode 速度;(2) 无需定制核,FlashAttention 兼容(差分用两个标准 softmax 核算再相减);(3) 去掉每头 RMSNorm(改用更稳的初始化与 λ 参数化),长训练稳。结果:生产预训练可用,decode 延迟不输基线。
💡 关键概念:降噪耳机类比。降噪耳机采环境噪声、反相播放、与原噪声相消。差分注意力同理:A₂ 学到「噪声底」模式,从 A₁ 减掉,剩纯净信号。信噪比随上下文长提升更显著。
def softmax_noise_floor(seq_len): # 查询与一个 token 强匹配,其余 N-1 个无关 q = np.random.randn(d); keys = np.random.randn(seq_len, d) keys[0] = q * 5 # 位置 0 强匹配 scores = keys @ q / np.sqrt(d) weights = np.exp(scores - scores.max()); weights /= weights.sum() signal = weights[0]; noise = weights[1:].sum() # 信号 vs 噪声 return signal, noise
跑 seq_len=1k、10k、100k:噪声从约 8% 涨到 30%+,直观看到 softmax 噪声底随上下文长增长。
def differential_attention(Q1, K1, Q2, K2, V, lam): A1 = softmax(Q1 @ K1.T / np.sqrt(d)) # 第一张图(信号+噪声) A2 = softmax(Q2 @ K2.T / np.sqrt(d)) # 第二张图(信号'+噪声) return (A1 - lam * A2) @ V # 相减消共享噪声
构造合成查询:一个真信号位置加 999 个噪声位置。跑标准 softmax 与差分注意力各算信噪比(信号位权重 / 噪声位平均权重)。差分的信噪比应显著高——验证噪声底被减掉。
def compute_lambda(lq1, lk1, lq2, lk2, lam_init=0.5): return np.exp(lq1 @ lk1) - np.exp(lq2 @ lk2) + lam_init # 可学,可负
λ 控制减多少 A₂。初始化 λ_init≈0.5 让训练初期接近标准注意力,随训练学最优减幅。
DIFF V2 已进 Microsoft 的生产栈,与 FlashAttention 兼容意味着 vLLM、SGLang 可直接集成。对比标准注意力:V2 的 decode 延迟匹配基线(翻倍 Q 头升算术强度),训练时两次标准 softmax 核算(无定制核)。对比其他注意力变体:差分是唯一显式建模并消除噪声底的,滑动窗口/稀疏注意力是回避问题(减少注意力范围)而非消噪声。
本节产出 outputs/prompt-diff-attention-evaluator.md——一个提示,接收模型 config 与目标上下文长,评估是否该用差分注意力:算预期信噪比提升、长上下文检索精度增益、decode 延迟影响,给出 V2 集成建议。
(Easy) 扩展噪声底演示:绘信噪比 vs 上下文长(1k 到 256k)曲线,量化 softmax 噪声底如何淹没信号。
(Medium) 在「针在草堆」任务上对比标准与差分注意力:256 token 里藏一条事实,测两者找出针的准确率随草堆长的变化。
(Medium) 实现完整 V2 头:翻倍 Q 头保持 KV 头,算 decode 算术强度提升,验证匹配基线 decode 速度的数学。
(Hard) 训一个小差分 Transformer 与标准 Transformer(同等参数),在 WikiText 上比困惑度,验证论文的低困惑度、少幻觉声明。
(Hard) 实现 V1 的每头 RMSNorm 版本,在长训练(1000+ 步)上对比有无 RMSNorm 的稳定性,复现 V1 在 70B+ 的不稳,验证 V2 去掉它的合理性。
下一节,原生稀疏注意力(DeepSeek NSA):三分支并行,64K 解码比 FlashAttention 还快,质量持平或更优。