☰
MXFP4量化KV Cache:大模型长上下文推理优化实践
2026/9/29 7:23:27 网站建设 项目流程

做长上下文推理优化的时候,估计不少人跟我一样被KV Cache这块硬骨头卡过:上下文一拉长,显存先爆,接着访存带宽拖后腿,decode阶段慢得让人抓狂。我最近在一个内部代号为“202606”的项目里试了试MXFP4方案,把注意力计算里的KV缓存改成了MX格式的4-bit浮点,再配合自定义的融合算子,效果比预想的要稳。这里把我踩过的坑、调参的心得和实测数据整理出来,给同样在搞大模型推理、量化部署或者算子开发的朋友一个参考。

先说清楚这文章聊的是什么:MXAttention是我在项目里做的一个面向Transformer解码器的融合注意力实现,核心低精度格式是MXFP4,也就是OCP MX规范里的4-bit浮点格式。它解决的是三个问题:KVCache的显存占用、解码阶段的内存带宽瓶颈、以及低精度算子在注意力这个场景下的精度退化。适合规模部署大模型、在长上下文下做推理优化、或者研究低精度训练/推理的工程师读。

1. 先理解MXFP4到底是什么,它跟INT4有什么区别

1.1 FP4不是INT4的替代品,它走的是另一条路

很多人一看到4-bit精度,第一反应是“这不就是INT4吗”。我一开始也是这么想的,真正用下去才发现两者思路完全不同。INT4是固定小数点格式,把数值范围均分成16个阶梯,均匀量化。FP4是浮点格式,走的是科学计数法的路子:1位符号、若干位指数、若干位尾数。典型配置有两种,e2m1和e3m0。e2m1意味着2位指数加1位尾数,能表达的最大值到+6,最小的正常浮点数在0.5附近,非正常值可以下探到0.125左右,动态范围跨度远大于INT4。e3m0则是3位指数加0位尾数,相当于只有8个不同的量级阶,没有尾数精度,只能表示1、2、4、0.5、0.25这些二的幂次。

这个差异看起来只是格式问题,放在量化场景里就是本质区别。注意力里Q、K、V向量经过LayerNorm之后,分布虽然相对稳定,但随着位置编码叠加和跨层传播,数值范围还是会波动的,尤其在长上下文下,某些token位置的K/V范数可能比平均值大好几倍。INT4用均匀阶梯去切这个分布,遇到极值就容易产生严重的相对误差。FP4的指数编码天然提供了更大的动态范围,即使尾数只有1位,也能通过指数偏移把量程拉开。

我实际测过一组数据:在一个7B模型上,同一批KV向量,INT4量化带来的attention score误差在部分head上能到4%以上,MXFP4的e2m1配置普遍压在1.5%以内。差异主要来自对异常值的刻画能力。如果向量内部元素都比较接近,INT4精度更好,因为尾数位宽更足;但只要分布出现拖尾,FP4的优势就出来了。注意力里的KV向量恰恰经常有拖尾,这就注定了MXFP4在这个场景比INT4更适用。

1.2 MX共享Scale的设计逻辑

MXFP4通常不会单独使用,而是配合MX规范里的block scale机制一起出现。所谓block scale,就是每固定数量的元素共享一个缩放因子,论文和硬件实现里常见的是32个元素作为一个block,每个block配一个FP32的scale。实际数值等于FP4尾数乘以对应的scale。这个设计的本质是为了解决浮点格式在数值太大或太小时的表达空白。

如果你不用block scale,直接拿e2m1去量化任意向量,假设某个元素是300,而FP4最大只能表达6,那就只能截断成6,信息直接丢失。但如果每32个元素配一个scale,比如把scale设成50,那300除以50等于6,正好落在FP4的表达范围内,反量化时乘回50,300就保住了。FP4负责表达相对大小和精细差值,scale负责适配量级,两者结合就让4-bit格式拥有了接近动态范围无限的表达能力。

