AI 中文总结
提出移位-累加注意力,用二次幂量化将Transformer解码中的QK^T乘法转为位移与累加,在RTX 4090上实现4.60倍加速和2.5倍KV缓存压缩,揭示PoT优势在于密度而非吞吐量。
AI 中文摘要
二次幂(PoT)量化将乘法转换为位移,但迄今为止仅适用于softmax后的注意力-值乘积。较早且规模更大的乘积S=QK^T尚未以相同方式重构。我们将键缓存量化为有符号二次幂定点码,使得QK^T中的每个标量乘法变为符号翻转、位移和整数累加。定点余量F>=e_max将每个位移索引r=F-e转换为非负左移,使累加在整数算术中精确。我们添加了次幂尾数扩展,通过一次额外移位-加法将分数误差减半;一种移位精确的在线softmax,其运行最大值位于整数log2域中,使得每次重新缩放本身都是精确移位;以及对数量化的AV乘积。在RTX 4090上的1.1B Llama解码器中融合CUDA内核:在B=8、T=32k时,4位内核比FP16缩放点积注意力快4.60倍,KV缓存小2.5倍,批量解码在B=64时交叉至1.22倍FP16吞吐量,峰值内存为11.02对13.77 GB。相同分块上的匹配INT8乘累加内核达到5.29倍:在具有硬件4路INT8点积的GPU上,移位-累加并不比其替代的MAC更快。等存储控制——每个键恰好一个字节的8位移位码,通过相同内核运行——将纯算术替换的代价定为2.02倍,精度相当。因此,PoT表示的优势在于其密度和去除乘法器,而非原始GPU吞吐量。对指数宽度、尾数级别、粒度和舍入的扫描给出了设计教训:最近PoT相对误差是尺度无关的,因此额外指数位无益(在3、4和5位时eps_S=0.1302),精度必须来自尾数项。
英文摘要
Power-of-two (PoT) quantisation turns a multiplication into a bit shift, so far only for the post-softmax attention--value product. The earlier and larger product, S=QK^T, has not been reformulated the same way. We quantise the key cache to a signed power-of-two fixed-point code, so that every scalar multiplication in QK^T becomes a sign flip, a bit shift and an integer accumulation. A fixed-point head-room F>= e_max turns every shift index r=F-e into a non-negative left shift, making the accumulation exact in integer arithmetic. We add a sub-power mantissa extension that halves the score error for one extra shift-add, a shift-exact online softmax whose running maximum lives in an integer log2 domain so that every rescaling is itself an exact shift, and a log-quantised AV product. Fused CUDA kernels in a 1.1B Llama decoder on an RTX 4090: at B=8, T=32k the 4-bit kernel is 4.60x faster than FP16 scaled dot-product attention with a 2.5x smaller KV cache, and batched decoding crosses over at B=64 to 1.22x FP16 throughput at 11.02 against 13.77 GB peak memory. A matched INT8 multiply--accumulate kernel on the same tiling reaches 5.29x: on a GPU with a hardware 4-way INT8 dot product, shift-accumulate is not faster than the MAC it replaces. An iso-storage control--an 8-bit shift code of exactly one byte per key, run through the identical kernel--prices the arithmetic substitution alone at 2.02x at comparable accuracy. The advantage of the PoT representation is therefore its density and the removal of the multiplier, not raw GPU throughput. A sweep over exponent width, mantissa levels, granularity and rounding gives the design lesson: nearest-PoT relative error is scale-free, so extra exponent bits buy nothing (eps_S=0.1302 at 3, 4 and 5 bits) and accuracy must come from mantissa terms.