模型优化与部署

5 minAdvanced2026/6/14

模型量化、剪枝、知识蒸馏、ONNX、TensorRT部署优化。

1. 模型优化概述

模型优化旨在在保持精度的前提下,减少模型大小、降低延迟、减少功耗。

精度与效率是模型优化的两端,量化、剪枝、蒸馏等技术在两者之间寻找平衡:压缩模型体积、加速推理速度,以尽可能小的精度损失换取效率提升。

优化维度目标方法
模型大小减少存储和内存量化、剪枝、蒸馏
推理延迟减少响应时间算子融合、量化
吞吐量增加QPS批处理、并行
功耗降低能耗量化、稀疏计算

2. 量化

2.1 量化原理

将浮点参数映射到低精度整数:

xq=round(xS)+Zx_q = \text{round}\left(\frac{x}{S}\right) + Z

其中 SS 为缩放因子(Scale),ZZ 为零点(Zero Point)。

反量化

xS(xqZ)x \approx S \cdot (x_q - Z)

2.2 量化

说明精度损失难度
训练后量化(PTQ)训练完成后量化
量化感知训练(QAT)训练时模拟量化
动态量化运行时量化权重
静态量化离线校准量化参数

2.3 PTQ(训练后量化)

import torch.quantization as quant

# 静态量化
model.eval()
model.qconfig = quant.get_default_qconfig('fbgemm')
prepared = quant.prepare(model)
# 用校准数据运行
with torch.no_grad():
    for batch in calib_loader:
        prepared(batch)
quantized = quant.convert(prepared)

2.4 QAT(量化感知训练)

model.train()
model.qconfig = quant.get_default_qat_qconfig('fbgemm')
prepared = quant.prepare_qat(model)

# 正常训练
for epoch in range(num_epochs):
    for batch in train_loader:
        loss = train_step(prepared, batch)

quantized = quant.convert(prepared)

伪量化:在前向传播中模拟量化效果,反向传播使用STE(Straight-Through Estimator)。

2.5 量化精度对比

精度每参数比特相对FP32内存典型精度损失
FP3232bit
FP1616bit0.5×<1%
INT88bit0.25×1~3%
INT44bit0.125×3~8%
二值1bit0.03×>10%

3. 剪枝

3.1 剪枝

粒度说明
非结构化剪枝单个权重灵活但硬件不友好
结构化剪枝通道/层硬件友好,实际加速

3.2 剪枝策略

幅度剪枝:移除绝对值最小的权重

maskij={1wij>τ0wijτ\text{mask}_{ij} = \begin{cases} 1 & |w_{ij}| > \tau \\ 0 & |w_{ij}| \leq \tau \end{cases}

迭代剪枝

1. 训练模型至收敛
2. 剪枝:移除p%最小权重
3. 微调:恢复精度
4. 重复2-3直到目标稀疏度

3.3 结构化剪枝

# 基于L1范数的通道剪枝
import torch.nn.utils.prune as prune

# 非结构化
prune.l1_unstructured(model.fc1, name='weight', amount=0.3)

# 结构化(整个通道)
prune.ln_structured(model.fc1, name='weight', amount=0.3, n=2, dim=0)

3.4 稀疏训练

RigL:周期性移除和生长权重

1. 前向+反向传播
2. 移除最小幅度的权重
3. 在梯度最大的位置生长新权重
4. 保持恒定稀疏度

4. 知识蒸馏

4.1 基本原理

用**大模型(教师)指导小模型(学生)**训练:

L=αLCE(y,ys)+(1α)T2LKL(σ(zt/T),σ(zs/T))\mathcal{L} = \alpha \cdot \mathcal{L}_{CE}(y, y_s) + (1-\alpha) \cdot T^2 \cdot \mathcal{L}_{KL}(\sigma(z_t/T), \sigma(z_s/T))

符号含义
ysy_s学生模型输出
zt,zsz_t, z_s教师/学生logits
TT温度参数
α\alpha损失权重

4.2 温度Softmax

σ(zi/T)=ezi/Tjezj/T\sigma(z_i/T) = \frac{e^{z_i/T}}{\sum_j e^{z_j/T}}

  • T>1T > 1:软化概率分布,传递更多”暗知识”
  • T=1T = 1:标准Softmax

4.3 蒸馏变体

方法蒸馏内容说明
响应蒸馏最终输出logits最基本
特征蒸馏中间层特征传递结构信息
关系蒸馏样本间关系传递相似性结构
自蒸馏同一模型不同层无需教师模型

5. ONNX与部署

5.1 ONNX导出

import torch.onnx

dummy_input = torch.randn(1, 3, 224, 224)
torch.onnx.export(
    model, dummy_input, "model.onnx",
    input_names=['input'],
    output_names=['output'],
    dynamic_axes={'input': {0: 'batch'}, 'output': {0: 'batch'}},
    opset_version=14
)

5.2 ONNX Runtime

import onnxruntime as ort

session = ort.InferenceSession("model.onnx")
input_name = session.get_inputs()[0].name
output = session.run(None, {input_name: input_data})

5.3 TensorRT

# ONNX → TensorRT
import tensorrt as trt

logger = trt.Logger(trt.Logger.WARNING)
builder = trt.Builder(logger)
network = builder.create_network(1 << int(trt.NetworkDefinitionCreationFlag.EXPLICIT_BATCH))
parser = trt.OnnxParser(network, logger)

with open("model.onnx", "rb") as f:
    parser.parse(f.read())

config = builder.create_builder_config()
config.set_memory_pool_limit(trt.MemoryPoolType.WORKSPACE, 1 << 30)  # 1GB
engine = builder.build_serialized_network(network, config)

5.4 部署方案对比

方案延迟优化精度易用性适用场景
PyTorch原生基准FP32开发调试
ONNX Runtime2~3xFP32/FP16/INT8通用部署
TensorRT3~10xFP16/INT8NVIDIA GPU
OpenVINO3~8xFP16/INT8Intel CPU/GPU
TFLite2~5xINT8移动端
Core ML2~5xFP16/INT8Apple设备