MX规范里对scale的选取还有讲究,通常用的是矩阵或向量的绝对最大值或者二阶范数。我在实现里用的是amax机制:先取block内32个元素的绝对值最大值,再按FP4可表示的最大浮点值去反推scale。有一点容易漏:因为e2m1的最大值是6,所以实际scale要取max_value/6,而不是直接取max_value。如果直接用max_value当scale,编码时所有值除以scale之后落在0到1之间,白白浪费了FP4在1到6之间的全部表达空间,精度亏得很明显。我最初就是在这里没注意,导致Kernel跑起来困惑度一下掉了0.3,查了半天才发现是scale计算公式少除了一个6。

1.3 什么时候选MXFP4而不是FP8或者INT8

用MXFP4之前,我其实先试过标准的FP8 KV Cache。FP8的精度确实不错,训练和推理都比较友好,但KV Cache的显存开销只减了一半,与4-bit方案相比没有质的飞跃。在GQA(Grouped Query Attention)结构下,KV Cache的体量本来就比MHA小一些,但如果上下文长度上到32K、64K,FP8依然会吃掉大量HBM,decode带宽问题也依然存在。

另一种思路是INT8加per-channel scale,或者做比较激进的INT4加per-group scale。INT8方案误差控制好,但显存压缩率不够;INT4加group scale的方案能压下去,可对一个group内的异常值鲁棒性差。我的经验是:如果模型本身的激活分布方差比较大,或者目标上下文超过16K,MXFP4在精度和压缩率的平衡上是更优解;如果模型输出非常稳定,或者主要跑短上下文,FP8或INT8反而实现成本更低,生态工具也更成熟。

2. 为什么注意力计算盯上了KV Cache

2.1 真正贵的是访存,不是计算

看Attention的FLOPs,很多人第一反应是“这玩意儿算力需求很大”,但从实际Kernel profiling来看,decode阶段的瓶颈几乎都在内存带宽上。单token生成时,Q只有一个向量,K和V是整个序列的缓存,规模是序列长度乘以隐藏维度再乘以层数。计算量很小,但K/V的读取量是逐token线性增长的。假设序列长度是4096,隐藏维度是4096,单层KV缓存就是4096乘以4096,FP16下要32MB,模型有32层就是1GB。这个数据要从HBM读进SM,比例远超计算量本身。

把KV压成MXFP4之后,同样的数据量变成原来的四分之一,HBM读压力直接下降。GPU上很多算子的效率卡在DRAM带宽上,只要压了内存搬运量,带来的加速比是实打实的。我在A100上实测,decode阶段把KV从FP16换成MXFP4,在其他条件不变的情况下,单token延迟下降了大约35%到40%,换算成端到端吞吐提升非常可观。

还有一个容易被忽略的点:L2 Cache。FP16的KV数据占L2空间太大,导致下一层要读的数据经常被挤出L2,反复miss回HBM。压成4-bit后,同样大小的L2能装下四倍序列长度的KV,cache命中率明显改善。尤其在中短序列下,这种L2层面的改善比HBM带宽的改善更明显。

2.2 融合算子要从数据流角度重新拆解

MXAttention不是我简单地把KV cache换了个存储格式,然后把反量化逻辑塞进现有的FlashAttention里,而是从数据流层面重新拆了一遍Attention实现。标准Attention流程是:读Q和K,算QK^T,Softmax缩放,再乘V,得到输出。FlashAttention把这个过程做了分块和在线Softmax优化,减少了中间矩阵的落盘。MXAttention在这个基础上又加了一层:KV的读取和反量化与后续矩阵乘融合,让FP4数据直接进寄存器或共享内存。

