使用tf_mot进行QAT时,如何在__call__外为Keras模型层设置training=False?
解决Keras 3.x中量化感知训练(QAT)与骨干网络训练模式固定的兼容问题
针对Keras 3.x结合tensorflow-model-optimization(tf_mot)进行量化感知训练时,无法同时固定骨干网络训练模式(training=False)和兼容QAT的问题,提供两种可行方案:
方案一:构建模型时指定training=False,针对性量化可训练层
先按迁移学习要求构建包含training=False调用的模型,再仅对可训练的头部层进行量化标注,避免骨干网与QAT的兼容性冲突:
import keras from keras import applications, layers, models import tensorflow_model_optimization as tfmot # 1. 构建基础迁移学习模型 inp = layers.Input((None, None, 3)) backbone = applications.vgg16.VGG16(include_top=False, weights=None) # 调用骨干网时强制设置training=False,固定BN等层的推理模式 x = backbone(inp, training=False) # 冻结骨干网权重 backbone.trainable = False # 添加任务专属头部 x = layers.GlobalAveragePooling2D()(x) x = layers.Dense(10, activation='relu')(x) out = layers.Dense(1, activation='sigmoid')(x) base_model = models.Model(inp, out) # 2. 仅对可训练层进行量化标注 def quantize_trainable_layers(layer): if layer.trainable: return tfmot.quantization.keras.quantize_annotate_layer(layer) return layer # 克隆模型并应用标注 annotated_model = keras.models.clone_model( base_model, clone_function=quantize_trainable_layers ) # 转换为量化感知训练模型 quantized_model = tfmot.quantization.keras.quantize_model(annotated_model) # 编译并启动训练 quantized_model.compile( optimizer='adam', loss='binary_crossentropy', metrics=['accuracy'] ) # 训练时模型整体处于训练模式,但骨干网因调用时指定training=False,不会更新批次统计 quantized_model.fit(train_dataset, epochs=10, validation_data=val_dataset)
原理说明
- 骨干网通过
training=False调用,强制BN等层使用推理阶段的移动均值/方差,且trainable=False固定权重 - 仅对可训练的头部层进行量化标注,避免tf_mot处理骨干网时的兼容性问题
- 训练时头部层正常参与量化感知训练,骨干网保持冻结状态
方案二:替换骨干网中需固定训练模式的层
通过自定义层替换骨干网中的BatchNormalization等层,强制其始终使用推理模式,无需依赖training=False参数:
import keras from keras import applications, layers, models import tensorflow_model_optimization as tfmot # 自定义冻结版BN层,始终使用推理模式 class FrozenBatchNormalization(keras.layers.BatchNormalization): def call(self, inputs, training=False): # 强制忽略training参数,固定使用移动均值和方差 return super().call(inputs, training=False) # 递归替换骨干网中的所有BN层 def replace_bn_layers(layer): if isinstance(layer, keras.layers.BatchNormalization): return FrozenBatchNormalization.from_config(layer.get_config()) # 处理嵌套层结构 if hasattr(layer, 'layers'): for idx, sub_layer in enumerate(layer.layers): layer.layers[idx] = replace_bn_layers(sub_layer) return layer # 构建并修改骨干网 backbone = applications.vgg16.VGG16(include_top=False, weights=None) frozen_backbone = replace_bn_layers(backbone) # 冻结骨干网权重 frozen_backbone.trainable = False # 构建完整模型 inp = layers.Input((None, None, 3)) x = frozen_backbone(inp) x = layers.GlobalAveragePooling2D()(x) x = layers.Dense(10, activation='relu')(x) out = layers.Dense(1, activation='sigmoid')(x) model = models.Model(inp, out) # 直接转换为量化感知训练模型 quantized_model = tfmot.quantization.keras.quantize_model(model) # 编译训练 quantized_model.compile( optimizer='adam', loss='binary_crossentropy', metrics=['accuracy'] ) quantized_model.fit(train_dataset, epochs=10, validation_data=val_dataset)
原理说明
- 自定义BN层强制使用推理模式,无需在调用时指定
training=False - 骨干网权重冻结后,不会参与参数更新,完全符合迁移学习要求
- 整个模型可直接应用tf_mot的量化流程,无兼容性冲突
注意事项
- 确保使用与Keras 3.x兼容的tf_mot版本(建议>=0.7.5)
- 若骨干网包含Dropout等其他需固定训练模式的层,可参照BN层的方式自定义替换
- 训练过程中无需额外设置全局学习阶段,Keras 3.x已弃用
set_learning_phase,依赖层的training参数或自定义逻辑控制
内容的提问来源于stack exchange,提问作者Никита Шубин
相关产品推荐
相关产品推荐

