Skip to content

模型优化实战:量化、蒸馏与剪枝的落地路线图

本页速览 模型上线前怎么瘦身提速?本文给出实操路线图:先量化(PTQ 起手)→ 精度不够再上 QAT/蒸馏 → 需要时剪枝,每步配工具与检查点,最后用一张决策表收尾。

模型优化(model optimization)是在不牺牲可接受精度的前提下,让模型更小、更快、更省资源的一组工程技术。 为什么要先动这个手?因为绝大多数模型部署项目里,瓶颈根本不是算力不够,而是「模型太大、太慢」的默认状态。一张 224×224 的 MobileNetV2 FP32 需要 14MB 和约 6 GFLOPs,量化到 INT8 后体积减 4 倍、GPU 上速度通常提升 1.5-2 倍、显存占用降 4 倍;而很多时候你什么都不用学,只要按顺序试一遍工具,就能白拿 50% 的收益。这篇解决「从哪下手、按什么顺序、怎么判断该停」——核心结论:先 PTQ,精度不达标再 QAT/蒸馏,需要极致性能再上引擎,每个环节都有明确的检查点与回退条件。理论基础见 量化,经典论文见 量化经典论文,优化效果要上线验证,配套 压测与容量规划

一句话记住优化顺序

90% 的优化收益来自前两步:定基线 + PTQ 量化。剩下的是增量工程。

一、先定目标,再动手:优化目标拆解

优化的本质是四选一(或按权重组合),不同目标方向完全相反,不先定目标必然做废:

目标关键指标典型手段常见反面教训
体积模型文件 MB、容器镜像大小量化(FP32→INT8)、剪枝为了省 30MB 折腾剪枝,导致精度掉 5 个点
延迟P99 / 单请求耗时量化、TensorRT、批处理调优只优化 P50 不管 P99,长尾请求照样超时
吞吐QPS、每秒 tokens连续批处理(batching)、并发调优单请求快但并发一高就崩
功耗/成本每千次推理成本、瓦特量化 + 低功耗硬件选择GPU 吃满 300W 跑一个本可用 CPU 扛的小模型

三个硬性原则:

  1. 先量化再剪枝:量化几乎零工程成本(几分钟),剪枝要重训;量化通常先试。
  2. 延迟和吞吐分开测:量化对单请求延迟帮助有限,对吞吐(batch 变大了)帮助巨大——别拿单请求延迟的失望去否定量化。
  3. 目标是满足 SLO,不是压榨到极限:P99 < 100ms 达标了就停,多出来的优化是负 ROI。参考 性能指标 里的 SLO 定义。

二、优化路线图:一条可回退的流水线

text
  基线测量 ──▶ PTQ 量化 ──▶ 精度评估 ──▶ 达标? ──▶ 上线压测验证
    │              │             │          │
    │              │             └──不达标──▶ QAT / 蒸馏 ──▶ 再评估
    │              │                                   │
    │              │                                   └──需要极致性能──▶ TensorRT / 引擎
    │              └──部分算子不支持──▶ 混合精度(FP16/INT8 分段)
    └──延迟不达标但精度敏感──▶ 从剪枝/蒸馏入手

每个节点都有「回退」:量化后精度掉了 → 回退到 FP32 基线并换策略;QAT 训练成本太高 → 回到蒸馏。这条路线图的价值在于把「拍脑袋优化」变成「有证据的决策流」。每一步的记录格式见本文「优化效果追踪表」。

三、步骤一:基线测量——没有基线,一切优化都是自嗨

优化前先花 30 分钟做一次规范测量,产出三张数字:延迟(P50/P99)、显存/内存峰值、精度指标。这份基线就是后面所有对比的锚点。

python
# baseline.py —— 用标准脚本测基线,避免手工计时误差
import time
import torch
import onnxruntime as ort
import numpy as np