具体来说,我按block加载FP4格式的K,先把scale和尾数都读进寄存器,在寄存器里做反量化,还原出FP16或者FP32的值,立刻参与QK^T计算。这样FP4数据只在HBM里占4-bit空间,一旦进入芯片内部,就在寄存器层面恢复成高精度浮点,避免了把整个KV Cache先反量化到一个临时FP16缓冲区再计算的额外访存。这个流程看着没什么,但在长序列下差别很大:如果不做融合,反量化后写回临时缓冲区就得多一次HBM写入和一次HBM读取,一来一回等于多搬了8倍数据。

2.3 从Prefill到Decode的不同策略

Prefill阶段和Decode阶段对MXFP4的收益差异也很大。Prefill阶段并行处理一整段prompt,Q矩阵很大,计算密度高,此时KV Cache访存占比相对低,单纯的访存优势不明显,反而是精度损失带来的质量影响更容易暴露。因此我在实际部署时,Prefill阶段默认走FP16 KV,只有Decoder阶段才启用MXFP4。两个阶段共用同一份KV存储会有点麻烦,我的做法是写了个kernel入口,根据phase参数决定是否启用MXFP4 load,Prefill写FP16,Decode读的时候如果发现存储格式是MXFP4就走反量化路径,如果有必要,甚至可以在Prefill完成之后把FP16的KV原地转成MXFP4。

Decode阶段为什么值得用MXFP4?因为此时每一步只有一个token的Q,K/V的读取成为绝对主导,访存压缩的收益最大。实际跑起来,同一个模型同一个prompt,混合策略的pipeline在吞吐上比全程FP16高出约1.8倍,而困惑度几乎持平,相差不到0.05。

3. 核心实现细节与量化方案

3.1 按Block量化KV的在线Scale计算

MXFP4的量化并不是提前离线算好的,而是在推理过程中实时计算。每次KV Cache需要写入时,计算它的block scale,然后把尾数写进缓存。这个流程适合叠加到原始的KV Cache写入kernel里,不额外增加一次全量扫描。

具体步骤我整理成了这个流程:

  1. 取当前KV向量的一段,长度32,按block切分。
  2. 计算该block的amax,即绝对值最大元素。
  3. 根据amax和FP4最大可表示值6,算出scale = amax / 6。
  4. 向量里的每个元素除以scale,就近取整到FP4可表示的浮点集合。
  5. 把scale用FP32保存下来,与4-bit尾数一起存储。

这里有几个容易踩的细节。一是只要遇到amax为0的block,scale直接置为0,反量化时把整个block视为0,不用走除法,性能和安全都兼顾。二是在计算scale时如果直接用amax而不是amax除以6,会损失尾数精度,前面说过了。三是block切分最好按最后一个维度连续切,这样才能在硬件上连续读取,避免cross-stride的scatter操作。

从误差角度看,per-block scale比per-tensor scale强得多。per-tensor scale相当于让整条KV共享一个缩放因子,为了让极端值不溢出,scale会偏大,大部分普通元素编码后都挤在很小的量级上,精度损失明显。per-block scale只有32个元素共享一个scale,局部动态范围适配得好,这也是MX规范推荐这种粒度的原因。

3.2 QK^T、Softmax和Aggregation的全流程

量化不是只换个存储格式就完了,关键是计算过程怎么组织。我实现的MXAttention流程是这样的:

  • Q保持FP16或者FP32,不量化。
  • K和V在HBM里以MXFP4格式存储。
  • 计算QK^T时,按block读取K,反量化到FP16或FP32,然后和Q做点积。
  • Softmax的结果在线计算,我用的是FlashAttention风格的online softmax,维护running max和running sum。
  • 之后O = softmax(QK^T)V,V同样从MXFP4格式按block读取并反量化。
  • 累加过程全程用FP32。

这个流程里最需要注意的点是:Q不要压成FP4。我试过把Q也压到MXFP4,推理质量下滑得厉害,尤其是一些对数值敏感的head,几乎直接失效。原因是Q和K做点积时,Q的精度直接影响每个注意力权重的精度,一旦Q量化出现相对误差,相当于给score加了个噪声。把Q保持在FP16,K用FP4,相当于只对内存体积大的那一侧做压缩,而计算精度主要依赖Q。这是性价比非常高的折中方案。

