QuantizeWrapperV2功能解析:QAT下参数与数学运算疑问
我正在训练一个包含Conv2D、MaxPool2D、Activation、Dense等TensorFlow基础层的小型CNN模型,需要部署到存储空间有限、不支持浮点运算的嵌入式系统中。为此我采用QAT(感知量化训练)将权重最终量化为8位,使用的API是tfmot.quantization.keras.QuantizeWrapperV2。但我无法理解该API针对不同层的参数数量变化及对应的数学操作,希望了解这些未文档化的运算细节。以下是模型应用QAT前后的结构摘要,Param #列的差异令我困惑,恳请帮助。
未应用QAT的模型摘要
Layer (type) Output Shape Param # ================================================================= input_2 (InputLayer) [(None, 56, 56, 1)] 0 conv2d_3 (Conv2D) (None, 54, 54, 30) 300 activation_3 (Activation) (None, 54, 54, 30) 0 max_pooling2d_2 (MaxPooling (None, 27, 27, 30) 0 2D) conv2d_4 (Conv2D) (None, 25, 25, 16) 4336 activation_4 (Activation) (None, 25, 25, 16) 0 max_pooling2d_3 (MaxPooling (None, 12, 12, 16) 0 2D) conv2d_5 (Conv2D) (None, 10, 10, 16) 2320 activation_5 (Activation) (None, 10, 10, 16) 0 global_average_pooling2d_1 (None, 16) 0 (GlobalAveragePooling2D) dense (Dense) (None, 8) 136 activation_6 (Activation) (None, 8) 0 dense_1 (Dense) (None, 1) 9 activation_7 (Activation) (None, 1) 0 ================================================================= Total params: 7,101 Trainable params: 7,101 Non-trainable params: 0 _________________________________________________________________
应用QAT后的模型摘要
Layer (type) Output Shape Param # ================================================================= input_2 (InputLayer) [(None, 56, 56, 1)] 0 quantize_layer_1 (QuantizeL (None, 56, 56, 1) 3 ayer) quant_conv2d_3 (QuantizeWra (None, 54, 54, 30) 301 pperV2) quant_activation_3 (Quantiz (None, 54, 54, 30) 3 eWrapperV2) quant_max_pooling2d_2 (Quan (None, 27, 27, 30) 1 tizeWrapperV2) quant_conv2d_4 (QuantizeWra (None, 25, 25, 16) 4337 pperV2) quant_activation_4 (Quantiz (None, 25, 25, 16) 3 eWrapperV2) quant_max_pooling2d_3 (Quan (None, 12, 12, 16) 1 tizeWrapperV2) quant_conv2d_5 (QuantizeWra (None, 10, 10, 16) 2321 pperV2) quant_activation_5 (Quantiz (None, 10, 10, 16) 3 eWrapperV2) quant_global_average_poolin (None, 16) 3 g2d_1 (QuantizeWrapperV2) quant_dense (QuantizeWrappe (None, 8) 137 rV2) quant_activation_6 (Quantiz (None, 8) 3 eWrapperV2) quant_dense_1 (QuantizeWrap (None, 1) 14 perV2) quant_activation_7 (Quantiz (None, 1) 1 eWrapperV2) ================================================================= Total params: 7,131 Trainable params: 7,101 Non-trainable params: 30 _________________________________________________________________
新增的30个参数全是非训练参数,都是QAT模拟8位量化所需的缩放因子(scale)、**零点(zero_point)**以及跟踪激活动态范围的统计参数,下面逐层拆解差异原因:
1. 输入层后的QuantizeLayer(+3参数)
这是输入量化的前置层,3个参数对应:
- 输入张量的
scale:浮点值转8位量化值的缩放比例 - 输入张量的
zero_point:对应量化后0值的浮点偏移量 - 跟踪输入动态范围的移动平均统计参数(用来自动计算scale和zero_point)
2. Conv2D/Dense层(原参数+1)
比如conv2d_3从300变为301,新增的1个参数是权重的量化scale。QAT中权重的zero_point固定为0(因为权重分布通常对称,无需偏移),只需要一个scale参数来完成权重的伪量化:训练时先把浮点权重量化为8位,再反量化回浮点参与计算,模拟部署时的量化误差。
3. Activation层(除最后一层外,+3参数)
原Activation层无参数,包裹后新增的3个参数对应:
- 激活输出的
scale和zero_point:用于激活值的伪量化转换 - 跟踪激活动态范围的移动平均统计参数:训练中实时更新,用来自动校准scale和zero_point
4. MaxPool2D/最后一层Activation(+1参数)
这类层的量化逻辑相对简单:
- MaxPool2D是线性操作,量化时复用输入的zero_point,仅需1个参数记录输出的scale
- 最后一层Activation是模型输出,通常只需要scale参数(zero_point固定为0,简化部署时的推理)
核心量化数学操作
QAT的伪量化逻辑是训练时模拟8位量化的误差,核心公式如下:
权重伪量化
# 量化 量化权重 = clip(round(浮点权重 / scale), 0, 255) # 反量化(回到浮点参与训练) 反量化权重 = 量化权重 * scale
激活伪量化
# 量化 量化激活 = clip(round(浮点激活 / scale + zero_point), 0, 255) # 反量化(回到浮点参与训练) 反量化激活 = (量化激活 - zero_point) * scale
这些操作都是在训练过程中自动完成的,新增的非训练参数就是用来存储这些scale、zero_point和统计值,部署时这些参数会用来完成实际的浮点<->8位量化值转换。
内容的提问来源于stack exchange,提问作者Jhon Margalit

