第 12 题:如何实现模型的量化感知训练(QAT)?torch.ao.quantization的workflow是什么?
题目
如何实现模型的量化感知训练(QAT)?torch.ao.quantization的workflow是什么?
完整讲解
一、QAT 在做什么?
量化感知训练(QAT):在训练阶段就模拟量化(float → 量化再反量化),让权重和激活在「舍入 + 缩放」的误差下更新,这样训出的模型在真正部署成 INT8 时精度损失更小。和 PTQ(训练后量化) 的区别是:PTQ 不训练、只校准;QAT 会反向传播、更新参数,专门适应量化误差。
二、核心组件(torch.ao.quantization)
- QuantStub / DeQuantStub:在模型入口放 QuantStub、出口放 DeQuantStub,表示「从这里开始模拟量化」「到这里结束模拟量化」;推理时会被替换成真实 quantize/dequantize。
- QConfig:描述「用什么量化配置」:如
activation 用 MinMaxObserver 还是 MovingAverageMinMaxObserver、weight 用 per-channel 还是 per-tensor、dtype 选 qint8 等;通过 torch.ao.quantization.get_default_qconfig('qnnpack') 或自定义。
- prepare_qat:把 float 的 Module 改成「QAT 模式」:在 Conv/Linear 等前后插入 fake quantize 模块(前向时做 round+scale 模拟量化,但用 float 算,梯度可传),并挂上 observer 收集 min/max 等统计量。
- 训练:正常 forward/backward,梯度会穿过 fake quantize;observer 会更新 running min/max(若用 moving average)。
- convert:训练完后
convert(model),把 fake quantize 和 observer 换成真实的 quantize / dequantize / packed params(如 int8 权重的 Linear),得到可部署的量化模型。
三、典型 workflow(PyTorch 官方风格)
- 定义 float 模型,在入口/出口加
QuantStub()、DeQuantStub()(若用 eager 模式)。
- 设 qconfig:
model.qconfig = get_default_qconfig('qnnpack')(或 ‘fbgemm’ for server)。
- prepare_qat:
model_prepared = prepare_qat(model, inplace=False),得到带 fake quant 和 observer 的模型。
- 训练若干 epoch:正常
loss.backward() 等,observer 会更新。
- 转成推理量化模型:
model_prepared.eval() 后 model_quantized = convert(model_prepared),得到真正 int8 的模型,可保存、在 C++/ONNX 里用。
(若用 FX Graph Mode,则用 prepare_qat_fx、convert_fx,不依赖 Stub,按图自动插入 fake quant。)
四、要点小结
- Fake quantize:前向 = 量化再反量化(round + scale),用 float 算,所以能求导;这样 loss 会感受到量化误差并更新参数。
- Observer:在 QAT 里顺带统计 min/max(或直方图),convert 时用这些统计量确定 scale/zero_point;QAT 常用 moving average,避免初期统计抖动。
- Per-channel / per-tensor:权重常用 per-channel 减轻通道间分布差异;激活常用 per-tensor 或 per-channel 视后端支持。
面试要点
- QAT = 训练时用 fake quantize 模拟量化,让梯度在量化误差下更新;部署时 convert 成真实 int8。
- workflow:float 模型 + qconfig → prepare_qat(插入 fake quant 和 observer)→ 训练 → convert 得到量化模型。
- torch.ao.quantization:QuantStub/DeQuantStub、get_default_qconfig、prepare_qat、convert;FX 用 prepare_qat_fx/convert_fx。
记忆要点
- QAT = 训练阶段模拟量化(fake quant),反向传播适应量化误差;convert 后变真实 int8。
- 流程:qconfig → prepare_qat → 训练 → convert。
- Fake quant = 前向 round+scale 用 float 算可导;observer 收集统计量供 convert 用。