3.3 关键Kernel伪代码示例

我贴一个简化版的kernel逻辑,便于理解整体结构。这个版本忽略了很多边界条件和硬件细节,但数据流是对的。实际开发时,我是先用这个简单版本跑通正确性,再去写TMA和双缓冲优化版本。

// 简化版:MXFP4 KV Cache + Attention // Q: [num_heads, head_dim] FP16 // K_mx: [seq_len, head_dim/2] MXFP4 (因为head_dim一般128, block_size=32) // K_scale: [seq_len, head_dim/32] FP32 __global__ void mx_attention_decode_kernel( const half2* Q, const uint8_t* K_mx, // 4bit packed const float* K_scale, const uint8_t* V_mx, const float* V_scale, float* O, int seq_len, int head_dim) { int tid = threadIdx.x; int block_seq = blockIdx.x; // 按序列分块 float q_local[HEAD_DIM]; // load full Q to registers (FP16 -> FP32) load_half2_to_float(Q, q_local, HEAD_DIM); float acc[HEAD_DIM] = {0.0f}; float m_i = -1e30f; float l_i = 0.0f; for (int kb = 0; kb < seq_len; kb += BLOCK_SIZE) { // 读取K的MXFP4 block uint8_t k_packed[32 / 2]; float k_scale = K_scale[kb / 32]; load_mxfp4_block(K_mx + kb * head_dim / 2, k_packed, BLOCK_SIZE); // 反量化K float k_vec[HEAD_DIM]; dequant_mxfp4_block(k_packed, k_scale, k_vec, HEAD_DIM); // 计算QK^T 点积 float score = dot_product(q_local, k_vec, HEAD_DIM); // online softmax float m_new = fmaxf(m_i, score); float alpha = __expf(m_i - m_new); float p = __expf(score - m_new); l_i = l_i * alpha + p; // 读取V block 并累积 uint8_t v_packed[32 / 2]; float v_scale = V_scale[kb / 32]; load_mxfp4_block(V_mx + kb * head_dim / 2, v_packed, BLOCK_SIZE); float v_vec[HEAD_DIM]; dequant_mxfp4_block(v_packed, v_scale, v_vec, HEAD_DIM); for (int d = 0; d < HEAD_DIM; d++) { acc[d] = acc[d] * alpha + p * v_vec[d]; } m_i = m_new; } // 归一化 for (int d = 0; d < HEAD_DIM; d++) { O[d] = acc[d] / l_i; } }

实际工程版本里有两处变化比较大:一是head_dim通常是128,一个矩阵块同时覆盖多个block,scale就不是单个值了,而是一个小数组;二是TMA可以一次性把FP4数据按block拉进SMEM,避免逐元素读取。代码里的dequant函数就是按block scale对32个元素逐一乘回去。反量化本身开销很低,就是一次乘法加一次格式转换,关键是别把这步放到HBM里做。

3.4 显存布局和Packing细节

MXFP4有一个工程上麻烦的地方:它不是字节对齐的格式。一个元素4-bit,一个字节能塞两个元素。如果直接按字节存取,必须考虑packing顺序。我用的是低位在前,即第一个元素存在字节的低4位,第二个元素存在高4位。这个约定需要在写KV和读KV两边保持一致,否则一路错到底。

另外,block size选32不只是为了精度,也是为了对齐。32个FP4元素正好占16个字节,等于一个float4的尺寸,读取时可以做128-bit memory transaction,带宽利用最充分。如果block size改成16,字节数变成8,读取效率差一些;如果选64,scale粒度变粗,精度下降。我对比过32和64的偏差,64的困惑度比32平均差0.08左右,但访存效率反而没明显提升,所以最终固定为32。

4. 实测数据与参数选型

4.1 与不同KV缓存方案的精度对比