def measure_ort(onnx_path: str, n: int = 200, warmup: int = 20):
    """测 ONNX Runtime 的延迟分布(单位 ms)。"""
    sess = ort.InferenceSession(onnx_path, providers=["CPUExecutionProvider"])
    x = np.random.rand(1, 3, 32, 32).astype(np.float32)
    for _ in range(warmup):                     # 预热,跳过 JIT/缓存热身
        sess.run(None, {"input": x})
    lats = []
    for _ in range(n):
        t0 = time.perf_counter()
        sess.run(None, {"input": x})
        lats.append((time.perf_counter() - t0) * 1000)
    lats.sort()
    return {"p50": lats[n // 2], "p99": lats[int(n * 0.99) - 1], "mean": sum(lats) / n}

基线记录表模板(建议存成 perf_baseline.md 随仓库提交):

记录示例
环境CPU/GPU 型号、驱动、引擎版本Intel i7-12700 / ONNX Runtime 1.19.2 CPU
输入规格shape、batch、数据类型(1,3,32,32) FP32
延迟 P50/P99ms8.2 / 14.5
峰值显存/内存MB220 MB
模型体积MB14.0
离线精度验证集指标accuracy 0.812
线程/批处理配置threads、batch4 threads, batch=1

别用「看起来挺快」代替测量

手工 time.time() 包一次推理,误差能到 ±30%,因为没预热、时钟精度不够。用上面的脚本形式(预热 + 多次采样 + 分位数)才是可对比的数字。

四、步骤二:PTQ 量化实操(先试这个)

PTQ(post-training quantization)不需要重训,是性价比最高的一步。两个主流入口:ONNX Runtime 的 quantize_static(CPU 快)和 PyTorch 的 torch.ao.quantization

4.1 静态量化与 calibration 数据

静态量化(static quantization)需要一小批「校准数据」来确定每个张量的数值范围(scale/zero_point)。校准数据必须来自真实输入分布,用随机噪声校准会让量化误差凭空翻倍——这是新手最容易做错的细节。

python
# quantize_ptq.py —— ONNX Runtime 静态量化
import onnx
from onnxruntime.quantization import quantize_static, CalibrationMethod
from onnxruntime.quantization.calibrate import CalibrationDataReader

class CalibReader(CalibrationDataReader):
    """从验证集取 200 张图作为校准数据,输入名必须与导出时一致。"""
    def __init__(self, dataloader, input_name="input"):
        self.iter = iter(dataloader)
        self.input_name = input_name
    def get_next(self):
        try:
            x, _ = next(self.iter)
            return {self.input_name: x.numpy()}
        except StopIteration:
            return None

# 200 张校准图,覆盖各类别分布;数量不是越多越好,256 张以内足够
calib = CalibReader(DataLoader(ds_eval, batch_size=32)[:200])

quantize_static(
    "models/model.onnx", "models/model_int8.onnx",
    calib, per_channel=True,
    calibration_method=CalibrationMethod.MinMax,   # MinMax 默认,ENTROPY 对分布敏感模型更稳
)
print("已导出 INT8 模型,体积应约为 FP32 的 1/4")

4.2 量化后立刻做三件事

bash
# 1) 体积:应当 ≈ FP32 的 25%(14MB -> 3.5MB)
ls -lh models/model.onnx models/model_int8.onnx
# 2) 精度:跑一遍验证集,和基线表对比
python evaluate.py --model models/model_int8.onnx   # 期望 accuracy 掉 < 1%
# 3) 延迟/显存:重跑 baseline.py,和基线表对比
python baseline.py --model models/model_int8.onnx

量化收益的预期值(经验数据)

  • 体积:INT8 ≈ FP32 的 1/4,FP16 ≈ 1/2;
  • 延迟:CPU 上 INT8 通常快 1.5-2 倍;GPU 上取决于算子是否被加速库覆盖;
  • 精度:分类/检测模型 PTQ 后通常掉 0.5-2%;如果掉 5% 以上,先查校准数据是不是用错了,而不是急着上 QAT。

五、步骤三:精度评估与回退判断

量化后唯一重要的问题是:这个精度损失可以接受吗? 判定标准不是「和 FP32 比差多少」,而是「是否跌破业务阈值」。

场景判定动作
accuracy 掉 ≤ 1%达标直接上 INT8,进入压测验证
掉 1%-5%边界先换校准方法(MinMax→Entropy)、校准数据,无效再上 QAT
掉 > 5%不达标回退 FP32,走蒸馏或剪枝路线
python
# evaluate.py —— 精度对比的自动化版本,输出结构化结果
import onnxruntime as ort
import numpy as np

def eval_onnx(path, loader):
    sess = ort.InferenceSession(path, providers=["CPUExecutionProvider"])
    correct = total = 0
    for x, y in loader:
        pred = sess.run(None, {"input": x.numpy()})[0].argmax(1)
        correct += (pred == y.numpy()).sum().item()
        total += y.size(0)
    return correct / total

print("FP32:", eval_onnx("models/model.onnx", loader))
print("INT8:", eval_onnx("models/model_int8.onnx", loader))

回退是正常流程,不是失败。把评估脚本固化进 CI,量化改版时自动对比基线——这也是 MLOps 流水线 里「模型验证」环节的雏形。

六、步骤四:QAT 与蒸馏(精度不够时的两条路)

PTQ 不达标时,两条主线:

6.1 QAT(量化感知训练)

QAT 把「量化误差」模拟进训练过程(fake-quant 算子),让权重学会在量化后仍保持精度。PyTorch 官方路径(torch.ao.quantization):

python
# qat.py —— 只列关键流程,完整代码见 PyTorch 文档
import torch
from torch.ao.quantization import (
    get_default_qat_qconfig_mapping, QATQuantizer, prepare_qat, convert)

qat = QATQuantizer(model)                       # 或手动 prepare_qat(model, qconfig)
qat.qconfig = get_default_qat_qconfig_mapping("x86")   # 按部署后端选 qconfig
model_qat = prepare_qat(model, inplace=False)
# 用训练集再训 3-5 个 epoch(学习率建议缩到原来的 1/10)
# ...
model_int8 = convert(model_qat)                 # 训练完后转成真正量化模型

QAT 的代价是训练时间 + 需要训练数据与 GPU,收益是比 PTQ 多保住 1-3% 精度。经验判断:PTQ 掉 3% 以上时 QAT 值得上,掉 1% 时 QAT 的 ROI 很低。

6.2 蒸馏(知识蒸馏)

用大模型(teacher)的软标签教小模型(student),让小模型在更小体积下逼近大模型精度。蒸馏和量化是正交的:可以蒸馏出一个更小的 FP32 模型,再对它做 INT8 量化。原理与工具见 模型压缩。典型收益:同样 80% 精度,蒸馏后的模型体积再小一个量级(ResNet-50 级 teacher 蒸馏出 MobileNet 级 student)。

python
# distill.py —— 关键代码:loss 同时吃 hard label 和 teacher 软标签
alpha, T = 0.7, 3.0          # 软标签权重 0.7,温度 3
kl = nn.KLDivLoss(reduction="batchmean")
with torch.no_grad():
    t_logits = teacher(x)
student_loss = F.cross_entropy(s_logits, y)
distill_loss = kl(F.log_softmax(s_logits / T, dim=1),
                  F.softmax(t_logits / T, dim=1))
loss = alpha * T * T * distill_loss + (1 - alpha) * student_loss

蒸馏别跳过目标定义

蒸馏前先问:student 的部署约束是什么(多少内存?什么引擎?什么延迟预算?)。先定约束再挑 teacher/student,否则蒸馏完发现 student 在目标硬件上优势不成立,白费一周。

七、步骤五:TensorRT 等引擎级优化

量化是「模型层面」的瘦身,引擎是「执行层面」的加速:层融合(layer fusion)、算子内核选择、显存规划。TensorRT 通常在 NVIDIA GPU 上比 ONNX Runtime 再快 30-60%,INT8 + 引擎双管齐下是 GPU 上的常见终态。完整流程(trtexec 命令行、engine 序列化、与硬件绑定)见 TensorRT 边缘部署

bash
# trtexec 从 ONNX 生成 FP16 引擎并测吞吐
trtexec --onnx=models/model.onnx \
        --saveEngine=models/model_fp16.engine \
        --fp16 \
        --shapes=input:1x3x32x32 \
        --avgRuns=100

引擎级优化的三个前提:

  1. 先量化再引擎:TensorRT 内部也可以做 INT8,但外部先量化的模型更可控;
  2. 引擎与硬件绑定.engine 文件与 GPU 型号、TensorRT 版本强相关,换卡必须重新 build——这是 模型格式 里「engine 是不可移植格式」的直接后果;
  3. 引擎只解决「算得快」,不解决「喂得饱」:输入管道的开销(解码、预处理)才是 CPU 服务的新瓶颈,压测时看整体 P99,别只看引擎内耗时。

八、每步检查清单(该记录什么指标)

步骤必须记录常见遗漏
基线测量环境、P50/P99、显存、精度、模型体积忘了记录 batch size 和线程数
PTQ 量化校准数据来源与数量、校准方法、per-channel 与否校准数据用了随机噪声
精度评估FP32 vs INT8 对比、业务阈值只看相对差,不看是否破业务底线
QAT/蒸馏训练 epoch、学习率、teacher 配置没记录训练成本(GPU 小时数)
引擎优化引擎版本、GPU 型号、build 参数没记录引擎与硬件的绑定关系

九、优化效果追踪表模板

把整条流水线的数字汇总到一张表,随模型版本走(对应 MLOps 流水线 的模型注册信息):

版本方案体积P50(ms)P99(ms)显存(MB)精度决策
v1FP32 基线14.08.214.52200.812基线
v2INT8 (MinMax)3.64.88.9960.806上线
v3INT8 + TensorRT3.62.95.4880.806上线

判定规则:只有「精度在业务阈值内」且「延迟/吞吐满足 SLO」的方案才进上线队列;两个条件缺一,就回退上一行。

十、决策表:什么情况选哪种优化

你的处境推荐路径理由
首次优化、时间紧PTQ 量化 → 压测零重训、几小时出结果
PTQ 后精度掉 > 3%QAT(有训练资源)比 PTQ 多保 1-3% 精度
模型体积是硬约束(边缘/移动端)蒸馏 → 量化体积能再降一个量级
精度敏感且 PTQ/QAT 都失效剪枝 + 重训结构性减肥,代价最大放最后
NVIDIA GPU、延迟敏感量化 + TensorRT引擎级 30-60% 额外提速
LLM 场景GPTQ/AWQ + vLLM权重量化是 LLM 的标配,见 LLM 推理

剪枝为什么放最后

结构化剪枝(结构化稀疏)需要重训、工具链成熟度低、收益往往被量化覆盖。先量化后剪枝是社区共识;把剪枝当作「量化解决不了体积约束」时的兜底方案,而不是起点。

权衡与取舍

优化的本质是拿「精度/工程成本」换「体积/延迟/成本」。三个容易被忽略的取舍:校准数据是训练数据的一部分,涉及数据合规;引擎绑定硬件会削弱可移植性(换云厂商机型就得重新 build);量化后的模型难以再微调,模型迭代流程要重跑一遍量化管线——所以量化管线本身要脚本化、进 CI,而不是每次手工操作。

延伸阅读

参考资料