MQBench QAT量化实践指南

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

一、什么是QAT

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

二、安装MQBench

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

bash 复制代码
git clone https://github.com/ModelTC/MQBench.git
cd MQBench
pip install -r requirements.txt

三、Naive QAT基本流程

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

1. 准备FP32模型

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

python 复制代码
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:

python 复制代码
# 后端选项
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并插入伪量化节点:

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

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

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

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

5. QAT训练阶段

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

python 复制代码
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导出可部署的量化模型:

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

四、完整代码示例

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

python 复制代码
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)和模型权重参数可以设置不同的学习率,以获得更好的收敛效果:

python 复制代码
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章节。

相关推荐
Είναι η κοπέλα1 小时前
llama.cpp 与 GGUF 格式:本地大模型的“裸引擎“
开发语言·人工智能·pytorch·python·conda
今夜有雨.1 小时前
图像运算、掩膜/ROI 与绘制交互
c语言·数据结构·c++·qt·算法·计算机视觉
UIU1141 小时前
P1028 [NOIP 2001 普及组] 数的计算
c++·算法
桃西西呀1 小时前
Spring AI Alibaba 之二:Spring AI 底座与原子抽象一笔带过,看清 SAA 在之上加了什么编排
人工智能·spring·llm
段一凡-华北理工大学1 小时前
高炉炼铁机器视觉与智能识别十八讲~系列文章05:炉顶料面识别:装料分布判读与布料制度优化
大数据·人工智能·机器视觉·工业智能化·高炉炼铁智能化·高炉炉顶料面识别
sbjdhjd1 小时前
智能体开始“动手”之后:OpenAI越权事件、Anthropic算力资本化与开放权重模型竞逐 | AI与SI行业日报整理(9月29日—10月6日)
大数据·人工智能·经验分享·笔记·ai·chatgpt·开源
hsfxuebao1 小时前
Loop Engineering 已死? 一文带你了解Graph Engineering
人工智能·后端
秦先生在广东1 小时前
Atlassian 携手 OpenAI 深化合作:以企业知识图谱驱动 AI 智能体在开发全流程中的应用
人工智能
品牌常新1 小时前
GEO哪家好?2026年服务商选型与能力对比
大数据·人工智能