AWQ 如何通过基于缩放的方法降低量化误差

1 量化误差从哪里来

权重量化可以表示为:

Q(w)=ΔRound(wΔ)Q(w)=\Delta\cdot Round\left(\frac{w}{\Delta}\right)

其中 Δ\Delta 是量化缩放因子。量化过程可以理解为:先用 Δ\Delta 将浮点权重映射到量化网格,再进行取整,最后乘回 Δ\Delta 得到量化后的权重。

原本前向传播中的计算:

y=wxy=wx

量化后变为:

y=Q(w)xy=Q(w)x

因此,权重量化产生的误差最终还会乘上对应的激活值 xx。论文将舍入误差记为:

RoundErr(wΔ)=Round(wΔ)wΔRoundErr\left(\frac{w}{\Delta}\right) = Round\left(\frac{w}{\Delta}\right) - \frac{w}{\Delta}

于是,量化对输出造成的误差大小可以近似写成:

ErrΔRoundErrxErr\approx\Delta\cdot RoundErr\cdot x

这里最重要的一点是:相同大小的权重量化误差,在激活值 xx 较大的通道上,会对最终输出产生更大的影响。

这也是 AWQ 为什么不仅关注权重本身,还要利用 activation 来寻找显著权重(salient weights)。

2 AWQ 的关键想法:放大重要权重

AWQ 对显著权重进行缩放:

wwsw\rightarrow ws

同时对对应的激活值进行反向缩放:

xxsx\rightarrow\frac{x}{s}

因为:

(ws)xs=wx(ws)\frac{x}{s}=wx

所以在没有量化的情况下,这种变换不会改变原模型的计算结果。

那么为什么缩放之后,量化误差反而会下降?

缩放后的误差可以写成:

Errnew=ΔRoundErrx1sErr_{new} = \Delta'RoundErr\cdot x\frac{1}{s}

而原始误差为:

Errold=ΔRoundErrxErr_{old} = \Delta RoundErr\cdot x

两者相除:

ErrnewErrold=ΔΔ1s\frac{Err_{new}}{Err_{old}} = \frac{\Delta'}{\Delta}\frac{1}{s}

如果缩放显著权重后,没有明显改变该 group 的最大权重,因此:

ΔΔ\Delta'\approx\Delta

那么就可以得到:

Errnew1sErrold\boxed{ Err_{new}\approx\frac{1}{s}Err_{old} }

也就是说,当 s>1s>1 时,显著权重对应的量化误差可以近似降低到原来的 1/s1/s

这也是 AWQ 基于缩放保护显著权重的核心原理。

3 为什么 ss 不能无限大

在论文实验中可以看到,当 ss 增大到一定程度之后,PPL 不再继续下降,甚至可能重新上升。

原因在于,一旦 salient weight 被放得太大,它可能改变所在 group 的最大权重:

max(w)\max(|w|)\uparrow

而量化缩放因子 Δ\Delta 又与 group 内最大权重有关,因此会导致:

Δ\Delta'\uparrow

此时整个 group 的量化间隔都会变粗,使其他 non-salient weights 的量化误差增大。

所以本质上存在一个平衡:

  • ss 太小 → salient weight 保护不够
  • ss 合适 → salient weight 误差下降,同时 Δ\Delta 基本不变
  • ss 太大 → Δ\Delta 被拉大,其他权重的量化误差增加

4 为什么 AWQ 最后要 Search to Scale

既然 ss 不能简单地无限增大,就需要找到一个合适的缩放系数。

因此 AWQ 利用 calibration data 中的 activation 统计信息,为不同输入通道搜索合适的 scaling factor,使量化后的层输出尽可能接近原始 FP16 层的输出。

简单来说,AWQ 要寻找的是这样一个平衡点:

保护 salient weights避免增大 non-salient weights 的误差\text{保护 salient weights} \quad\Longleftrightarrow\quad \text{避免增大 non-salient weights 的误差}

而 Search to Scale,就是用来寻找这个平衡点的过程。