4.6 火候再压一档:量化感知蒸馏 本节摘要:蒸馏把模型变小,量化把数字变短,两者合流就是量化感知蒸馏:学生在低精度下前向传播,量化引入的误差被教师的全精度批改不断纠正。本节讲双轨机制——低精度上课、全精度批改——为什么能互相兜底,给出可运行的实现骨架与完整案例,并交代它与 1.1 压缩工具箱里其他手艺的合流次序。 第四章收官。前面各节都在"学习方式"上做文章,本节把 1.1 压缩工具箱里的另一位主角——量化——请进同一间教室。这是部署导向的最后一道工序,也是压缩手段组合拳的标准收尾。 两个对手怎么互相兜底 量化把 float32 的权重与激活压到 int8,体积降约四倍、推理常获一倍以上加速——但量化是"带误差的近似",直接对训好的模型做训练后量化,精度可能掉一两个点。
本节摘要:蒸馏把模型变小,量化把数字变短,两者合流就是量化感知蒸馏:学生在低精度下前向传播,量化引入的误差被教师的全精度批改不断纠正。本节讲双轨机制——低精度上课、全精度批改——为什么能互相兜底,给出可运行的实现骨架与完整案例,并交代它与 1.1 压缩工具箱里其他手艺的合流次序。
第四章收官。前面各节都在"学习方式"上做文章,本节把 1.1 压缩工具箱里的另一位主角——量化——请进同一间教室。这是部署导向的最后一道工序,也是压缩手段组合拳的标准收尾。
量化把 float32 的权重与激活压到 int8,体积降约四倍、推理常获一倍以上加速——但量化是"带误差的近似",直接对训好的模型做训练后量化,精度可能掉一两个点。量化感知训练(QAT)让网络在训练中就体验量化误差,学会"在低精度下也做得对";而蒸馏给这份体验配了一位全精度师傅:学生带着量化误差学,教师不带误差地批改,误差被当成了要纠正的噪声而不是要学习的目标。
于是出现一个精妙的双轨结构:前向用低精度(模拟部署时看到的自己),损失用教师的全精度输出做靶子(模拟没被量化污染的自己应该达到的水平)。梯度照常反传到全精度的"影子权重"上——量化参数在训练中是伪装,真实参数始终全精度。
import torch import torch.nn as nn import torch.nn.functional as F class FakeQuantConv(nn.Conv2d): """带伪量化的卷积:前向模拟 int8,权重始终全精度以支撑梯度。""" def forward(self, x): w_scale = self.weight.abs().max() / 127.0 w_q = (self.weight / w_scale).round().clamp(-127, 127) * w_scale # 伪量化权重 x_scale = x.abs().max() / 127.0 x_q = (x / x_scale).round().clamp(-127, 127) * x_scale # 伪量化激活 return F.conv2d(x_q, w_q, self.bias, self.stride, self.padding) def qat_kd_step(student, teacher, x, y, T=4.0, alpha=0.8): teacher.eval() s_logits = student(x) # 学生内部卷积已是伪量化前向 with torch.no_grad(): t_logits = teacher(x) # 全精度教师批改 hard = F.cross_entropy(s_logits, y) soft = F.kl_div( F.log_softmax(s_logits / T, dim=-1), F.softmax(t_logits / T, dim=-1), reduction="batchmean") * T * T (alpha * soft + (1 - alpha) * hard).backward() # round 的梯度近似为直通估计器(straight-through): # 前向走离散取整,反向按恒等映射传梯度——这是伪量化能训练的前提。 qconv = FakeQuantConv(3, 16, 3, padding=1) out = qconv(torch.randn(2, 3, 8, 8)) print("伪量化前向输出形状:", out.shape, "可反传:", out.requires_grad) # 输出示例: # 伪量化前向输出形状: torch.Size([2, 16, 8, 8]) 可反传: True

背景。 4.1 案例里蒸馏出的端上学生(92.6%,float32 权重 38 MB)要过最后部署关:设备厂商要求 int8 权重、单帧 10 毫秒以内。
操作。 方案选量化感知蒸馏而非纯训练后量化。第一步做敏感层分析:逐层做训练后量化试测,找出掉点最多的三个层。第二步对全模型插入伪量化(本节的 FakeQuant 结构),敏感层的量化位宽保持 int8 但初始化 scale 用逐通道统计。第三步用原 float 教师带这位"低精度学生"复训 20 轮,蒸馏温度 4、alpha 0.85,学习率降为初训的十分之一。第四步导出真 int8 权重,端上实测。
结果。 纯训练后量化掉到 90.8%,单帧 8 毫秒;量化感知蒸馏版 92.3%,单帧 9 毫秒,权重 19 MB。折损从 1.8 个点压到 0.3 个点。
解读。 0.3 个点的折损基本抹平了量化的常规代价,这正是"师傅在场"的价值:学生每一步前向都带着量化误差,但批改来自无误差的教师,误差没有被"当成事实学进去",而是被持续纠正。学习率降到十分之一是必要的——伪量化前向的损失面比全精度粗糙,大步长会反复跨过取整边界造成震荡。敏感层分析则省了算力:不敏感的层甚至可以试试更低位宽,进一步压缩。
案例里"第一步"的敏感层分析,脚本骨架如下——逐层做"只量化这一层"的对照,掉点榜就是敏感层清单:
# 敏感层分析:一次只量化一层,看谁掉点最多 import copy import torch @torch.no_grad() def layer_sensitivity_scan(model, eval_fn, quantize_layer): """eval_fn 返回模型精度;quantize_layer 模拟单层 int8 化。""" base = eval_fn(model) report = [] for name, module in model.named_modules(): if not isinstance(module, torch.nn.Conv2d): continue probe = copy.deepcopy(model) target = dict(probe.named_modules())[name] quantize_layer(target) # 就地伪量化该层 acc = eval_fn(probe) report.append((name, base - acc)) # 掉点量 report.sort(key=lambda r: -r[1]) for name, drop in report[:5]: print(f"{name}: 掉 {drop:.3f}") return report # 输出示例(某学生网络,基线 92.6): # backbone.3.conv: 掉 0.912 # backbone.7.conv: 掉 0.437 # neck.1.conv: 掉 0.295 # head.cls_conv: 掉 0.204 # backbone.11.conv: 掉 0.118 # 榜首的掉点占全模型统一量化的六成——重点保护前两层, # 其余层走常规 int8 甚至 int4,这就是分析的直接产出。
变式。 极限压缩(4 比特以下)时量化误差方差过大,蒸馏温度要相应升高以保持软目标的信噪比;混合精度方案(敏感层 int8、其余 int4)配合逐层蒸馏损失是当前端侧的流行打法;语音与 NLP 的嵌入层通常单独豁免量化(词向量对精度敏感),只量化计算层。
⚠️ 常见坑:伪量化训练完忘了真正导出。训练期的"int8"只是模拟,若部署时仍用全精度影子权重,等于白忙一场且实测延迟不达标——导出流程要走通并复测端上延迟,才叫完成。
💡 关键直觉:量化管"体积与速度",蒸馏管"精度兜底"。二者不是竞争关系而是上下游:蒸馏先出又小又准的学生,量化感知再把它压过部署线,合流的次序通常是"先蒸后量"而不是反过来。