发表机构
Virginia Commonwealth University(弗吉尼亚联邦大学)
机构由 AI 辅助整理,请以论文原文为准。AI 中文总结
FoldAttention通过固定参考的加性softmax公式,在扫描KV缓存前确定权重,实现快速解码和确定性反向传播,在H100上解码加速最高3.09倍,反向传播加速最高1.84倍。
AI 中文摘要
自回归解码反复流式传输不断增长的KV缓存,使得注意力在长上下文场景下成为主要成本。现有的高性能内核使用在线softmax,在扫描键时发现行的归一化参考。因此,早期贡献仍然是临时的,可能需要重新缩放。我们认为参考无需发现:softmax对公共平移不变,因此参考只需保持权重在范围内。我们提出FoldAttention,一种加性softmax注意力公式,在扫描KV缓存之前固定有限参考$Z_i$。每个权重$2^{s_{ij}-Z_i}$在计算时即为最终值,因此贡献跨不相交键范围相加,其商在实数算术中等于softmax注意力。我们利用此属性为Hopper解码开发两种技术:(1) 最终权重在字节获取之前门控键和值读取,每次调用的深度$T$将低于$2^{-T}$的键截断,同时保留其质量;(2) 加性部分和组合分裂KV和共享前缀级联,无需重新缩放。在H100上,$T=16$时,FoldAttention解码七个真实模型生成比最快的BF16基线快1.36-2.30倍,在MHA和GQA形状上最高快3.09倍,七个中六个的错误率在最低BF16错误的1.5%以内;读取每个键时,在匹配错误率下快1.14-1.30倍。我们在Qwen3-8B上验证,整个解码步骤最高快1.46倍,同时似然和长上下文准确性与BF16内核匹配。相同原理使反向传播确定性:CTA将有界部分梯度舍入到归约前声明的整数网格上,并以任意顺序相加。FoldAttention因此消除了确定性税:其确定性反向传播比确定性FlashAttention-3/4快1.84倍,比最快的非确定性内核快1.05倍。
英文摘要
Autoregressive decode repeatedly streams a growing KV cache, making attention a major cost at long context. Existing high-performance kernels use online softmax, which discovers a row's normalization reference as it scans keys. Earlier contributions therefore remain provisional and may require rescaling. We argue that the reference need not be discovered: softmax is invariant to a common shift, so the reference only has to keep the weights in range. We present FoldAttention, an additive formulation of softmax attention that fixes a finite reference $Z_i$ before scanning the KV cache. Each weight $2^{s_{ij}-Z_i}$ is then final when computed, so contributions add across disjoint key ranges and their quotient equals softmax attention in real arithmetic. We use this property to develop two techniques for Hopper decode: (1) final weights gate key and value reads before the bytes are fetched, and a per-call depth $T$ cuts keys below $2^{-T}$ while keeping their mass, and (2) additive partials compose split KV and shared-prefix cascades without rescaling. On H100 at $T=16$, FoldAttention decodes seven real-model generations 1.36-2.30$\times$ faster than the fastest BF16 baseline, and up to 3.09$\times$ faster across MHA and GQA shapes, at an error within 1.5% of the lowest BF16 error on six of the seven; reading every key, it is 1.14-1.30$\times$ faster at matched error. We validate on Qwen3-8B that a whole decode step is up to 1.46$\times$ faster while likelihood and long-context accuracy match those under BF16 kernels. The same principle makes the backward deterministic: CTAs round bounded partial gradients onto an integer grid declared before the reduction and add them in any order. FoldAttention thereby removes the determinism tax: its deterministic backward is up to 1.84$\times$ faster than deterministic FlashAttention-3/4 and 1.05$\times$ faster than the fastest nondeterministic kernel.