TensorFlow 2.4+子类模型中如何正确冻结/解冻预训练模型?
问题根因
- 核心问题出在你给
train_step、test_step、call手动加了@tf.function装饰器,导致可训练权重列表被静态缓存:
TensorFlow的@tf.function会在第一次执行函数时追踪计算图,将函数内部用到的self.trainable_weights捕获为当前的静态列表。第一次调用fit完成后,这些函数的计算图已经固化,就算你后续修改层的trainable属性、重新调用compile,已追踪的函数也不会自动更新用到的可训练权重列表,所以冻结/解冻操作看起来完全不生效。
而函数式API是静态图结构,Keras可以感知到层可训练状态的变化,重新构造计算图,因此没有这个问题。 - 次要问题:你将预训练卷积基
conv_base作为类初始化的默认参数,Python的函数默认参数在定义阶段就会被实例化,多次创建Test_Model实例时会复用同一个VGG16对象,容易出现权重串用的异常。
修复方案
- 移除
train_step、test_step上手动加的@tf.function装饰器:Keras的fit方法默认会自动将这两个方法包裹在tf.function中执行,且每次compile后会自动重新追踪计算图,能自动适配修改后的可训练权重列表。 call方法如果没有特殊的性能需求,也可以移除@tf.function,如果要保留的话,每次修改完可训练状态后调用model.call.get_concrete_function(tf.TensorSpec((None,70,70,3), tf.float32))触发重新追踪,或者直接调用tf.keras.backend.clear_session()清理全局缓存。- 调整类初始化逻辑,避免默认参数复用VGG16实例:
def __init__( self, num_unfrozen_head_layers, num_unfrozen_base_layers, num_classes, conv_base = None, ): super(Test_Model, self).__init__(name = "Test_Model") if conv_base is None: conv_base = tf.keras.applications.VGG16(include_top = False, weights = 'imagenet', input_shape = (70,70,3)) self.base = conv_base # 其余初始化逻辑保持不变
- 每次调用
set_trainable_layers修改完可训练状态后,必须重新调用compile,触发Keras重新关联可训练权重和优化器。
修改完成后再跑测试用例,就能看到冻结/解冻逻辑和函数式API表现一致。
内容的提问来源于stack exchange,提问作者Agnosie
相关产品推荐
相关产品推荐