我在一个内部微调过的7B模型上做了对比实验,评测时用了困惑度指标和一组人工设计的极端prompt集。同一模型、同一输入,KV Cache分别用FP16、FP8-E4M3、INT4-Group和MXFP4-e2m1存储。

KV Cache格式每元素比特数相对FP16显存困惑度(越低越好)Attention Score最大偏差
FP16161x8.720%
FP8 E4M380.5x8.760.3%
INT4 Group=3240.25x9.052.8%
MXFP4 e2m140.25x8.811.4%

FP16的困惑度最低,这是预期中的。FP8的精度损失极小,但如果只看显存收益,它没有4-bit方案的吸引力。INT4的困惑度恶化最明显,而且位置靠后的层误差有累积效应。MXFP4的困惑度比FP8高了0.05左右,比INT4低了0.24,同时拿到4倍压缩,综合看是最划算的。我后来又测了e3m0配置,困惑度掉到了9.3,直接放弃,e3m0虽然表示范围更大,但没有尾数位,对普通数值的刻画太粗糙了。

4.2 吞吐和显存收益

集成MXAttention后,我在同样的推理框架下对比了端到端的效果。模型配置是7B、GQA、32层,输入prompt 4096 token,输出512 token,测了FP16 baseline和MXFP4两种情况。

场景KV显存占用Decode吞吐(token/s)单层Kernel耗时占比
FP16 baseline12.6GB64.5100%
MXFP4 KV3.2GB89.272%
MXFP4 + 融合3.2GB104.858%

这里有三个结论值得细看。第一,显存从12.6GB降到3.2GB,意味着同样的显存预算下,上下文长度可以拉长到原来的约4倍。第二,光换格式不融合,decode吞吐提升大约38%,这部分提升主要来自HBM带宽压力的下降。第三,加上融合kernel后,在相同存储格式下又涨了将近18%,这部分收益来自避免反量化中间缓冲区的额外写入和读取。第三个结论可能出乎不少人的预期,但实际上长序列下中间缓冲区的读写量非常大,省掉这一步带来的加速比一点不比压缩带宽少。

4.3 参数选型建议

根据我实验的积累,MXFP4用于注意力的参数选型顺序应该是:block_size优先设32,格式优先选e2m1,Q不量化,scale用FP32保存。这四个参数基本不用再纠结。

block_size如果设成16,scale多了,精度微升,但读取效率下降,实测吞吐损失约7%。设成64,精度下降明显,但吞吐提升只有2%左右,不值。e2m1和e3m0之间不用犹豫,除非你对显存有极端要求且能接受较大精度损失,否则e2m1是唯一合理的选择。scale的存储精度我试过FP16和FP32,FP16 scale会让某些大头位置出现0.5%左右的额外误差,而FP32 scale占用只多了一点点,因为scale数量是元素数的三十二分之一,成本可忽略,所以无脑选FP32。

另外还有一个容易被忽略的点:某些模型做了Sliced Attention或者GQA分组之后,KV在不同group之间重复存储,让量化收益打了折扣。如果遇到这种结构,建议在KV写入层就把分组的共享关系利用起来,相同数据只存一份MXFP4。

5. 常见问题与排查技巧实录

5.1 数据流中莫名其妙的NaN和Inf

我第一次把MXAttention跑起来的时候,前向直接出现NaN,排查了整整一个下午。后来发现根因不在kernel本身,而是scale计算的一个边界条件:某些padding位置的向量全为0,amax是0,0除以6等于0,scale是0,反量化时x乘以0等于0,这本身没问题。但问题是有些硬件路径在scale为0的block里如果恰好有非零尾数,做除法时会触发invalid operation,直接产生NaN。

解决方案有两层。第一层是在量化阶段,遇到amax是0的block直接写一个特殊的scale值(我用的-1.0做标记),反量化时看到-1.0就直接返回0,不走乘法路径。第二层是在kernel里加clamp,所有反量化后的值都限制在一个合理范围内。这两层让我在后面几个模型上再也没见过NaN。

