ARTICLE DETAIL

资讯详情

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

模型推理优化实战:量化、剪枝、蒸馏与算子融合的落地指南

模型推理优化实战:量化、剪枝、蒸馏与算子融合的落地指南 1. 从模型能跑到模型跑得省Model-Optimizer 到底在解决什么问题做过模型部署的人大概都有过这种体验训练阶段一切顺利指标也好看可一旦要把模型塞进实际业务环境问题就全冒出来了。推理延迟高得离谱、显存占用把显卡撑爆、批量请求一上来服务直接雪崩。这时候你会发现训练出一个能跑的模型只是万里长征第一步真正决定它能不能落地的是推理阶段的效率。Model-Optimizer 这个方向本质上就是冲着这个痛点去的。它不是一个具体的库或者框架而是一类工具和方法的总称——通过量化、剪枝、蒸馏、算子融合、图优化等手段把训练好的模型瘦身和提速让它在保持精度基本不变的前提下跑得更快、占得更少、成本更低。我接触这个方向是从一次线上事故开始的。当时一个视觉模型在测试环境跑得好好的单张图片推理 80ms结果上线后 QPS 一上来P99 延迟直接飙到 2 秒以上。排查下来发现是显存不够导致频繁的显存交换而根因就是模型本身太大、算子太碎。后来做了一轮量化和算子融合模型体积缩了将近 4 倍延迟降到 30ms 以内精度只掉了 0.3 个百分点。从那以后我就意识到模型优化不是锦上添花而是部署环节的必修课。这篇文章适合几类人看一是刚把模型训出来、准备部署但不知道怎么优化的工程师二是被推理成本和延迟折磨、想系统性了解优化手段的技术负责人三是对量化、剪枝这些概念有耳闻但没实操过、想搞清楚底层逻辑的开发者。我会尽量把每个手段的为什么讲透而不是只丢一堆 API 让你照抄。需要先明确一个认知模型优化没有银弹。量化、剪枝、蒸馏各有各的适用场景和代价选错了不仅不省事还可能把精度搞崩。所以下面我会按先搞清楚瓶颈在哪、再选手段、最后验证效果这个顺序来展开这也是我踩过坑之后总结出的最靠谱的路径。2. 优化之前先别急着动手定位瓶颈比选工具重要十倍2.1 为什么上来就量化是最常见的错误我见过太多人一提到模型优化第一反应就是上量化。这个思路不能说错但顺序反了。量化的收益取决于你的瓶颈到底在哪——如果瓶颈是显存带宽而不是计算量那量化带来的收益可能远不如你预期如果瓶颈是算子调度开销那量化甚至可能因为引入了额外的反量化操作而变慢。正确的做法是先做 profiling搞清楚时间到底花在哪。常见的瓶颈分三类计算密集compute-bound、访存密集memory-bound、调度密集launch-bound。这三类的优化策略完全不同。计算密集GPU 算力被打满典型表现是 SM 利用率高、计算单元繁忙。这时候量化尤其是 INT8能直接提升吞吐因为整数运算单元吞吐更高。访存密集算力没打满但显存带宽吃紧典型表现是大量时间花在读写权重和激活值上。这时候减少数据位宽量化或者减少数据量剪枝都有效。调度密集小算子太多GPU 大部分时间在等 kernel launch。这时候算子融合fusion比量化更管用。我一般会用torch.profiler或者 Nsight Systems 跑一遍看 kernel 的时间分布。如果发现一堆耗时几十微秒的小 kernel 排着队那基本就是调度密集先做融合如果发现几个大 kernel 占了大头且 SM 利用率高那就是计算密集量化优先。2.2 一套可复用的 profiling 流程具体怎么操作我通常分三步走。第一步用torch.profiler抓一次前向的算子级耗时import torch from torch.profiler import profile, ProfilerActivity model model.eval().cuda() dummy torch.randn(1, 3, 224, 224).cuda() with profile(activities[ProfilerActivity.CUDA, ProfilerActivity.CPU]) as prof: with torch.no_grad(): for _ in range(10): model(dummy) print(prof.key_averages().table(sort_bycuda_time_total, row_limit20))这张表能直接告诉你哪些算子最耗时。如果前 20 个算子里有一大半是elementwise、add、mul这种小算子那融合的空间就很大。第二步看显存占用和带宽。用torch.cuda.memory_summary()看峰值显存用 Nsight 看 DRAM 吞吐。如果 DRAM 吞吐接近硬件上限那就是访存密集。第三步算一下理论下限。比如你的模型是 100M 参数、FP16 存储那光读一遍权重就要 200MB按 A100 约 2TB/s 的带宽算理论下限就是 0.1ms。如果你的实际延迟是 5ms那说明有大量时间没花在有效访存上优化空间很大。提示profiling 一定要在真实输入尺寸和 batch size 下做。我见过有人用 batch1 调优结果线上是 batch32优化策略完全失效。2.3 把瓶颈量化成可对比的指标定位完瓶颈要把它转化成可量化的目标。我习惯用三个指标延迟latency、吞吐throughput、显存峰值peak memory。优化前先记录基线优化后逐项对比。指标基线优化目标测量方式单次推理延迟80ms 30ms100 次取 P50吞吐batch32120 QPS 400 QPS持续压测 60s显存峰值6.2GB 2GBmemory_summary 峰值有了这张表后面每做一步优化都能清楚知道收益多少、代价多少而不是凭感觉说好像快了点。3. 量化收益最大但也最容易翻车的一环3.1 量化的本质是用精度换效率量化的核心思想很简单把 FP32 或 FP16 的权重和激活值用更低的位宽INT8、INT4 甚至更低来表示。位宽降下来显存占用和带宽需求成比例下降同时整数运算单元吞吐更高所以速度也上去了。但代价是精度损失。FP32 有 23 位尾数能表示非常精细的数值INT8 只有 256 个离散值必然有信息丢失。关键在于怎么把损失控制在可接受范围内。量化的数学形式是x_int round(x / scale) zero_point其中 scale 是缩放因子zero_point 是零点偏移。scale 的选择直接决定了量化误差。常见做法是用校准数据集统计激活值的动态范围取 min/max 或者用更精细的直方图方法如 KL 散度来确定 scale。3.2 PTQ 和 QAT两条路线的取舍量化分两大流派训练后量化PTQ和量化感知训练QAT。PTQ 是模型训练完之后直接量化不需要重新训练成本低、上手快。缺点是精度损失相对大尤其是对量化敏感的模型比如检测、分割这类对数值精度要求高的任务。PTQ 又分动态量化和静态量化动态量化在推理时实时计算 scale灵活但慢静态量化提前校准好 scale快但需要校准数据。QAT 是在训练过程中模拟量化误差让模型提前适应低精度。精度保持得更好但需要重新训练成本高。我一般的策略是先试 PTQ如果精度掉得在可接受范围内比如分类任务掉 1 个点以内就直接用如果掉太多再上 QAT。import torch.quantization as tq # 静态 PTQ 的典型流程 model.eval() model.qconfig tq.get_default_qconfig(fbgemm) model_prepared tq.prepare(model, inplaceFalse) # 用校准数据跑一遍统计激活值范围 with torch.no_grad(): for data in calib_loader: model_prepared(data) model_quantized tq.convert(model_prepared, inplaceFalse)这段代码看起来简单但坑很多。比如qconfig选fbgemm还是qnnpack取决于你的目标硬件校准数据的分布必须和真实输入接近否则 scale 会偏。3.3 量化翻车的三个典型场景我踩过的坑里有三个特别典型。第一个是激活值动态范围过大。某些层的激活值存在极端离群点outlier导致 scale 被拉得很大大部分正常值被压缩到很小的量化区间里精度崩掉。解决办法是用 per-channel 量化每个通道独立 scale或者对离群点做裁剪clipping。第二个是首尾层敏感。模型的第一个卷积层和最后的分类层对精度特别敏感量化后精度掉得厉害。常见做法是这两层保持 FP16 不量化只量化中间层。第三个是校准数据不匹配。我用一批干净图片做校准结果线上输入是带噪声的监控画面激活值分布完全不同量化后精度惨不忍睹。后来改用真实业务数据做校准问题才解决。注意量化后一定要在完整的验证集上跑一遍精度不能只看几个样本。我见过有人抽样测了 10 张图觉得没问题上线后整体精度掉了 5 个点。3.4 量化精度的快速验证方法验证量化效果我一般用逐层对比的方法把量化模型和原始模型在同一个 batch 上跑逐层对比输出的余弦相似度。如果某一层的相似度突然掉到 0.9 以下那这层就是量化敏感层需要特殊处理。def compare_layer_outputs(fp_model, int_model, x): fp_acts, int_acts {}, {} # 注册 hook 抓中间层输出 # ... 省略 hook 注册代码 with torch.no_grad(): fp_model(x) int_model(x) for name in fp_acts: cos torch.nn.functional.cosine_similarity( fp_acts[name].flatten(), int_acts[name].flatten(), dim0) if cos 0.99: print(f敏感层: {name}, 相似度: {cos:.4f})这个方法能快速定位问题层比盲目调参高效得多。4. 剪枝与蒸馏当量化不够用时的补充手段4.1 剪枝不是删了就行结构化与否差别巨大剪枝的思路是去掉模型里不重要的权重或结构。听起来简单但实操里最大的坑是非结构化剪枝把单个权重置零虽然能压缩存储但在通用 GPU 上几乎带不来加速因为 GPU 的并行计算模式对稀疏矩阵并不友好除非你用专门的稀疏计算库。真正能带来加速的是结构化剪枝——直接删掉整个通道、整个注意力头或者整个层。这样模型的稠密结构变小了GPU 能实打实地跑得更快。结构化剪枝的关键是怎么判断哪些结构不重要。常见的重要性度量有权重的 L1/L2 范数、BN 层的缩放因子BN scaling factor、梯度信息等。我比较常用的是基于 BN 缩放因子的方法因为 BN 的 gamma 参数本身就反映了该通道对最终输出的贡献训练时加个 L1 正则让不重要的通道 gamma 趋近于零然后剪掉这些通道即可。# 训练时给 BN 的 gamma 加 L1 正则 def bn_l1_loss(model): loss 0 for m in model.modules(): if isinstance(m, torch.nn.BatchNorm2d): loss m.weight.abs().sum() return loss # 训练循环里 total_loss task_loss 1e-4 * bn_l1_loss(model)剪枝率一般从 10% 到 30% 起步逐步增加每剪一次都要重新微调fine-tune恢复精度。一次性剪太多精度很难救回来。4.2 蒸馏让小模型继承大模型的能力蒸馏的思路是让一个小模型学生去模仿一个大模型教师的输出。学生模型不仅学真实标签还学教师的软标签soft label即 softmax 输出的概率分布。软标签包含了类别之间的相似性信息比硬标签信息量更大所以学生模型能学到更多。蒸馏在什么场景下最有用当你需要极致的小模型、且能接受重新训练成本时。比如要把一个 BERT-base 蒸馏成 6 层的小模型或者把一个大检测模型蒸馏成轻量 backbone。蒸馏的损失函数通常是两部分加权# 硬标签损失 软标签蒸馏损失 loss alpha * ce_loss(student_logits, labels) \ (1 - alpha) * T * T * kl_div( F.log_softmax(student_logits / T, dim1), F.softmax(teacher_logits / T, dim1))其中 T 是温度系数控制软标签的平滑程度。T 越大概率分布越平滑类别间的相对关系信息越丰富。一般 T 取 3 到 10 之间alpha 取 0.3 到 0.7 之间需要根据任务调。4.3 三种手段的组合策略实际项目里量化、剪枝、蒸馏往往不是单选而是组合使用。我常用的组合顺序是先蒸馏得到一个结构更小的模型再剪枝进一步压缩最后量化做部署级优化。这个顺序的逻辑是蒸馏改变的是模型结构剪枝改变的是模型宽度量化改变的是数值精度从粗到细逐层优化。但组合也有风险误差会累积。每做一步都要验证精度如果某一步掉得太多就回退或者调整参数。我一般会设一个精度红线比如相对原始模型掉不超过 2%超过就停。手段压缩效果精度损失是否需要重训适用瓶颈量化2-4x小到中PTQ 不需要计算/访存密集结构化剪枝1.5-3x中需要微调计算密集蒸馏2-10x中到大需要重训结构冗余5. 图优化与算子融合被低估的免费加速5.1 为什么小算子多是性能杀手前面提到调度密集的瓶颈根源就是小算子太多。每个算子都要单独启动一个 GPU kernel而 kernel launch 本身有开销几微秒到几十微秒。如果模型里有几百个小算子光启动开销就累积成毫秒级。更糟的是小算子往往访存效率低。比如一个add操作读两个张量、写一个张量计算量几乎为零但显存读写量是实打实的。GPU 的算力完全被浪费在等数据上了。算子融合就是把多个小算子合并成一个大算子一次读入、一次计算、一次写出。比如conv bn relu这三个操作可以融合成一个 kernel中间结果不落显存直接寄存器里传递。这样既减少了 kernel launch 次数又减少了显存读写。5.2 常见的融合模式和收益我整理了几种最常见的融合模式Conv BN ReLU最经典的融合几乎所有推理框架都会自动做。收益通常是 10%-20% 的延迟下降。Element-wise 链式融合比如add mul sigmoid这种连续的点操作融合成一个 kernel。收益取决于链的长度。Attention 融合Transformer 里的 QKV 计算、softmax、加权求和可以融合减少中间张量的显存占用。对大模型尤其重要。LayerNorm 融合把均值、方差、归一化、缩放合并减少多次遍历。这些融合大部分推理框架TensorRT、ONNX Runtime、TorchScript都能自动做但前提是你的模型图是干净的——没有奇怪的动态控制流、没有不支持的算子。我遇到过模型里有个自定义的if分支导致整个子图无法融合性能直接腰斩。后来把分支逻辑挪到图外融合才生效。5.3 手动融合的实操要点有些框架自动融合覆盖不到的地方需要手动改。比如把连续的view transpose contiguous合并或者把多个小矩阵乘合并成一个大矩阵乘。手动融合的核心原则是减少中间张量的物化。每产生一个中间张量就多一次显存读写。能在一个 kernel 里算完的绝不拆成两个。# 融合前三次显存读写 x self.linear1(x) x self.relu(x) x self.dropout(x) # 融合后如果框架支持合并成一个 fused kernel x self.fused_linear_relu_dropout(x)不过手动融合要小心数值一致性。融合后的计算顺序变了浮点误差可能累积导致输出和原来有细微差异。一般差异在 1e-5 量级可以接受但如果你的模型对数值极其敏感比如某些科学计算就要谨慎。提示融合后一定要做数值对比测试逐层对比融合前后的输出差异确保没有引入 bug。6. 优化效果的验证与回归别让优化变成劣化6.1 精度验证不能只看一个指标优化做完最怕的就是速度上去了、精度悄悄掉了。我见过太多案例优化后 top-1 精度只掉了 0.5 个点看起来没事但实际业务里某些关键类别的召回率掉了 10 个点直接导致线上事故。所以精度验证要分层做整体指标top-1、mAP、分类别指标、关键样本指标。尤其是业务里最重要的那些类别要单独看。我一般会维护一个关键样本集包含几百个业务上最关心的样本每次优化后都跑一遍确保这些样本的结果没有明显变化。6.2 性能验证要贴近真实场景性能验证的坑也不少。最常见的是实验室快、线上慢。原因可能是测试用的 batch size 和线上不一致、测试数据分布和线上不同、线上还有其他服务争抢资源。我的做法是性能测试一定要在接近线上的环境里做用真实的请求分布和 batch size并且要测 P50、P95、P99 多个分位数不能只看平均值。平均值好看但 P99 爆炸的情况太常见了。验证维度验证内容通过标准精度整体 分类别 关键样本相对基线掉幅 阈值延迟P50/P95/P99全部达标吞吐持续压测 QPS满足峰值需求显存峰值占用留 20% 余量稳定性长时间运行无内存泄漏、无精度漂移6.3 建立可回滚的优化流程最后一点经验优化一定要可回滚。每做一步优化都保存一个 checkpoint记录清楚做了什么、精度和性能变化多少。这样一旦发现问题能快速定位是哪一步引入的直接回退。我习惯用一个简单的表格记录每次优化的账本步骤操作精度变化延迟变化显存变化是否保留0基线-80ms6.2GB-1INT8 量化-0.3%35ms1.8GB是2剪枝 20%-0.8%28ms1.5GB是3激进量化 INT4-4.2%20ms0.9GB否回退这张表让整个优化过程透明可控也方便团队协作时交接。说到底Model-Optimizer 这类工作的核心不是会用某个工具而是知道什么时候该用什么、代价是什么、怎么验证。工具会更新换代但这套定位瓶颈、选择手段、验证效果的思路是通用的。我自己最大的体会是优化之前多花时间 profiling比优化过程中反复试错要省事得多。先把问题看清楚再动手往往事半功倍。
返回列表