MQBench QAT量化实践指南

0 阅读4分钟

MQBench(Model Quantization Benchmark)是由ModelTC团队维护的开源模型量化评估与调优框架,旨在简化深度学习模型在各种硬件平台上的量化过程。它基于PyTorch构建,支持多种后端(如TensorRT、SNPE、OpenVINO等),并提供从量化感知训练(QAT)到模型部署的完整工具链。

一、什么是QAT

QAT(Quantization Aware Training,量化感知训练)是一种模型量化手段,通过在训练过的浮点模型中插入伪量化节点来实现后续的精度微调。与训练后量化(PTQ)相比,QAT通常能获得更高的精度,因为模型在训练过程中就能感知量化带来的影响,并相应地调整权重。

二、安装MQBench

首先确保Python环境为3.7或更高版本,然后执行以下命令:

git clone https://github.com/ModelTC/MQBench.git
cd MQBench
pip install -r requirements.txt

三、Naive QAT基本流程

MQBench提供了简洁的API来实现QAT,整个过程相比普通微调只多出少量额外操作。

1. 准备FP32模型

首先加载预训练的浮点模型:

import torchvision.models as models
from mqbench.prepare_by_platform import prepare_qat_fx_by_platform, BackendType
from mqbench.convert_deploy import convert_deploy
from mqbench.utils.state import enable_calibration, enable_quantization

# 加载预训练模型
model = models.__dict__["resnet18"](pretrained=True)
model.train()

2. 选择后端

MQBench支持多种硬件后端,根据部署目标选择合适的BackendType:

# 后端选项
backend = BackendType.Tensorrt      # NVIDIA TensorRT
# backend = BackendType.SNPE        # Qualcomm SNPE
# backend = BackendType.OPENVINO    # Intel OpenVINO
# backend = BackendType.Vitis       # Xilinx Vitis
# backend = BackendType.Tengine_u8  # Tengine
# backend = BackendType.ONNX_QNN    # ONNX QNN
# backend = BackendType.PPLCUDA     # PPL CUDA

3. 准备量化模型

使用prepare_qat_fx_by_platform对模型进行trace并插入伪量化节点:

# 基于选定后端为模型添加量化节点
model = prepare_qat_fx_by_platform(model, backend)

4. 校准阶段(可选但推荐)

在进行正式QAT训练前,通常先进行校准以初始化量化参数:

model.eval()
enable_calibration(model)  # 开启校准模式
for i, batch in enumerate(calibration_data):
    # 执行前向传播,收集统计信息
    model(batch)

5. QAT训练阶段

校准完成后,切换至量化训练模式,进行正常的训练循环:

model.train()
enable_quantization(model)  # 开启量化训练模式
for i, batch in enumerate(train_data):
    output = model(batch)
    loss = criterion(output, target)
    loss.backward()
    optimizer.step()

6. 导出量化模型

训练完成后,使用convert_deploy导出可部署的量化模型:

# 定义用于模型导出的虚拟输入形状
input_shape = {'data': [10, 3, 224, 224]}
convert_deploy(model, backend, input_shape)

四、完整代码示例

以下是一个完整的Naive QAT示例:

import torch
import torchvision.models as models
from mqbench.prepare_by_platform import prepare_qat_fx_by_platform, BackendType
from mqbench.convert_deploy import convert_deploy
from mqbench.utils.state import enable_calibration, enable_quantization

# 1. 准备FP32模型
model = models.__dict__["resnet18"](pretrained=True)
model.train()

# 2. 选择后端
backend = BackendType.Tensorrt

# 3. 准备量化模型
model = prepare_qat_fx_by_platform(model, backend)

# 4. 校准阶段
model.eval()
enable_calibration(model)
# 假设calibration_loader是校准数据加载器
for images, _ in calibration_loader:
    model(images)

# 5. QAT训练
model.train()
enable_quantization(model)
optimizer = torch.optim.SGD(model.parameters(), lr=0.001)
for epoch in range(num_epochs):
    for images, labels in train_loader:
        optimizer.zero_grad()
        outputs = model(images)
        loss = criterion(outputs, labels)
        loss.backward()
        optimizer.step()

# 6. 导出模型
input_shape = {'data': [1, 3, 224, 224]}
convert_deploy(model, backend, input_shape)

五、高级用法:分离量化参数优化

在QAT中,量化参数(如scale和zero_point)和模型权重参数可以设置不同的学习率,以获得更好的收敛效果:

from mqbench.nn.intrinsic.qat.modules import SomeFakeQuantize

normal_params = []
quantization_params = []

for name, module in model.named_modules():
    if isinstance(module, SomeFakeQuantize):
        quantization_params.extend(module.parameters())
    else:
        normal_params.extend(module.parameters())

optimizer = torch.optim.SGD([
    {'params': quantization_params, 'lr': quant_lr},  # 量化参数使用特定学习率
    {'params': normal_params, 'lr': normal_lr}        # 权重参数使用另一学习率
], lr=default_lr)

六、目标检测模型的QAT

对于目标检测等复杂模型,MQBench在United-Perception项目中提供了完整的QAT配置示例。核心步骤包括:

  1. 在self.build_model()中构建浮点模型
  2. 在self.load_ckpt()中加载预训练权重
  3. 使用torch.fx在self.quantize_model()中trace模型
  4. 在self.calibrate()中执行PTQ校准和评估
  5. 在self.train()中进行QAT训练

配置文件中的关键参数包括:

  • deploy_backend:选择部署后端
  • ptq_only:设为False以执行QAT
  • extra_qconfig_dict:量化配置
  • resume_model:预训练模型路径

七、注意事项

  1. 模型分离:对于目标检测等模型,应将网络主体与后处理分离,torch.fx仅trace网络部分
  2. 检查点保存:量化模型应以qat为键保存,便于后续恢复
  3. EMA处理:QAT中建议禁用EMA;若检查点包含EMA状态,会在加载时将其合并到模型中
  4. 可学习参数:若量化模型包含额外可学习参数(如LSQ),需在优化器中正确配置

通过以上步骤,你可以使用MQBench高效地完成模型的QAT量化,在保持精度的同时获得显著的推理加速和模型体积缩减。如需更详细的配置说明,可参考MQBench官方文档中的Learn MQBench configuration章节。