量化感知训练(QAT)与量化微调:伪量化、STE 与 QLoRA 实战

训练后量化(PTQ)在精度敏感场景常常失手,量化感知训练(QAT)则让模型在训练阶段就『学会』适应量化噪声,是工业界把 INT8 精度损失压到 1% 以内的关键手段。本文从伪量化算子与直通估计器(STE)讲起,给出 QAT 完整流程,讲解蒸馏+量化组合,并深入 LoRA 量化微调(QLoRA)与工业实践。一、QAT 背景与动机

训练后量化(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 指令微调中尤为显著。

对比PTQQAT
训练需求无需要标注/无监督数据
时间成本分钟级训练时长(数小时~数天)
精度损失1~5%(敏感层更高)<1%(多数场景)
适用快速验证、非敏感任务精度敏感、INT4、端侧量产
工具TFLite Converter / ORTPyTorch 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)**把两条技术叠加:

  1. 基础模型权重冻结并量化(如 NF4 4-bit),大幅降低显存占用
  2. 只训练低秩适配器(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 QATPyTorch → ONNX/TensorRT/ExecuTorch生态最全、与蒸馏/LoRA 集成
TensorFlow QATTF → TFLite与 TFLite Delegate 无缝
TensorRT QAT 校准NVIDIA 部署Q/DQ 节点直读
Intel Neural Compressor多框架自动搜索量化方案
bitsandbytes / AutoGPTQLLM 4-bitQLoRA 训练、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 → 部署
蒸馏+量化软目标拟合教师,双管齐下恢复精度
QLoRA4-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 微调一个大模型验证显存收益。

继续阅读

探索更多技术文章

浏览归档,发现更多关于系统设计、工具链和工程实践的内容。

全部文章 返回首页

「ai」更多文章

  1. 时序预测实战:从 ARIMA 到时序基础模型
  2. 模型压缩:量化、剪枝、蒸馏与部署优化实战
  3. 模型评估与基准:从分类指标到 LLM-as-Judge