SmoothQuant实战:W8A8量化让大模型推理速度翻倍不掉精度
前阵子帮朋友调一个70B模型的线上服务他第一反应还是把权重压到4bit来省显存结果decode速度没提多少精度还心疼。我跟他说你缺的不是压缩比是W8A8。很多人一听到“8bit量化”第一反应是能省一半显存第二反应是精度肯定掉但真正把SmoothQuant这套等效缩放迁移的W8A8方案吃透之后你会发现精度不掉、速度翻倍、单节点塞下超大模型这三件事是可以同时成立的。这篇文章就是把我从原理到落地、再到踩坑的完整过程捋一遍给想在大模型推理上做W8A8量化的朋友一个可以直接参考的路线图。1. W8A8的问题从来不在权重而在激活值里的一小撮“钉子户”1.1 权重可以慢慢磨激活只有一次机会熟悉量化的人都知道把权重从FP16压到INT8本质上是离线任务权重在部署前就已经固定你可以花任意长的时间去统计它的数值分布按per-channel或者per-group选择缩放因子甚至把量化误差反反复复迭代消掉。这种静态量化几乎没有工程压力。你今天量化出来的权重不好那就重新调一版再出一版反正模型不会突然变脸。但换到激活值情况完全不一样。激活是前向推理过程中实时算出来的每个token的激活向量都不同你不可能提前知道它的最大值、最小值。所以常见的做法是per-token量化对每个token对应的一行向量统计当前行的绝对值最大值用这个scale把这一行压到[-127,127]。听起来天经地义对不对问题在于LLM的激活分布并不是“整行统一”的。这里先给一个结论W8A8真正的难点从来不在“把权重存成8bit”而在“把激活也压成8bit时精度还能不能保住”。权重量化你随时可以回退、可以调参激活量化在线上只有那么一次机会scale稍微定错整个GEMM的输出都跟着歪。这也解释了为什么很多团队做了很久W8A16权重8bit、激活16bit后一上W8A8就被卡住。顺带说一句W8A16和W8A8的差别。W8A16在GEMM时权重已经是INT8但激活还是FP16算子层往往要额外做一次dequant或者走混合精度的路径计算效率上打折扣。W8A8就干净得多两个操作数都是INT8直接进TPU或GPU的INT8 TensorCore累加在INT32里完成最后再一次性还原。不仅省显存GEMM本身也快了。所以真正的性能红利在W8A8这里不是W8A16。1.2 激活异常值channel维度的“房间里的大象”我把Llama一类模型的激活分布打出来看过画面非常直观激活矩阵的每一行是一个token每一列是一个channel维度。绝大部分channel的数值都很温和可能就在-5到5之间但总有少数几个channel的数值动辄上百、甚至上千而且这些channel不是随机出现的——它们在几乎所有token上都保持大值位置相对固定。这一小撮channel在量化论文里叫activation outliers。它们的麻烦是灾难性的你用per-token量化一个token行里只要有一个channel数值是800那这一行的量化scale就按800来定。行内大多数正常值被均匀切到[-127,127]的离散格点上后大量低位信息直接变成0或1档之差。说白了一颗老鼠屎坏了一锅粥。你可能会想那把per-token改成per-channel每个channel单独一个scale不就行了学术上可以工程上很难受。激活是运行时才产生的如果每个channel一个scaleGEMM之前需要对输入做一次逐通道的缩放而且这个缩放和后续的INT8 GEMM很难融合内存布局和kernel都会被打乱。那个开销足以抵消量化带来的收益。为什么这些异常值会集中在固定channel而不是随机漂移我的理解是这跟attention里某些维度承担了特殊的语义聚合功能有关。模型训练完成后它形成的“计算偏好”是固定的个别维度专门用来放大或者记录强信号于是这些维度的激活值天然比其他维度高一个量级。这不是bug而是模型容量分配的结果。也正因为如此这个问题在大模型上几乎普遍存在谁都没法绕开。这就形成了一个僵局activation的outlier导致量化范围被拉垮per-token量化救不了per-channel量化又做不起。SmoothQuant的“等效缩放迁移”恰恰是从这个僵局里找出一条绕过去的路径。2. 等效缩放迁移一个恒等式把量化难度从激活挪到权重2.1 一个不起眼的矩阵恒等式SmoothQuant的核心出发点特别朴素既然出问题的是激活侧的channel维outlier那我能不能让激活“看起来”没有outlier直接的想法是截断、clip但那是有损的。SmoothQuant想到的方式是用一个数学上完全等价的变换。矩阵乘法 Y XW 有一个恒等性质对任意一个对角矩阵 s都有Y XW (X · s^{-1}) · (s · W)展开到channel维度看把每个channel j 上的激活列 X[:, j] 除以 s_j同时把权重行 W[j, :] 乘上 s_j。因为 s^{-1} · s 1两次操作在浮点世界里相互抵消最后的矩阵乘结果一模一样连一个bit都不差。这个变换在数学上没有任何“技巧含量”但在量化上价值极大你可以给每个channel挑选一个合适的 s_j把激活里那些大的outlier压下去同时把权重拉上来一点。激活原本难量化的“体质”被削弱了而权重侧即使被拉大由于权重是离线的、可以用per-channel scale来刻量这部分负担完全兜得住。打个我很喜欢的比方同样的货物从A地搬到B地货量没变但你给卡车换了一条绕开拥堵的大路车队能不能顺利通行就有了天壤之别。这里“货物”是矩阵乘法的结果“路况”是数值分布“换路”就是缩放迁移。2.2 scale怎么定平滑因子α是核心旋钮剩下的问题就一个s_j 到底取多少直觉上我们希望经过平滑后激活的每个channel最大幅度和权重的每个channel最大幅度都在同一个量级这样两边的8bit量化都不会因为某个channel太大而吃亏。SmoothQuant论文给了一个非常直接的公式s_j max(|X_j|)^α / max(|W_j|)^(1-α)其中 X_j 是第j个channel的激活列W_j 是第j列的权重行。α是一个0到1之间的平滑因子控制“迁移强度”α0s1/max(|W_j|)相当于不迁移激活还是原样权重被归一。此时激活的outlier还在。α1smax(|X_j|)把激活完全压平权重按激活的尺度放大理论上激活最好量化但权重可能被撑得很夸张。α0.5两边各退一步平滑后的激活max和权重max都在sqrt(maxX * maxW)量级这是论文默认的起点。实际操作中α怎么选不是纯理论问题。我试过7B、13B、70B不同规模的模型一个比较稳的经验是规模越大、outlier越明显的模型α0.5往往就够了而小模型3B以下或者某些训练不太充分的模型outlier分布不集中可能需要把α往大了调比如0.7到0.8让激活那边更“干净”代价是权重侧更吃力。调α的时候最好跑一组ppl曲线别凭感觉定。还有一个容易被忽略的点公式里的 max(|X_j|) 不是随便拿一条样本统计出来的它在离线校准阶段确定这个校准集的质量直接决定整个方案的好坏。后面落地章节我会展开讲。2.3 “迁移”为什么不是有损近似很多人第一次听“把激活的难度迁移到权重”会觉得不靠谱这不是拆东墙补西墙吗关键要意识到这个迁移本身发生在浮点运算层面是精确的重参数化跟最终量化误差不是一回事。在SmoothQuant的落地流程里平滑后的激活 X·s^{-1} 会被per-token量化成INT8平滑后的权重 s·W 会被per-channel量化成INT8后面GEMM完再把scale乘回来。等式左边是原始精度等式右边是量化后的结果误差只出现在“量化X和量化W”这个环节。而SmoothQuant做的事情恰恰是让X和W都变成“最容易被8bit表示”的形态激活侧的outlier被压下去了权重侧虽然range变大但per-channel权重量化天生就支持每个输出通道一个scale本来就设计来处理这种通道间差异大的情况。所以SmoothQuant不是靠有损近似去“赌”量化误差而是把量化误差按最小化目标分配到了两个量化器都能承受的位置。相比之下直接对激活做clip或者换一个很大的per-token scale属于让所有通道一起受罪SmoothQuant则是让“难搞的通道”单独消化掉自己的问题。这也是它能在175B甚至530B这个体量上依然保持精度的底层原因。3. 从论文公式到推理引擎SmoothQuant落地的完整链路3.1 离线校准统计激活分布、确定平滑系数要拿到公式里的 max(|X_j|)必须做一次离线校准。我自己的流程是这样的准备校准集。推荐用目标场景或者分布接近的文本几十到几百条就够序列长度取模型常用的比如512或1024。论文里一般用几百条文本片段我实测下来100条左右开始稳定。用FP16或者FP32加载原始模型在校准集上做前向把每个Linear层以及Attention里的QKV投影、O投影输入侧的激活按channel统计绝对值最大值。同时读取对应权重矩阵每个输出channel的绝对值最大值。按公式算出每个层的s向量α先取0.5保存成量化配置文件。这一步有两个细节必须注意。第一统计激活max的时候不要直接取全局max这个我后面会专门讲建议取99.99%左右的分位数给运行时偶发的极端值留一点余量。第二校准过程必须关掉随机性比如关掉dropout、固定seed否则两次统计出来的scale会不一致复现性很差。伪代码大概是这个风格# 伪代码校准前向中收集每层Linear的输入激活 for name, module in model.named_modules(): if isinstance(module, nn.Linear) or qkv_proj in name: x activation_inputs[name] # [batch*seq, hidden] w module.weight.detach() # [out_features, hidden] x_max x.abs().amax(dim0) # 按channel统计激活max w_max w.abs().amax(dim1) # 按输出channel统计权重max alpha 0.5 s x_max.pow(alpha) / w_max.pow(1 - alpha) smoothing_scales[name] s拿到s之后权重做离线平滑再量化W_int8 quantize(s[:, None] * W)。激活侧的平滑留在在线阶段处理。3.2 在线推理一个图层只有一个“除以s”的动作部署阶段每个Linear层在进入GEMM之前多了一个操作输入X先除以s即乘以 s^{-1}。之后做per-token的INT8量化得到X_int8权重直接用预先算好的W_int8参与INT8 GEMM。GEMM输出是INT32累加结果再乘上token_scale × channel_scale还原成浮点。这里有个我很想强调的工程细节s这个向量只需要在进入GEMM之前作用一次千万不要在GEMM之后再乘一次也不要拆成两半分别乘。有些同学复现的时候习惯性把“平滑缩放”和“每token量化scale”搞混结果误差翻倍。我的建议是把两张scale分开管理smooth_scale离线算好per-channel向量用来做平滑迁移token_scale在线根据当前输入行的max算出来的per-token标量用来做INT8量化。这两套scale在内部实现里必须拆开再在最后的反量化阶段合并成一次浮点乘法否则很容易出错。另外不是所有层都需要参与平滑。LayerNorm、GELU、SiLU、Softmax这些非线性层通常保留在FP16/BF16里跑因为它们的输入输出范围不稳定硬上8bit很容易掉点。你要平滑的是Linear层以及QKV/O投影的输入激活目的是让后面那个GEMM的两个操作数都可以是INT8。KV cache如果要压到8bit那是另一套per-token/per-channel scale别和SmoothQuant的平滑scale混用。3.3 和主流推理框架的兼容情况SmoothQuant提出的这套范式因为效果好后来成了不少推理框架里W8A8量化的默认底座。TensorRT-LLM、vLLM、DeepSpeed-Inference、LMDeploy这些项目里你搜weight-activation quantization或者w8a8选项背后多多少少都是“先对激活做等效缩放、再走INT8 GEMM”这个思路。有的直接叫smooth_quant有的做成了API参数但核心概念是一致的。所以我一直跟人讲你不需要从零去写kernel完全可以先用开源实现跑通一个小模型确认ppl和下游指标符合预期再切到生产框架的优化算子。理解原理的价值在于一旦框架里的量化参数不符合预期你知道该去查哪一环——是校准集不对、平滑α没调对还是scale重复乘了。4. 单节点承载530B这笔“显存账”究竟怎么算4.1 显存账本拆开看标题里“单节点承载530B”听起来很唬人但算一下账你就知道SmoothQuant为什么敢说这话。以标称530B参数的稠密模型为例只看权重FP16存储530B × 2字节 ≈ 1060GB也就是超过1TB。单节点哪怕全是H100 80GB也装不下必须跨节点。W8量化530B × 1字节 ≈ 530GB。8卡H100 NVL单卡94GB或者GH200单卡141GB的节点权重本身已经可以放下再搭配量化后的KV cache单节点可行。对比W8A16和W8A8W8A16虽然权重省了但激活和KV还在FP16显存大头并没有完全压下来真正把激活、KV、GEMM全链路都往8bit压内存余量才足够容下530B这个级别的“巨无霸”。我把典型占用整理成了表格方案权重占用激活/KV典型占用单节点可行性FP16约1060GB高不可能W8A16约530GB高视KV策略勉强或仍需多节点W8A8SmoothQuant约530GB比FP16低一半以上高显存单节点可承载这里的数字只算了权重侧加上KV cache和推理请求的动态buffer实际需求会更高。工业界说的“能跑”一般还包含beam search、连续批处理等场景的内存余量所以我会建议大家按“权重占节点显存50%到60%”来划线留足KV和调度缓冲。4.2 为什么“单节点”是工业落地的分水岭可能有人说能跑在多节点和单节点有什么区别区别太大了。跨节点意味着模型张量要被切到多台机器上每一层前向都依赖NVLink甚至IB网络来回传中间结果。网络带宽再高也远低于显存带宽延迟、故障面、运维成本全部被放大。单节点内走NVLink全互联调度简单、故障域缩小部署运维的成本断崖式下降。SmoothQuant最初的论文在OPT-175B上做过完整展示用INT8 GEMM把175B跑进单节点相比FP16实现吞吐提升接近1.56倍显存占用也大幅下降。530B是同一个逻辑的延伸——把权重减半、激活减半相当于给每个节点多腾出一整个模型的余量原来要4台甚至8台的活一台挤一挤就放下。对私有化部署、机房内网受限的场景这个差距直接决定项目能不能立项。还有一笔很多人忽略的账decode阶段LLM的瓶颈不是算力而是显存带宽。你每生成一个token都要把全部权重从头到尾读一遍。权重从FP16变成INT8读权重的时间直接减半在不改任何kernel的情况下Tokens/s就能接近翻倍。这才是W8A8比W4-only方案在服务端更香的根本原因——W4-only可以省显存但decode时还得dequant、还要把其他结构一并搬运带宽红利没有W8A8来得直接。SmoothQuant同时踩中了“省显存”和“降带宽”两个红利。4.3 工业落地还差哪些砖不过话说回来“单节点能承载”距离“工业环境好用”还隔着一层。真要在生产环境跑530B你需要考虑prefill阶段的吞吐输入序列很长时GEMM是计算密集的INT8 TensorCore能不能吃满取决于kernel质量连续批处理支持显存碎片、动态KV cache分配都得和量化模型配合好多机多卡的张量并行/流水线并行per-channel的s向量在切片后如何保持正确上线后的监控量化模型偶尔会因为极端prompt出现输出退化需要能快速回滚到FP16版本。SmoothQuant解决的是“数学和算子层”的问题剩下这些是工程系统问题。我的经验是先在小规模验证精度和速度都达标再逐步上量最后再考虑多节点扩展。5. 复现SmoothQuant时踩过的坑与排查思路5.1 校准集一换ppl就崩我第一次复现SmoothQuant用的是公开的wikitext校准模型是Llama-2-7B跑完ppl几乎不涨我当时非常开心。结果把同一个量化模型切到代码生成任务上一测ppl涨了快2个点代码生成的输出肉眼可见地变傻。排查过程是这样的先怀疑权重量化有问题把权重换回FP16跑发现没问题再怀疑per-token的token_scale算错单测验证数学也没问题最后回到校准集那一环才发现问题出在scale是按wikitext的激活统计的而代码数据会触发完全不同的attention模式激活outlier的位置和幅度都变了。这个坑的根因不难理解平滑scale是离线校准的产物它只在“校准分布约等于线上分布”时最优。解法也很直接校准集必须包含目标场景的数据或者干脆混合多个领域。我在实际项目中会准备一个“领域均衡”的校准集确保代码、对话、文档都有覆盖再在多个下游任务上验证。5.2 平滑后权重出现“超大通道”INT8直接失守另一个我印象很深的坑α调大了之后激活那边确实漂亮了但权重侧出问题了。回想公式 s_j max(|X_j|)^α / max(|W_j|)^(1-α)如果某个通道的 max(|W_j|) 特别小、接近0而 max(|X_j|) 正常偏大算出来的s_j会非常大导致平滑后的权重在这个通道整体被放大到几十甚至上百。权重per-channel量化虽然能按通道刻量但一个通道内的数值被放得太大8bit的离散化步长也跟着变大权重这一路的精度反而崩了。排查的时候我是在逐层cosine相似度里发现的某几个层cosine掉到0.98以下其余层都正常。定位到具体层之后打开该层平滑前后的权重分布发现正是那几个“分母极小”的通道在作怪。处理办法有三个按成本从低到高给s_j设一个合理上界比如不超过 max(|W_j|) 的100倍避免极端放大把α调回0.5或者更小让平滑不那么激进对这类“分母太小”的通道单独走混合精度不平滑、保持FP16只对正常通道做INT8。实践中我倾向先设上界再微调α混合精度作为兜底。5.3 校准统计max和百分位数差别比想象大有一阵子我图省事直接拿校准集里的激活max来算s_j结果效果时好时坏。后来换成分位数统计稳定性明显提升。原因也很好理解校准集是有限的一次极端样本比如一个超长token组合产生的最大值未必代表线上的稳态分布。取99.99%或99.999%分位数等于主动放弃最野的那一小撮噪声为正常的量化范围省下宝贵的档位。我常用的几个档位对比统计方式优点缺点max绝对覆盖极端值容易被个别离群样本绑架scale偏大99.99%分位保留整体分布、抗离群可能漏掉真正的尾部极端案例99.999%分位尾部保护更好需要的校准样本更多计算也略重我的默认选择是99.99%分位数配合混合校准集使用。如果上线后发现极端prompt引起量化退化再往99.999%加或者额外加一个clip层做补偿。5.4 上线前的验证方法最后说说我怎么确认一个SmoothQuant量化模型真的能上线。最基础的是ppl对比基线模型和量化模型在同一份校准集之外的验证集上算ppl掉点一般不应该超过0.1到0.3再大的话要回到前面的坑去查。其次是几个常见下游benchmark比如常识推理、问答跑一轮看综合掉点。更细的做法是逐层算激活输出的cosine相似度比对FP16版本和W8A8版本在某一层的输出有多接近。一般非量化层能接近1.0量化层维持在0.99以上就问题不大如果哪一层掉了就针对那层的scale、α、校准统计方式去调。这三个维度合起来基本能避免“看起来ppl正常、线上却变傻”的尴尬。批量上线之前我还会把量化模型挂在影子流量上跑几天和FP16主模型并线对比输出这是最稳的兜底。6. 一些实在话SmoothQuant给量化带来的思路转变我自己的感受是SmoothQuant最大的贡献不只是给出一个具体的W8A8方案而是把“量化难题”的视角从“如何更准地量化难搞的数”转向“如何通过数学变换让数变得不难搞”。这个思路后来在量化领域影响非常大AWQ做权重量化时用到了类似的“按通道重要性缩放”思想SmoothQuant、Outlier Suppression等后续工作也都是在这个框架上继续修修补补。如果你现在要上手大模型落地部署我的建议是先别急着追W4先把你手里的模型完整跑一遍W8A8。W8A8在decode场景的带宽收益和显存收益都足够实在SmoothQuant带来的精度损失在绝大多数模型上都可以忽略。真正把W8A8吃透了再去看更激进的压缩方案也不迟。最后再分享一个实用小技巧如果你的生产环境不方便改kernel可以先只做“权重INT8 激活平滑后FP16”也就是只利用SmoothQuant的平滑思想降低激活的量化难度再逐步过渡到全W8A8。这样风险最小也能马上拿到一部分收益。量化的路很长但把SmoothQuant这一个点啃透你后面看所有8bit方案都会觉得豁然开朗。