第 5 章 QAT伪量化STE

第 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, None

STE 的本质:前向传播老老实实模拟量化(该截断截断、该取整取整),反向传播假装量化不存在(梯度直接穿过)。这是 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 版本。"""
    ...

关键设计:

  1. Observer(观察器):每次前向更新权重/激活的 min/max,进而算出当前的 S/Z。Observer 的更新不参与梯度(detach),它只是"统计"。
  2. FakeQuantized 层:前向把数据过一遍伪量化,让模型在训练时就"看见"量化误差。
  3. 训练流程:先用 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 必须处理的技术点:

  1. BN 折叠(BatchNorm folding):推理时 BN 是线性变换,可以折进 Conv 的权重里。QAT 前必须先折叠,否则训练和推理的数值行为不一致。fold_bn_into_conv 和 fold_all_bn 实现了这个。
  2. Observer 冻结(freeze):训练后期要把 Observer 的 scale 固定住(apply_observer_freeze),否则 scale 抖动会影响收敛。
  3. 学习率调度:QAT 是微调,学习率要小,且配 warmup。
  4. 渐进位宽: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[评估验证]
  1. STE 是灵魂:前向模拟量化,反向梯度直通,模型学会"带伤作战"。
  2. 配方七件套:BN 折叠、伪量化层、Observer、per-channel 缩放、渐进位宽、Observer 冻结、小学习率微调。
  3. 先诊断再决策:用 PTQ 失败诊断定位问题层,能选择性 QAT 就不全模型 QAT。
  4. 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人工智能时代,转载请注明出处。