第 5 章 量化感知训练(QAT)
第 5 章 量化感知训练(QAT)
第 4 章的 PTQ 有个天花板:模型是训练好的,量化只是"事后加工"。当位宽降到 4bit、或模型对误差极敏感时,PTQ 撑不住。这一章进入量化感知训练(Quantization-Aware Training,QAT)——把量化误差当成训练的一部分,让模型在前向传播里"模拟量化",通过反向传播学会补偿量化带来的损失。
原书仓库 ch5/ 目录的脚本非常系统:ch5_fake_quantization_ste.py、ch5_per_channel_qat.py、ch5_ptq_failure_diagnostics.py、ch5_qat_schedule.py、ch5_transformer_qat.py。
5.1 核心难点:量化不可导,怎么反向传播?
QAT 的第一步是伪量化(fake quantization):前向时把权重/激活量化再反量化(dequant),让模型"看到"量化误差;但存储的仍然是 FP32 的浮点值,这样梯度还能正常流动。
问题来了:round() 函数的导数几乎处处是 0——量化器的梯度传不过去,模型根本学不动。
解法是 STE(Straight-Through Estimator,直通估计器):反向传播时,把 round() 的梯度"绕过"——量化器的输入梯度直接等于输出的梯度。
# 来自 ch5/ch5_fake_quantization_ste.py(整理)
class FakeQuantizeSTE(torch.autograd.Function):
@staticmethod
def forward(ctx, x, scale, zero_point, q_min, q_max):
# 前向:量化 → 反量化,模拟量化误差
x_q = torch.clamp(torch.round(x / scale) + zero_point, q_min, q_max)
x_deq = (x_q - zero_point) * scale
return x_deq
@staticmethod
def backward(ctx, grad_output):
# STE:梯度直通,绕过 round() 的零梯度
return grad_output, None, None, None, NoneSTE 的本质:前向传播老老实实模拟量化(该截断截断、该取整取整),反向传播假装量化不存在(梯度直接穿过)。这是 QAT 的基石,几乎所有 QAT 框架(PyTorch、TensorFlow、ONNX 的 QDQ)都是这个套路。
ch5_fake_quantization_ste.py 里还演示了一个细节:STE 的梯度截断。如果梯度完全不截断,那些已经被 clamp 到边界以外的权重还会继续收到梯度,可能导致权重"冲出"量化范围。所以更稳妥的 STE 会在梯度上做 clip,让越界的权重不再被更新(对应 visualize_clipped_ste_gradient 实验)。
5.2 QAT 的完整配方:一个 CNN 例子
ch5_fake_quantization_ste.py 里用一个 QATSmallCNN 展示了 QAT 的全部环节:
# 来自 ch5/ch5_fake_quantization_ste.py 的结构
class MinMaxObserver(nn.Module):
"""观察权重的 min/max,算出 scale/zero_point。"""
...
class FakeQuantizedLinear(nn.Module):
"""把 nn.Linear 包一层伪量化。"""
def forward(self, x):
x_q = fake_quantize(x, ...) # 激活伪量化
w_q = fake_quantize(self.weight, ...) # 权重伪量化
return F.linear(x_q, w_q, self.bias)
class QATSmallCNN(nn.Module):
"""把 CNN 里的 Conv/Linear 换成 FakeQuantized 版本。"""
...关键设计:
- Observer(观察器):每次前向更新权重/激活的 min/max,进而算出当前的
S/Z。Observer 的更新不参与梯度(detach),它只是"统计"。 - FakeQuantized 层:前向把数据过一遍伪量化,让模型在训练时就"看见"量化误差。
- 训练流程:先用 FP32 预训练一个基线,再插入伪量化层做 QAT 微调。
visualize_weight_adaptation实验会展示:QAT 过程中权重会逐渐"适应"量化网格——权重分布会被推向量化格的"安全位置"。
5.3 per-channel QAT:给每个通道独立学缩放
ch5_per_channel_qat.py 把第 3 章的 per-channel 粒度搬进 QAT。它对比了 per-tensor 和 per-channel 两种伪量化在 Fashion-MNIST 小 CNN 上的表现:
# 来自 ch5/ch5_per_channel_qat.py(核心)
class FakeQuantizePerChannel(torch.autograd.Function):
"""per-channel 伪量化:每个输出通道独立 scale。"""
@staticmethod
def forward(ctx, w, scales, ...):
# w: [out, in],scales: [out, 1]
w_q = torch.clamp(torch.round(w / scales) + zero, q_min, q_max)
return (w_q - zero) * scales实验结论很有代表性:
- 4bit per-channel QAT 可以追平 8bit per-tensor QAT 甚至更好——细粒度缩放 + 训练补偿的组合拳。
- per-channel 的梯度比 per-tensor 更"健康":每个通道的权重有自己的缩放,训练时更新更平滑(
gradient_analysis和scale_evolution实验展示了 per-channel 的 scale 演化更稳定)。 conv_axis_analysis还研究了 per-channel 对卷积核取哪个轴——对 Conv2d,scale按输出通道(axis=0)取,这对齐了第 3 章"per-channel = 每个输出通道一个 scale"的定义。
5.4 什么时候该上 QAT:PTQ 失败诊断
QAT 有代价(要训练、要数据、要时间),所以不是无脑上。ch5_ptq_failure_diagnostics.py 做了一件非常务实的事:在 ImageNette(视觉)和 SST-2(NLP)上系统诊断 PTQ 为什么失败,从而给出"什么时候该上 QAT"的判断依据。
它把 PTQ 的失败分解成几个可定位的原因:
# 来自 ch5/ch5_ptq_failure_diagnostics.py 的诊断维度(结构)
# 1. 逐层量化误差:哪些层是误差大户?
# 2. 逐类别精度:量化后哪些类别崩了?(长尾类别通常最脆弱)
# 3. 置信度分布:量化后模型是不是变得"过度自信"或"自信崩塌"?
def quantize_single_layer(model, layer_name, config):
# 只量化某一层,观察它对整体精度的单独影响
...
def evaluate_per_class_accuracy(model, loader):
# 逐类别评估,定位崩掉的类别
...典型诊断结果:
- PTQ 失败往往是少数"问题层"造成的,不是全模型均匀恶化。定位出来后,可以只对这些层保留更高位宽,或只对它们做 QAT——这就是选择性量化,成本比全模型 QAT 低得多。
- 尾部类别最脆弱:量化后,模型对训练数据中出现少的类别(长尾)识别率下降最严重。如果你的业务依赖这些类别,PTQ 风险更高。
- 置信度漂移:量化后的模型置信度分布会变化,对需要校准概率输出的场景(如风险预估)是隐形风险。
5.5 渐进式量化调度:别一步到位
ch5_qat_schedule.py 提出了 QAT 的一个工程细节:渐进式量化调度(Progressive Quantization Schedule)。别一开始就全位宽量化,而是:
# 来自 ch5/ch5_qat_schedule.py(结构)
class ProgressiveQuantSchedule:
"""先粗后细:位宽从 8bit 逐步降到目标位宽。"""
def get_bits(self, epoch):
# 例如: epoch<3 → 8bit, epoch<6 → 6bit, 之后 → 4bit
...
def prepare_model_for_qat(model, bits=8):
# 1. BN 折叠(fold BN into Conv,QAT 前必须做)
# 2. 替换成 FakeQuantized 层
# 3. 注册 Observer
...这个脚本还包含几个 QAT 必须处理的技术点:
- BN 折叠(BatchNorm folding):推理时 BN 是线性变换,可以折进 Conv 的权重里。QAT 前必须先折叠,否则训练和推理的数值行为不一致。
fold_bn_into_conv和fold_all_bn实现了这个。 - Observer 冻结(freeze):训练后期要把 Observer 的 scale 固定住(
apply_observer_freeze),否则 scale 抖动会影响收敛。 - 学习率调度:QAT 是微调,学习率要小,且配 warmup。
- 渐进位宽:
ProgressiveQuantSchedule让模型先适应 8bit 的"温和误差",再逐步加压到 4bit,比直接 4bit 训练更稳。
5.6 Transformer 的 QAT:选择性策略
ch5_transformer_qat.py 把 QAT 带到了 BERT 上,并且做了一件和 5.4 呼应的事:选择性 QAT。它对 BERT 的不同子层(q_proj、k_proj、v_proj、o_proj、FFN 等)分别做敏感性分析:
# 来自 ch5/ch5_transformer_qat.py(结构)
def get_sublayer_paths():
# 返回 BERT 各子层的路径,如 "encoder.layer.0.attention.self.query"
...
def apply_selective_qat(model, strategy="all", bits=8):
# strategy: "all" 全部量化 / 只量化敏感层 / 只量化注意力投影
...
def run_sensitivity(save_plots=False, device="cpu"):
# 逐层量化,测每层单独量化对 SST-2 精度的影响
...实验结论:
- 不是所有层都值得 QAT。某些层(例如某些 attention 投影)量化后精度掉得厉害,QAT 收益大;另一些层本身很鲁棒,QAT 纯属浪费。
- 注意力投影 vs FFN 的敏感性不同:FFN 通常是参数量大户,但注意力投影对量化更敏感。选择性策略可以"把预算花在刀刃上"。
- Transformer 的 QAT 因为结构规整(一堆同样的 block),可以写一个通用的替换逻辑,按路径精确控制哪些层被量化。
5.7 本章小结
QAT 的决策框架:
graph LR
A[PTQ 诊断] -->|误差集中在少数层| B[选择性 QAT]
A -->|整体误差都大| C[全模型 QAT]
B --> D[BN 折叠 + 伪量化 + STE]
C --> D
D --> E[渐进位宽调度]
E --> F[冻结 Observer 收敛]
F --> G[评估验证]- STE 是灵魂:前向模拟量化,反向梯度直通,模型学会"带伤作战"。
- 配方七件套:BN 折叠、伪量化层、Observer、per-channel 缩放、渐进位宽、Observer 冻结、小学习率微调。
- 先诊断再决策:用 PTQ 失败诊断定位问题层,能选择性 QAT 就不全模型 QAT。
- per-channel 是免费午餐:4bit per-channel QAT 可以追平 8bit per-tensor。
下一章跳出"自己手写量化器"的视角,看看工业界的三条现成量化通路:PyTorch TorchAO、ONNX Runtime、TensorFlow Lite。
作者: itech001 来源: 公众号:AI人工智能时代(the-ai-era) 网站: https://www.theaiera.top/ 关注每日最新AI新闻和技术博客,主页有更多的文章的AI 技术参考:https://www.theaiera.top
本文首发于 AI人工智能时代,转载请注明出处。