先说说背景。我手上有一个叫 Model-Optimizer 的内部工程化项目,目标是解决模型训练完到上线之间那段“最后一公里”的问题。具体来说,就是训练好的 PyTorch 模型在 GPU 上跑得挺快,但一上生产环境、一放 CPU 推理、一塞进容器限了内存,各种问题就冒出来了:延迟超标、显存溢出、镜像太大拉取慢、QPS 上不去。Model-Optimizer 不是某个单一算法,而是一整套优化流水线,把量化、剪枝、蒸馏、算子融合、后端导出串成一条自动化的链路。
这篇文章把这套东西的完整思路、实操步骤、踩坑记录和最终效果整理出来。不管你是刚接触模型优化的新手,还是已经在做部署优化的老手,只要手上有需要上线的深度学习模型,都可以参考这里的方案。
1. 先想清楚:Model-Optimizer 到底在优化什么
1.1 优化不是单纯把模型变小
很多人一提模型优化,第一反应就是“把模型文件压缩一下”。实际做上线优化时,你会发现模型体积只是其中一个指标,甚至不是最重要的指标。线上服务卡不卡、用户等不等得起,才是核心。
一个典型的在线推理服务,关注的指标有三个:
- 延迟:单次请求从进来到返回结果的时间,要求通常 p99 要在几十到几百毫秒内。
- 吞吐:单位时间能处理的请求数,也就是 QPS。
- 资源占用:显存、内存、CPU 使用率,直接关系到机器成本和部署密度。
Model-Optimizer 要同时照顾这三者。举个最简单的例子:INT8 量化能把模型体积缩小到原来的四分之一,但这个收益在 GPU 上能不能兑现,取决于后端有没有对应的 INT8 kernel。如果只是把权重存成 INT8,推理时又转回 FP32 计算,那体积变小了,延迟一点没改善。所以优化必须从算法、计算图、运行时三个层面一起做,缺一环效果都会打折扣。
1.2 三个层面拆开看
我把 Model-Optimizer 的优化动作分成三层:
- 算法层:量化、剪枝、蒸馏,改变模型本身的表达方式。
- 计算图层:算子融合、常量折叠、图优化,减少计算图中的冗余节点。
- 运行时层:选择合适的推理后端、调整线程数、开启内存池复用、使用 TensorRT/ONNX Runtime 等加速引擎。
这三个层面的收益是叠加的。计算图优化打底,量化把精度降到 8bit,运行时再适配硬件指令集,最终才能同时拿到低延迟和高吞吐。如果只做其中一层,通常很难达到上线标准。
1.3 什么时候不适合优化
也不是所有模型都值得这么折腾。我一般先判断几个条件:模型是否要长期服役、请求量是否大到需要压资源、是否有硬实时要求。如果只是离线批量跑一次,CPU 时间多花几分钟也无所谓,那根本不需要做量化剪枝,直接拿原始模型预测就行。
反过来,如果模型要部署到手机端、嵌入式设备,或者线上 QPS 要求很高,那 Model-Optimizer 这套流程就是必需品。先确认边界,再投入精力,否则容易做了一堆优化,实际收益却不明显。
2. 核心技术栈与工具选型
2.1 量化:PTQ 起步,QAT 救场
量化是所有优化手段里投入产出比最高的一招。把 FP32 的权重从 32bit 压到 8bit,模型体积直接缩到四分之一,推理延迟在很多硬件上也能明显下降。Model-Optimizer 第一版只做了 Post-Training Quantization(PTQ),就是用一批校准数据跑一遍模型,统计激活值的分布,然后确定量化 scale 和 zero_point。
PTQ 遇到的最大问题是敏感层掉精度。某些层的激活值分布范围特别宽,用简单的 min/max 校准会损失大量信息。我的做法是先用每通道量化和百分位校准把这些层救回来,如果还不够,才考虑 Quantization-Aware Training(QAT)。QAT 需要在训练阶段就模拟量化的舍入误差,训练时间会变长,但精度恢复效果通常比 PTQ 好一个档次。
实操里我建议的推进顺序是:校准数据充足、模型对精度不敏感,直接用 PTQ。模型在 FP32 上本身就只差几个点就到目标精度,或者 PTQ 后掉点超过两个百分点,果断切 QAT。
2.2 剪枝:别只盯着非结构化稀疏
剪枝的目标是去掉不重要的权重或通道。非结构化剪枝把权重矩阵里接近零的元素置零,但得到的是一个稀疏矩阵,在通用硬件上很难直接加速,除非后端有专门的稀疏算子支持。Model-Optimizer 里我主要做结构化剪枝,也就是按通道或者按注意力头剪掉整个结构。
通道剪枝的流程是这样的:先用 BN 层的 gamma 系数作为重要性判断依据,gamma 接近零的通道对输出影响很小,可以剪掉。剪完之后必须做一次微调,否则精度会掉得很厉害。微调不需要太多 epoch,一般 10~20 个 epoch,学习率调小一点,让模型适应剪枝后的结构。
剪枝的收益在 CPU 和 GPU 上不太一样。CPU 上通道剪枝后的模型有比较明显的加速效果,因为计算量实打实减少了。GPU 上如果 kernel 没有针对新 shape 做优化,有时候反而变慢,因为 Tensor Core 对矩阵尺寸有对齐要求。所以剪枝之后一定要重新 benchmark,不要想当然。
2.3 蒸馏:让大模型教小模型
当量化救不回来、剪枝又剪不动的时候,蒸馏是最后一个大招。用一个精度高的教师模型去指导一个小学生模型训练,让小模型的输出尽量逼近教师模型。这里的输出不只是最后的 logits,还包括中间层的特征图,尤其是注意力矩阵,对小模型的帮助很大。
蒸馏的配置有几个关键点:教师模型直接用优化前的高精度模型,不需要重新训练。损失函数用 KD loss 加任务 loss 的组合,温度参数一般设 3~8。蒸馏之后的小模型如果继续做 PTQ 量化,精度能比直接量化大模型更好,因为小模型的决策边界更平滑,对量化误差更不敏感。
2.4 算子融合与后端适配
算法层的优化做完后,还要解决后端适配问题。PyTorch 的 eager mode 执行时每个算子单独调度,中间结果反复读写内存,浪费很大。Model-Optimizer 会把模型导出成 ONNX,然后用 ONNX Runtime 或者 TensorRT 做图优化。最典型的优化是 Conv+BN 融合、Conv+ReLU 融合、多头注意力里的 QKV 拼接合并。
ONNX 和图优化不是完全无痛的。有些 PyTorch 算子导到 ONNX 后可能不支持,或者导出的子图结构非常绕,后端优化不起来。这时就需要手工改模型结构,或者写自定义算子(custom op)把关键算子包起来。经验是:导出前在 PyTorch 里用 torch.jit 先 trace 一遍,把动态控制流和 Python 侧的逻辑尽量去掉,导出的图会更干净。
3. 实操记录:优化一个图像分类模型的完整流程
3.1 基线准备
我以 ResNet-50 在 ImageNet 子集上的分类任务为例。原始模型在单张 V100 上跑 FP32,输入分辨率 224x224,batch size 1,延迟大概 6.2ms,显存占用 246MB。这个数据就是基线,后面每一步优化都要跟它对比。
环境上推荐直接用英伟达官方镜像,我本地用的是 pytorch/pytorch:1.13.1-cuda11.7-cudnn8-runtime,补装 onnx、onnxruntime-gpu、torchmetrics 这几个包装好。模型代码单独封装成 module,保证训练和推理共用同一套前向逻辑,这样后续做 QAT 时不会因为代码分叉导致精度对不上。
3.2 第一步:PTQ 量化
我用了 PyTorch 官方的 torch.ao.quantization API。先把模型的每个模块按“可量化部分”和“不可量化部分”分开,比如卷积、全连接、ReLU 都可以量化,但最后的 softmax 不量化。校准数据从验证集里随机抽 200 张图,batch size 32,跑 5 轮,统计每层激活值分布。
核心代码逻辑大概是:
import torch from torch.ao.quantization import prepare, convert model_fp32 = load_resnet50().eval() model_fp32.qconfig = torch.ao.quantization.get_default_qconfig("x86") model_prepared = prepare(model_fp32, inplace=False) with torch.no_grad(): for images, _ in calibration_loader: model_prepared(images) model_int8 = convert(model_prepared)PTQ 后精度从 76.3% 掉到 74.1%,掉了一个多点。对于分类任务能接受,但离上线标准还差一点。我用 torch.ao.quantization 提供的 per_channel 配置重新量化,精度回到 75.4%,延迟降到 4.1ms,模型体积从 98MB 缩到 26MB。这是全流程里最直接的收益。
3.3 第二步:结构化剪枝加微调
量化之后精度没有恢复的空间,我开始做通道剪枝。使用 torch.nn.utils.prune 里基于 BN gamma 的全局剪枝策略,设置剪枝比例为 30%。
剪枝后模型结构从标准的 ResNet-50 block 变成了窄通道版本,计算量下降约 45%,但精度也掉到了 71.2%。所以紧接着做微调,用 ImageNet 子集训练 15 个 epoch,初始学习率 0.001,余弦退火调度。微调结束精度回到 74.8%。这里有个细节:剪枝后 BN 层的统计量需要重新估计,微调前先跑一遍 forward 收集 running_mean 和 running_var,否则微调前期 loss 会震荡。
剪枝后的模型继续量化,INT8 精度 73.6%,延迟进一步降到 3.2ms。体积因为稀疏化存储又小了一些,大约 19MB。
3.4 第三步:蒸馏补精度
剪枝加量化后距离原始精度还差将近 3 个点,没法接受。我引入蒸馏。教师模型就是最初那个 FP32 的 ResNet-50,学生模型是剪枝后的窄版。蒸馏训练 25 个 epoch,温度设 6,KD loss 权重 0.5。
这一轮下来,学生模型的精度追到了 75.9%。再走一遍 PTQ 量化,精度是 74.6%,比未蒸馏的剪枝模型高出 1 个点。到此,优化后的模型和原始 FP32 模型的精度差控制在 1.7 个点以内,延迟下降了近一半,体积缩小到原来的五分之一。
3.5 导出 ONNX 并用 TensorRT 加速
Python 侧的优化完成后,接下来是真正的部署环节。我导出 ONNX,然后在 TensorRT 里做 FP16 推理。ONNX 导出时直接把动态轴设为 batch 维度,这样线上可以根据实际流量调整 batch size。
dummy_input = torch.randn(1, 3, 224, 224) torch.onnx.export( model_int8_student, dummy_input, "model_optimized.onnx", opset_version=17, input_names=["input"], output_names=["output"], dynamic_axes={"input": {0: "batch"}, "output": {0: "batch"}}, )TensorRT 读取 ONNX 后的优化结果让我很满意。FP16 推理延迟 1.8ms,INT8 量化后约 1.5ms,吞吐比原生 PyTorch 高了接近 3 倍。显存占用也从 246MB 降到了 78MB,单个服务在 16G 显存卡上能同时跑 40 路进程,线上扩容压力大幅下降。
3.6 最终数据对比
| 阶段 | 精度 | 延迟 | 显存 | 模型体积 |
|---|---|---|---|---|
| 原始 FP32 | 76.3% | 6.2ms | 246MB | 98MB |
| PTQ INT8 | 75.4% | 4.1ms | 164MB | 26MB |
| 剪枝 + 微调 + PTQ | 73.6% | 3.2ms | 132MB | 19MB |
| 蒸馏 + PTQ + TensorRT | 74.6% | 1.5ms | 78MB | 18MB |
这份对比表我每次分享都会给出来。因为优化最怕的就是单看某个指标漂亮,忽略整体。精度、延迟、显存、体积必须放在一起评估,Model-Optimizer 最终目标是在这四个维度上同时拿回一个可接受的平衡点。
4. 优化路上的坑与排查记录
4.1 精度掉点最多的不是权重,是激活
做 PTQ 时大家习惯盯着权重量化误差看,但真正影响精度的是激活值。某些层的激活范围特别宽,直接 min/max 校准会把分布中间的密集区域压坏。我排查的时候会在校准前后分别把每一层的输出分布打出来,对比 KL 散度。发现的规律是:检测模型和分割模型里的 head 层,以及 Transformer 里的 attention 层,往往是量化敏感层,需要 skip 量化或者单独调高 bit 数。
4.2 导出的 ONNX 图有冗余子图
PyTorch 导出 ONNX 时,有时候会把一些 Python 控制流展开成奇怪的条件分支,或者生成形状推理不出来的动态 reshape。TensorRT 加载这样的图会直接报错或者性能奇差。我的排查思路是先用 onnx.shape_inference 做形状推断,再用 onnxsim 做简化,最后在 Netron 里肉眼检查关键子图。特别是直接拿 HuggingFace 模型导出时,必须记得先关掉不需要的输出头,再把 attention mask 之类的输入精简掉。
4.3 GPU 上的延迟不稳定
量化后延迟平均值很好看,但 p99 一直在跳。排查后发现是显存分配器的问题,PyTorch 默认的 caching allocator 在动态 shape 下频繁扩缩,导致碎片。解决办法是固定输入尺寸、开启 cudnn.benchmark,并且在服务启动时跑一遍 warmup,把显存池预占好。TensorRT 那边就简单一些,显存池只要设置好 workspace size 就行。
4.4 和线上 CPU 推理的兼容问题
有些场景跑在纯 CPU 上,ONNX Runtime 和 PyTorch 对 INT8 算子的支持不一致,量化后的模型在 ONNX Runtime 里会 fallback 成 FP32,延迟一点没降下来。解决方案是开启 ONNX Runtime 的 graph optimization level 到 ORT_ENABLE_EXTENDED,并检查 execution provider 日志里每个算子的实际执行后端。如果关键算子没有对应的 INT8 kernel,就得考虑手动写 custom op,或者换一个推理框架。
4.5 常见问题速查表
| 表现 | 可能原因 | 处理方式 |
|---|---|---|
| 量化后精度暴跌 | 激活分布过宽 | 检查敏感层,做 per_channel 量化或 skip |
| 剪枝后 loss 震荡 | BN 统计量未重置 | 微调前用校准集过一遍模型 |
| 导出 ONNX 失败 | trace 时包含动态控制流 | 用 torch.jit 先 script,再 export |
| TensorRT builder 报错 | 算子版本过新 | 升级 TensorRT 或改用 ONNX Runtime |
| INT8 延迟未下降 | 算子 fallback 到 FP32 | 检查 provider 日志,确认 INT8 kernel 生效 |
| 多进程显存超限 | 显存池没有复用 | 预热固定尺寸输入,打开显存池复用 |
5. 上线后的维护与扩展
优化不是一次性的事,模型上线之后还会继续迭代。我养成了一个习惯:每次训练新版本模型,都会跑一遍 Model-Optimizer 流水线,把量化、剪枝、蒸馏的配置参数和结果记录到一个 JSON 里,方便后面复盘。配置里除了常规的超参,还有每条优化路径的耗时和收益,这样能快速判断下一次改动是值得的还是该放弃。
另外值得说的是,这套优化链路不只适用于图像分类。我在文本分类、目标检测、语音识别这些小模型上也跑过同样的流程,基本思路完全一致,只是算子层面对齐的东西不太一样。检测模型要额外注意 anchor 和 NMS 的处理,文本模型要保留下 attention mask 的输入,语音模型对延迟和抖动更敏感。核心原则不变:先量化,再剪枝,还不行就蒸馏,最后统一导出到高性能后端。
Model-Optimizer 目前已经在我们内部沉淀成一套半自动化的流水线脚本,新手同学跑一个命令就能拿到优化后的模型和对比报告。整个项目给我最大的体会是:模型优化没有银弹,每一步都要用数据说话。别人跑出来的好结果不一定适合你的模型和硬件,只有把基线测准、把每个环节的收益和损失量化出来,才能真正越优化越有底气。