You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.16 21:14:54