5.2 收益没有想象中大:反量化开销补偿问题

有朋友抄了我的方案之后反馈说速度没提升,甚至变慢了。他用的实现是先开一个FP16缓冲区,把MXFP4的KV全部反量化到缓冲区,然后再跑标准Attention。这个做法等于你省了HBM读4-bit,但又把反量化后的FP16写回HBM,再读一遍。写入165MB、再读165MB,加在一起比直接读FP16还要多出一倍流量,自然慢。

正确的做法必须是反量化和计算融合在一个kernel里完成,FP4数据从HBM读出来后,直接留在寄存器或SMEM里,经过scale恢复成高精度值就立刻参与计算,绝无第二次落地的过程。这个道理说起来简单,但实际写TMA版本的时候很容易为了图省事先搞一个临时缓冲区。我自己的经验是,可以先用临时缓冲区版本验证正确性和精度,再花时间做真正融合的版本,因为后面这个版本的收益占到总收益的三四成。

5.3 长上下文尾部的精度退化

在长上下文场景下,我观察到decode到后半段时,模型输出质量比起前半段略有下降,困惑度升高了约0.15。排查后发现不是MXFP4的问题,而是attention得分在尾部极端情况下出现了较大的绝对值。当某个query和某个key的内积远大于正常分布时,softmax会把概率几乎全给它,此时这个score即使只有1%的误差,也会对最终输出产生明显影响。

解决方案是给MXFP4的K增加一个per-block的截断策略。当block的amax超过某个阈值时,依然用FP4存储,但在计算QK^T时对score做一个logit cap,限制极端score的继续增大。这个操作不影响绝大多数token的正常计算,只在极端情况下保持数值稳定。我在这里多花了大概200行代码,换来了长上下文下稳定不掉点,值得。

5.4 不同算子库之间的格式兼容

MXFP4目前还不是所有深度学习框架都能原生支持的类型。PyTorch里没有内置的MXFP4 dtype,TensorRT的某些版本对MX格式支持也仅限于部分算子。在工程落地时,我建议把MXFP4当作一种“私有的紧凑存储格式”来看待,而不是当作框架原生类型。KV Cache写入时由自定义kernel负责编码,读的时候由自定义kernel负责解码,框架本身只要保证能搬运uint8的tensor即可。

如果你要跟别的组件共享KVCache,比如投机解码或者多轮对话的缓存管理,需要特别注意格式标记。我在工程里用一个额外的int32 tensor记录了每个block的scale状态,并在KV缓存管理结构里增加一个format字段,值为0表示FP16、1表示MXFP4。在切换格式时,需要重新初始化KV Cache。这个字段虽然小,但避免了不同阶段误读数据导致的诡异错误。

6. 一些从实测打磨出来的工程经验

如果让我只挑几条最重要的经验说,我会挑这三条。第一,MXFP4方案不要试图在所有阶段统一启用,Prefill和Decode要区别对待,Prefill阶段老老实实用FP16,Decode阶段再用MXFP4,既保住质量又不牺牲速度。第二,block scale必须和反量化kernel严格绑定,中间任何一次格式转换都可能让scale语义失效,KV Cache的生命周期内不要动格式。第三,先做一个能跑通的低性能版本,把精度验证了,再去做TMA等高性能优化,不要一上来就追求极致效率,否则出问题时很难定位是格式问题还是kernel问题。

我在调MXAttention的过程中还有一个体会:注意力算子看似简单,实际上每个细节都在跟访存和数据布局较劲。MXFP4的4倍压缩是敲门砖,真正拉开差距的是把反量化、load、计算融合成一个完整的数据流水。这个思路不但适用于KV Cache,也适用于其他任何大模型推理中的内存密集算子。项目做到后面,我已经在考虑把同样的MXFP4策略用在FFN的中间激活上,不过那是另一个话题了。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询