ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

模型压缩实战:从ONNX量化到剪枝蒸馏的推理优化

模型压缩实战:从ONNX量化到剪枝蒸馏的推理优化 1. 一场由模型膨胀引发的线上事故Model-Optimizer的起点半年前我接手了一个推荐系统的在线推理服务模型文件从600MB一路涨到2.1GBGPU显存占用率逼近85%。业务方提了个要求月底前把单次推理成本压到原来的三分之一否则双十一流量洪峰来了就得加机器。加机器当然可以但一台A10的价格摆在那里预算根本不批。我不打算直接上多机部署因为问题根源很清楚模型本身太大了。当时的模型结构是DeepFM加上三套Transformer子网络embedding表占了75%的体积但线上流量分布极度稀疏大量embedding向量根本没有命中概率。这种情况继续堆参数边际收益几乎为零。于是我在内部立项了一个叫Model-Optimizer的小项目目标很简单——在不明显掉点的前提下把模型体积、推理延迟和显存占用同时压下去。项目最终跑通了模型从2.1GB压缩到286MBP99延迟从76ms降到31ms显存占用从7.8GB降到2.4GB准确率只掉了0.13个百分点。这个结果比预期好不少整条链路里的坑也值得单独拿出来复盘一遍。1.1 优化前必须回答的三个问题启动任何模型优化项目之前先别急着上工具花半天时间把下面三个问题想清楚否则后面全是白干。第一你的瓶颈到底是显存、延迟还是吞吐这三个指标的优化手段完全不同。显存瓶颈优先做量化顺便砍掉冗余的embedding维度延迟瓶颈优先看算子融合和IO开销因为大模型推理的很多时间其实浪费在数据搬运上吞吐瓶颈则要考虑动态batch和并发调度。我当时一上来就傻乎乎地先跑量化结果延迟没降多少反而把显存占用更低的问题给掩盖了后面重新做了性能剖析才纠正方向。第二你的评估指标能不能自动化优化过程中每轮迭代都要跑完整评测集人工看榜不现实。必须提前写好自动评测脚本把AUC、P99延迟、内存峰值这些指标固化下来每次优化完一键产出一张对比表。没有这个基线你根本分不清某个改动到底是变好了还是变坏了。第三上下游的兼容边界在哪里模型推理不是孤岛前面有特征工程后面有结果后处理。如果你把模型输出从float32改成int8后端的排序逻辑和打分逻辑能不能承接如果你剪掉了某些embedding桶特征线上对齐逻辑要不要改这些边界在优化前就要列清楚否则模型优化完了线上直接崩掉。1.2 我把Model-Optimizer拆成了三段式管线整个项目我没有做成一锅炖的工具而是按部署前的处理顺序拆成了三段第一步做模型分析搞清楚钱花在了哪里第二步做压缩和加速按量化、剪枝、蒸馏的优先级推进第三步做回归验证确保压缩后的模型在真实业务场景里依然可靠。管线设计成可插拔的每一段都接标准接口。输入模型统一转成ONNX格式后续所有处理都在ONNX图上做。这样做的好处是训练框架用什么写的已经不重要了PyTorch、TensorFlow还是Paddle只要导出成ONNX就能进入同一条优化流水线。2. 先别急着优化把模型结构里的“水分”挤干很多人拿到模型就急着上INT8量化这是典型的操作顺序错误。量化只是把数据的表示精度降下来如果模型结构本身存在冗余量化解决不了根本问题。我处理过的模型里至少有30%的体积和计算量是结构性的浪费。2.1 通过ONNX图分析找出真正的开销来源模型导出成ONNX后我习惯先用onnx.shape_inference跑一遍形状推断再结合onnxruntime的profiling工具看每个算子的耗时分布。这一步能直接暴露两件事一是哪些算子占了大部分时间二是哪些张量在反复进行无意义的复制和转换。实测下来我们的DeepFM部分问题最大。宽度方向堆了四十多个特征交叉模块每个模块都独立做了一次embedding lookup很多特征组的命中率不到0.5%。我把线上日志拉出来统计了一周的特征覆盖率把那些长期不命中的特征组标记为可裁剪。这一步操作纯靠数据分析不需要动模型训练逻辑把特征组从47个砍到18个后模型体积直接缩水了38%。2.2 Embedding表压缩用更小的维度装下同样的信息embedding表是宽表推荐模型的心头大患。我们的原始配置是每个特征桶维度64总表项接近两千万。砍完特征组之后我发现剩下的embedding矩阵依然存在大量稀疏维度。按奇异值分解的思路对每个特征桶的embedding矩阵做低秩近似把维度从64压到32再用两个小矩阵的乘积去近似原矩阵。这一招在代码上只需要给原embedding层加一个低秩适配器但收益极其明显。压缩后模型体积又降了24%线上AUC几乎没动因为原始高维空间里本来就有大量维度是噪声。这里有个前提条件必须确认特征值分布没有剧烈变化否则低秩近似的误差会放大。我在验证集上按天拆分做了三周的数据回放确认AUC波动不超过0.05%才敢上。注意低秩近似不是无脑压维度。如果某个embedding桶的取值基数本身就不高比如只有几百个值再砍维度会直接丢信息。我的做法是先统计每个桶的有效秩选那些有效秩远小于当前维度的桶下手。2.3 算子融合与冗余节点清理ONNX图里经常挂着推理阶段根本用不到的算子。我们在训练时为了数值稳定性加过一些统计量标准化节点部署到线上时这些节点每轮推理都在做无意义计算。用onnxsim做一次简化加上手写的图匹配规则把连续的两个Transpose合并、把Conv BatchNorm ReLU融合成单个算子图的节点数从2800多个直接降到900出头。这一步虽然没有改变模型的数学本质但推理延迟降了21%。原因是PyTorch的原生导出结果存在大量碎片化的kernel launch开销GPU对短任务的利用率极低算子融合之后每次kernel执行的有效计算占比显著提高。3. 量化落地从PTQ快速起步不够再上QAT模型结构清理完之后真正让体积和延迟同时大跳水的是量化。我用的是INT8量化把权重和激活值从FP32降到8位整数存储体积缩到四分之一整数算子还能在GPU上跑出比浮点更高的吞吐。但量化不是一行命令的事里面有几个决定成败的关键点。3.1 动态量化最容易出效果的入门方案如果你的模型里线性层和embedding层占大头动态量化Dynamic Quantization是性价比最高的起点。它的做法很简单权重提前转成INT8推理时再把激活值按需动态量化。因为不需要校准数据集也不用管激活分布的复杂性几乎是一键完成。from onnxruntime.quantization import quantize_dynamic, QuantType quantize_dynamic( model_inputmodel_clean.onnx, model_outputmodel_dynamic_int8.onnx, weight_typeQuantType.QInt8 )在CPU上测试动态量化把我们的纯DNN部分推理提速了2.3倍。但注意它对Transformer类模型的效果有限因为自注意力计算里的激活值量化没有解决瓶颈依然在浮点激活矩阵乘法上。3.2 静态量化必须做校准而且校准集要贴近线上分布想要把Transformer子网络也压到INT8就得用静态量化。它需要一个校准数据集统计每一层激活值的min/max范围提前算好缩放系数。我踩过一个大坑第一次做校准图省事直接用了训练集里的随机batch结果线上P99延迟确实降下去了AUC却掉了1.8个百分点。根本原因是训练集里的特征分布和线上实时特征分布差异很大校准出来的量化范围根本不匹配。后来我把线上过去七天的真实请求日志抽了5000条按特征来源分层采样做成校准集重新跑了量化掉点立刻收窄到0.2%以内。这个过程花了不少时间去清洗日志里的异常值和缺省特征但绝对值回票价。from onnxruntime.quantization import quantize_static, CalibrationDataReader, QuantFormat # 自定义CalibrationDataReader按batch从校准集加载数据 class CustomCalibrationDataReader(CalibrationDataReader): def get_next(self): # 逻辑略返回 {input_name: numpy_array} pass quantize_static( model_inputmodel_clean.onnx, model_outputmodel_static_int8.onnx, calibration_data_readerCustomCalibrationDataReader(), quant_formatQuantFormat.QDQ )静态量化之后整个模型跑在TensorRT上的INT8引擎里和FP16版本相比延迟又降了37%显存占用直接减半。3.3 混合精度不要让所有层都承受INT8的误差并不是所有层都对量化误差免疫。通过逐层敏感性分析我发现位置编码矩阵、FFN的第一层线性变换对量化特别敏感一旦压到INT8梯度式的误差会逐层放大。而embedding层、attention的投影层对量化容忍度很高。应对方案是混合精度把敏感层保持FP16其余层用INT8。Model-Optimizer里我实现了一个自动敏感度探测模块用一小批验证集逐层做“替换后对比输出差异”的测试按敏感度排序后自动决定哪些层保持高精度。最终只有17%的层保留了FP16但量化带来的掉点从0.7%收窄到0.13%。3.4 QAT当PTQ怎么调都调不动时的最后手段如果PTQ在你的模型上掉点始终超过0.5%就该考虑量化感知训练了。QAT的做法是在训练过程中模拟量化误差把量化的舍入操作当作一种噪声注入让模型在训练阶段就学会适应这种噪声。用PyTorch实现QAT时关键是在模型中插入torch.quantization.QuantStub和DeQuantStub并在fake_quantize模块里设置observer收集运行时的数值范围。训练时用较小的学习率通常是原学习率的十分之一做几轮微调。我们最终对attention里的敏感层单独做了QAT微调量化后精度反超PTQ版本0.1个百分点。重要QAT不是万能的。如果模型本身没有训练充分或者训练数据和线上数据分布差距过大QAT只会放大问题。先保证底模质量再考虑QAT。4. 剪枝与蒸馏在训练环节里省出来的计算能力量化解决的是算子层面的表示效率但如果网络本身的参数量就是冗余的量化之后依然有大量无效计算。剪枝和蒸馏就是在这个层面上做减法。4.1 结构化剪枝不要碰非结构化剪枝NVIDIA的TensorRT在处理稀疏权重时虽然能做2:4结构化稀疏加速但一般的推理框架对非结构化稀疏支持极差。我试过把模型里低于阈值的权重直接置零模型体积确实小了但推理延迟毫无变化因为稀疏矩阵在GPU上依然是密集排布计算的索引开销反而增加了。所以我的建议直接明确做剪枝就做结构化剪枝也就是按通道、按行、按注意力头为单位去整体删除。结构化剪枝能同时减少参数量、计算量和内存占用现有的深度学习框架都能完美承接。以Transformer为例我们可以对多头注意力机制做“头剪枝”。很多头在训练之后学到的高度相似尤其是深层网络冗余头比例可能达到30%。判断方法很简单在验证集上单独mask掉某个头如果输出变化极小说明这个头可以删。我们用了基于梯度的显著性分数对每个头计算它对loss的贡献量把排名靠后的头剪掉。4.2 剪枝后的微调周期不需要太长剪枝后必须微调否则模型精度会明显回落。但这里有个容易走极端的误区剪得越狠微调越久。实际上我观察到的规律是10%-20%的剪枝比例下微调一个epoch就能把指标拉回99%以上但剪枝超过50%微调训练时间指数级上升而且未必回得来。我们的实践是剪掉15%的Transformer通道和30%的冗余注意力头微调了2.5个epoch就恢复了。这里有一个技巧微调时不要冻结embedding表让它也跟着更新。embedding层的梯度本来就稀疏冻结它反而会让剪枝后的特征表示失配。4.3 知识蒸馏用小模型学大模型的“软知识”剪枝之后模型又从286MB涨了一点点我继续做了蒸馏。蒸馏的思路很简单让一个小模型学生去模仿大模型教师的输出分布而不仅仅是学习硬标签。关键参数是温度系数T。softmax的软化程度由T控制T越高分布越平滑学生能学到教师对相似样本之间的细微判别。我用T3.0把教师的logits和学生的logits做了KL散度损失再叠加一项硬标签的交叉熵损失权重分配是7:3。这个组合在精排模型上效果不错学生模型参数量不到教师的一半AUC只掉了0.06%。顺带说一句蒸馏和剪枝的顺序有讲究。我的顺序是先剪枝、再蒸馏。因为剪枝会引入一定的结构扰动蒸馏可以把教师模型的暗知识回流到学生模型里弥补剪枝损失的信息。反过来如果先蒸馏再剪枝等于学生模型跟着教师学完一套知识又被剪枝破坏了一遍效果会差很多。5. 验证与回归优化后不等于能用评测要过三关模型优化项目的最后一步才是真正见真章的地方。我见过太多项目优化完指标热闹一上线就出事。所以Model-Optimizer里我强制要求所有优化结果必须过三关验证缺一不可。5.1 第一关离线指标回归离线指标回归不只是看AUC。我习惯把业务核心指标拆成三个维度看排序质量、冷启动效果、尾部流量表现。因为模型压缩之后最容易先崩掉的往往是长尾部分。量化对低频特征的表示本来就有损再加上剪枝把一些低频通道删了长尾预测很容易失效。我跑完离线回归后发现整体的AUC只掉了0.13%但单独看低频用户群的AUC掉了0.62%。这个信号说明压缩模型对长尾信息的保留不足。于是我在量化校准集里增加了低频特征的采样权重让校准过程多关注那些分布稀疏但业务上重要的特征区间。5.2 第二关延迟与吞吐的压测离线指标再好看延迟超标就一切归零。压测时我用的是wrk配合推理服务自身的metrics接口分别测了P50、P95、P99三个分位数的延迟以及单卡并发吞吐。这里有一个特别容易忽略的点不能只测模型本身的算子里程碑要测完整推理链路的端到端延迟。特征加工、请求解析、结果排序在真实环境里占的耗时可能比模型本身还多。我们优化完模型后发现端到端延迟只降了18%但模型算子耗时降了37%说明IO和预处理成了新瓶颈。后来又额外优化了特征拼接的显存拷贝逻辑端到端延迟才真正达到目标。5.3 第三关灰度上线与稳定性观察灰度是最后一道保险。我的习惯是配置20%的流量先跑三天观察业务核心指标和模型预测分布的变化再逐步放量。发布期间还要监控显存碎片率、CPU/GPU利用率、推理服务的内存水位防止INT8量化后的特殊算子产生显存泄漏。我遇到过一次灰度期间的奇怪现象模型推理延迟整体稳定但每隔几分钟会出现一个300ms以上的尖峰。排查了半天发现是TensorRT的INT8引擎在遇到过长的输入序列时会自动触发fallback到FP32路径这个路径的编译和加载开销极大。解决方案是在服务层把超长序列截断并对截断策略做了AB测试确认业务损失可忽略后才放量。6. 踩坑记录四件我后悔没早做的事情Model-Optimizer整个项目做完最大的收获不是技术指标而是一堆用时间换来的教训。这里挑四个最有代表性的写出来希望能帮你少走弯路。6.1 没有在最开始就把线上日志接入校准集我前面提到校准集要用线上真实分布但这件事应该在项目第一天就做。因为训练集和线上分布之间的偏移不是短时间能清洗干净的数据筛选、采样、异常值过滤都有大量工程工作。如果从一开始就搭好校准集管道整个量化环节至少能省一周时间。6.2 过度依赖单个优化手段忽视组合收益单项技术的收益看着都不大算子融合降21%动态量化降15%静态量化再降37%剪枝降12%蒸馏又把体积降了20%。但如果我只做其中任何一项最终指标都到不了“三分之一成本”的目标。优化能力是指数级叠加的但前提是每一步都按正确顺序执行。先做结构瘦身再做算子融合接着量化最后蒸馏这个顺序不能乱。6.3 忽略了对embedding表的字节对齐优化embedding表在内存里是按行连续存储的但不同特征桶的embedding维度不一样导致访问时产生了大量非对齐读取。在GPU上非对齐访问会触发额外的内存事务拉高显存带宽的压力。我把所有embedding行都对齐到16字节边界并按照特征频率做了行重排把高频embedding放在连续内存段里。这个改动毫不起眼但端到端延迟又降了4%左右。6.4 团队对INT8结果缺乏信任导致上线决策拖延技术问题都好解决最难的是让业务方相信一个掉了0.13个百分点的INT8模型可以上线。后来我学到一个办法不只给指标对比而是把量化前后模型预测结果不一致的样本专门挑出来逐一分析业务合理性。大部分不一致样本都发生在特征缺失或极端长尾场景线上业务本来就很难对这些样本给出明确反馈。把分析报告给业务方看他们心里的疑虑自然就消了。Model-Optimizer这个项目的代码后来被我整理成了内部工具库团队里新同学也可以直接复用整条优化管线。每次碰到新模型只要把模型导出成ONNX丢进管线自动化报告就会输出体积压缩率、延迟变化、精度回退三张表省去了大量重复劳动。如果时间倒流一次我会在项目第一天就搭好这套自动化框架而不是等到项目中期才开始补课。
返回列表