训练后量化(PTQ)在精度敏感场景常常失手,量化感知训练(QAT)则让模型在训练阶段就"学会"适应量化噪声,是工业界把 INT8 精度损失压到 1% 以内的关键手段。本文从伪量化算子与直通估计器(STE)讲起,给出 QAT 完整流程,讲解蒸馏+量化组合,并深入 LoRA 量化微调(QLoRA)与工业实践。
一、QAT 背景与动机
1.1 PTQ 的精度瓶颈
https://plumephp.com/ai-quantization/ 中介绍的训练后量化(PTQ)流程是:训练好模型 → 校准激活范围 → 把权重转成 INT8。好处是无需训练、速度极快,但对三类场景经常失手:
- 敏感层:某些层的权重分布极不均匀,INT8 表示误差被放大
- 低比特:INT4/W4A16 场景下 PTQ 的误差急剧上升
- 长尾分布:激活中少量异常值把量化范围撑得很大,有效精度被稀释
PTQ 是"一次性近似":它无法让模型适应量化带来的信息损失。精度一旦掉太多,回天乏术。
1.2 QAT 的核心思想
**量化感知训练(Quantization-Aware Training,QAT)**的思路正好相反:在训练过程中就模拟量化误差,让模型权重在梯度下降中学习抵消这种误差。
PTQ:训练(FP32) ──> 校准 ──> INT8(一次近似,误差不可逆)
QAT:训练(FP32 + 伪量化模拟) ──> 移除伪量化 ──> INT8(模型已适应)
QAT 的关键不是"训练得更久",而是在正确的精度模拟下训练。它比 PTQ 的精度恢复幅度通常达到 1~5 个百分点(分类任务),在检测、语义分割、LLM 指令微调中尤为显著。
| 对比 | PTQ | QAT |
|---|---|---|
| 训练需求 | 无 | 需要标注/无监督数据 |
| 时间成本 | 分钟级 | 训练时长(数小时~数天) |
| 精度损失 | 1~5%(敏感层更高) | <1%(多数场景) |
| 适用 | 快速验证、非敏感任务 | 精度敏感、INT4、端侧量产 |
| 工具 | TFLite Converter / ORT | PyTorch QAT / TensorRT QAT |
二、伪量化算子(Fake Quant)
2.1 伪量化原理
QAT 在计算图中插入伪量化算子(FakeQuant):前向传播时执行"量化 → 反量化"回到浮点,从而让梯度"感知"量化误差;反向传播时保持浮点梯度。
FakeQuant(x) = dequant(quant(x))
= scale * clamp(round(x / scale), qmin, qmax)
scale = (x_max - x_min) / (qmax - qmin) # 对称量化
例如 INT8 对称量化:q = clamp(round(x / s), -128, 127),x_hat = q * s。前向时 x_hat ≈ x,但带有量化舍入误差——这正是希望训练去适应的"噪声"。
2.2 直通估计器(STE)
问题是:round() 的梯度几乎处处为 0,量化算子没法训练。**直通估计器(Straight-Through Estimator, STE)**的解法是:反向传播时把量化算子的梯度近似为恒等函数。
前向:y = quantize(x)(round 引入不可导)
反向:dy/dx ≈ 1(跳过 round,直接把梯度传回)
import torch
class StraightThroughQuant(torch.autograd.Function):
@staticmethod
def forward(ctx, x, scale):
# 前向:真正的量化 + 反量化
xq = torch.clamp(torch.round(x / scale), -127, 127)
return xq * scale
@staticmethod
def backward(ctx, grad_output):
# 反向:STE——直接把梯度传回
return grad_output, None
PyTorch 官方 torch.ao.quantization 内置了上述机制,无需手写 autograd 函数。
2.3 伪量化算子代码示例(PyTorch)
import torch
import torch.nn as nn
import torch.ao.quantization as tq
class QuantLinear(nn.Module):
"""带伪量化的线性层,模拟 INT8 权重+激活量化。"""
def __init__(self, in_f, out_f):
super().__init__()
self.linear = nn.Linear(in_f, out_f)
self.weight_quant = tq.FakeQuantize(
observer=tq.MinMaxObserver.with_args(dtype=torch.qint8, qscheme=torch.per_tensor_symmetric),
quant_min=-128, quant_max=127, dtype=torch.qint8,
)
self.act_quant = tq.FakeQuantize(
observer=tq.MovingAverageMinMaxObserver.with_args(dtype=torch.quint8),
quant_min=0, quant_max=255, dtype=torch.quint8,
)
def forward(self, x):
w = self.weight_quant(self.linear.weight) # 权重伪量化
y = nn.functional.linear(x, w, self.linear.bias)
return self.act_quant(y) # 激活伪量化
注意:激活量化用对称还是非对称、per-tensor 还是 per-channel,直接影响精度与硬件适配。量化方案选择细节见 https://plumephp.com/ai-quantization/。
三、QAT 完整流程
3.1 五步标准流程
① 定义量化配置(对称/非对称、范围、observer)
② prepare_qat:插入伪量化算子,模型变为"可量化训练"态
③ 微调训练(低学习率,数据可来自 PTQ 校准集或下游任务集)
④ convert:移除伪量化,产出生 INT8 权重与 scale
⑤ 端侧/引擎部署(TensorRT / TFLite / ORT 均可加载)
import torch.ao.quantization as tq
model = MyModel().eval()
model.qconfig = tq.get_default_qat_qconfig_mapping("x86") # ①
tq.prepare_qat(model, inplace=True) # ②
opt = torch.optim.Adam(model.parameters(), lr=1e-4)
for epoch in range(3): # ③
for x, y in loader:
loss = criterion(model(x), y)
opt.zero_grad(); loss.backward(); opt.step()
model = model.cpu()
tq.convert(model, inplace=True) # ④
torch.save(model.state_dict(), "int8_qat.pt") # ⑤
3.2 QAT 训练要点
| 超参/决策 | 推荐值 | 理由 |
|---|---|---|
| 学习率 | 比正常训练低 5~10 倍 | 防止扰动已收敛权重 |
| 训练轮次 | 1~3 epoch | 只需让模型"适应"而非"重学" |
| 数据 | 下游任务数据或校准集 | 覆盖生产分布即可 |
| 是否冻结 BN | 是 | 量化下 BN 统计易漂移 |
| 优化器 | Adam | 与普通微调一致 |
3.3 部署时的"移除伪量化"
convert 之后模型里不再有 FakeQuant,取而代之的是 torch.quantized 张量(保存 INT8 权重 + scale)。部署侧注意:
- INT8 权重 + scale 打包:导出的 ONNX/TFLite 会带
QuantizeLinear/DequantizeLinear节点 - 激活仍按浮点计算:真实运行时(TensorRT/TFLite 的 INT8 engine)会在 kernel 内部处理,不需要外部再量化
- 验证一致性:部署前后的 FP32 与 INT8 输出差异应小于预设阈值
四、蒸馏 + 量化组合
4.1 为什么要组合
QAT 单独用时,如果原始模型精度本就不高,量化误差仍可能突破阈值。蒸馏(Distillation)+ 量化是工业界最常见的组合:让 INT8 学生模型去拟合 FP32 教师模型的软输出,双管齐下。
FP32 教师模型 ──soft label──┐
├──> INT8 学生模型(QAT 训练)
原始硬标签 ──────────────────┘
蒸馏与剪枝的完整方法论见 https://plumephp.com/ai-model-compression/,这里重点看"蒸馏目标如何和量化损失叠加"。
4.2 蒸馏目标设计
def qat_distill_loss(student_logits, teacher_logits, hard_label, T=3.0, alpha=0.5):
import torch.nn.functional as F
# 软目标 KL 散度(带温度 T)
soft_loss = F.kl_div(
F.log_softmax(student_logits / T, dim=-1),
F.softmax(teacher_logits / T, dim=-1),
reduction="batchmean",
) * (T * T)
# 硬标签交叉熵
hard_loss = F.cross_entropy(student_logits, hard_label)
return alpha * soft_loss + (1 - alpha) * hard_loss
要点:
- 温度 T:软化教师输出分布,突出类间结构,一般 3~8
- alpha 平衡:0.5~0.9 侧重软目标;对量化误差大的模型可提高 alpha
- logits 对齐:学生与教师 logits 维度必须一致,若结构不同可用 feature distillation
五、LoRA 量化微调(QLoRA)
5.1 QLoRA 原理
大模型参数巨大,直接做 QAT 成本高昂。**QLoRA(Quantized LoRA)**把两条技术叠加:
- 基础模型权重冻结并量化(如 NF4 4-bit),大幅降低显存占用
- 只训练低秩适配器(LoRA),参数规模仅占原模型的千分之几
冻结 + 4-bit 量化基座权重(BF16 反量化到计算)
│
▼
┌─────────────┐ LoRA_A (低秩) ──┐
│ base 层输出 │ ───────────────> │ + 前向结果
└─────────────┘ LoRA_B (低秩) ──┘
显存收益:70B 模型 FP16 需 ~140GiB,NF4 QLoRA 只需 ~40GiB(单张 48G 或 2×24G 可微调)。
5.2 代码示例(transformers + peft + bitsandbytes)
from transformers import AutoModelForCausalLM, AutoTokenizer
from peft import LoraConfig, get_peft_model
import torch
model = AutoModelForCausalLM.from_pretrained(
"meta-llama/Llama-2-7B",
load_in_4bit=True, # 4-bit 量化(NF4)
bnb_4bit_compute_dtype=torch.bfloat16,
device_map="auto",
)
lora = LoraConfig(
r=16, lora_alpha=32, target_modules=["q_proj", "v_proj"],
lora_dropout=0.05,
)
model = get_peft_model(model, lora)
model.print_trainable_parameters() # 只训练 ~4M 参数
# 常规微调循环(冻结基座,只优化 LoRA 参数)
for x, y in dataset:
loss = model(x, labels=y).loss
loss.backward()
optimizer.step()
optimizer.zero_grad()
推理时同样把 LoRA 合并回基座,再进行 PTQ/QAT 压缩。注意:QLoRA 产出的还是 FP16 权重的 LoRA 适配器,下游仍要再做量化才能得到 INT8 模型。
5.3 QAT 与 QLoRA 的结合策略
| 场景 | 推荐做法 |
|---|---|
| 基座模型 FP32 → INT8 | 直接 QAT(prepare_qat → convert) |
| 基座太大,单卡装不下 | QLoRA 微调 → 合并 → 再做 QAT |
| 微调本身要量化感知 | QAT 与 LoRA 并行:基座伪量化 + 只训 LoRA |
| 精度仍不足 | 蒸馏 + QAT + 量化范围校准三管齐下 |
其中"量化感知 LoRA“是当前 LLM 低比特化的前沿做法:基座以 FakeQuant INT8/INT4 模拟参与前向,LoRA 在量化误差存在的情况下学习补偿,得到的适配器天然对量化更鲁棒。权重级量化(GPTQ/AWQ 类)可参考 https://plumephp.com/ai-llm-quantization/。
六、工业实践与工具
6.1 TensorRT QAT 流程
NVIDIA 官方推荐在 TensorRT 中走"QAT + 校准"双保险:
# PyTorch 训练后导出 QAT 模型
model.qconfig = tq.get_default_qat_qconfig_mapping("tensorrt")
tq.prepare_qat(model, inplace=True)
# ...训练...
tq.convert(model, inplace=True)
# 导出带 Q/DQ 节点的 ONNX 供 TensorRT 构建
TensorRT 在构建 INT8 engine 时识别 QuantizeLinear/DequantizeLinear 节点,直接使用训练得到的 scale,而非重新校准——这正是 QAT 比 PTQ 更可控的原因:量化范围由训练决定,而非校准集碰运气。引擎构建细节见 https://plumephp.com/ai-tensorrt/。
6.2 工具链全景
| 工具 | 适用 | 特点 |
|---|---|---|
| PyTorch QAT | PyTorch → ONNX/TensorRT/ExecuTorch | 生态最全、与蒸馏/LoRA 集成 |
| TensorFlow QAT | TF → TFLite | 与 TFLite Delegate 无缝 |
| TensorRT QAT 校准 | NVIDIA 部署 | Q/DQ 节点直读 |
| Intel Neural Compressor | 多框架 | 自动搜索量化方案 |
| bitsandbytes / AutoGPTQ | LLM 4-bit | QLoRA 训练、GPTQ 推理 |
6.3 经验总结
- 先 PTQ 摸底:PTQ 精度如果已经 >1% 损失,优先用 QAT 而不是折腾校准
- QAT 数据不求多,求准:覆盖生产分布远比数据量大重要
- 量化方案要与硬件对齐:NPU/GPU 支持 per-channel 与否,直接决定 QAT 配置(见 https://plumephp.com/ai-edge-inference-deployment/)
- 把 QAT 纳入 CI:每次训练迭代都应重跑 INT8 精度验证,防止量化退化悄悄混入
七、总结
| 知识点 | 核心要点 |
|---|---|
| QAT 动机 | 让模型在训练中适应量化噪声,精度损失 <1% |
| 伪量化算子 | 前向量化+反量化模拟,反向用 STE 传梯度 |
| 五步流程 | 配置 → prepare_qat → 微调 → convert → 部署 |
| 蒸馏+量化 | 软目标拟合教师,双管齐下恢复精度 |
| QLoRA | 4-bit 冻结基座 + 低秩适配器,单卡可微调 70B |
| 工程工具 | PyTorch QAT / TensorRT QAT / Intel Neural Compressor |
| 经验 | 先 PTQ 摸底、数据求准、与硬件量化方案对齐 |
QAT 是量化体系的"精度保险丝”:当 PTQ 达不到业务精度要求时,用几轮训练换 1% 以内的损失完全值得。理解伪量化算子与 STE,是理解一切"量化感知"机制(含量化感知 LoRA、量化蒸馏)的钥匙。建议先用 PyTorch 在 ResNet 上跑通五步流程,再对端侧模型(见 https://plumephp.com/ai-edge-inference-deployment/)实施 QAT,最后尝试用 QLoRA 微调一个大模型验证显存收益。
继续阅读
探索更多技术文章
浏览归档,发现更多关于系统设计、工具链和工程实践的内容。