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

使用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,提问作者Никита Шубин

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.25 12:32:45