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

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对象,容易出现权重串用的异常。
修复方案
  1. 移除train_step、test_step上手动加的@tf.function装饰器:Keras的fit方法默认会自动将这两个方法包裹在tf.function中执行,且每次compile后会自动重新追踪计算图,能自动适配修改后的可训练权重列表。
  2. call方法如果没有特殊的性能需求,也可以移除@tf.function,如果要保留的话,每次修改完可训练状态后调用model.call.get_concrete_function(tf.TensorSpec((None,70,70,3), tf.float32))触发重新追踪,或者直接调用tf.keras.backend.clear_session()清理全局缓存。
  3. 调整类初始化逻辑,避免默认参数复用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
            # 其余初始化逻辑保持不变
  1. 每次调用set_trainable_layers修改完可训练状态后,必须重新调用compile,触发Keras重新关联可训练权重和优化器。

修改完成后再跑测试用例,就能看到冻结/解冻逻辑和函数式API表现一致。

内容的提问来源于stack exchange,提问作者Agnosie

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.06 22:48:04