咱们平时做迁移学习,总爱把预训练模型的前面一堆层冻住,只让后面的层学新东西。这个概念很简单,但真上手TensorFlow微调时,坑一个接一个,尤其是当你动了批量归一化(Batch Normalization,后面都叫BN)层的时候,一个不小心,你的模型在训练完以后,出来的预测结果跟瞎猜差不多。今天咱们不聊那些花哨的理论,就用最直白的话,把这个“冻结层顺序”的事掰扯清楚,顺便看看那个躲在角落里的BN统计量,到底是怎么被我们无意间搞坏的。

一、冻结层和BN的“恩怨情仇”

预训练模型,你可以把它理解成一个“见识很广的老员工”。它已经在海量数据上吃过苦、流过汗,学会了怎么识别人脸、怎么分辨猫狗。我们微调它时,通常希望把它已经学好的基础能力保留下来,只针对自己的小任务做点调整。于是,“冻结层”这个操作就出现了:把前面那些学好的层给锁起来,不让它们继续变,只让后面新加的层去适应我们的数据。

听起来很简单对吧?你可能会写一行代码:

layer.trainable = False

然后拍拍手,认为完事了。但真正训练起来,很多人发现效果不对,模型甚至越来越差。问题出在哪儿?很大概率出在BN层身上。

BN层在平时干活时,不只是简单地把数据归一化。它心里还揣着两个“小账本”,一个记录着滑动平均均值(moving mean),一个记录着滑动平均方差(moving variance)。训练期间,它每看一批数据,就会用这一批的均值和方差去更新那两个账本。更新完以后,账本里的数字就像是“历史经验”,将来做推理时就用它们来归一化。

如果把BN层冻结了,你以为账本就不动了,可现实偏偏不是这样。在TensorFlow的Keras环境里,一个层的trainable开关,主要管的是“这个层的权重可不可以让优化器去更新”,而不是“这个层在训练时会不会干些私活”。于是,你兴高采烈地把BN层设成不可训练,结果它还在偷偷更新滑动平均。更头痛的是,如果你冻结层的顺序不对,比如先冻结了某些层,然后又对整个模型执行了一次trainable = True,那之前设好的“锁”会瞬间被全部撬开,所有层又变回可训练状态,之前那一通操作全白费。

这就是为什么咱们必须先理清楚:冻结层顺序,以及“层可训练性”和“变量可训练性”到底是怎么一回事。

二、拆开“trainable”这枚鸡蛋:层可训练性和变量可训练性

2.1 层上的可训练开关

每个Keras层都有一个trainable属性,它是一个布尔值。当你把某个层的trainable设为False,理论上这个层的所有“可训练变量”就不再参与梯度更新了。也就是说,优化器在计算完梯度之后,会忽略这些变量,不会改动它们的值。

这里说的“可训练变量”,是指层里面那些希望从数据中学习的参数,比如全连接层的权重矩阵和偏置、卷积层的卷积核、BN层的缩放系数gamma和平移系数beta

2.2 变量上的可训练标签

但事情没这么简单。Keras里,每个变量自己也有一个trainable标签。举个例子,tf.Variable(0.0, trainable=False)创建出来的变量,就算它所在的层宣布“我可训练”,这个变量也不会被优化器碰到。反之,一个变量即使trainable=True,如果它所在的层被整体设成了trainable=False,那它最终也不会被更新。

所以,层上的trainable就像是小区大门的保安,变量自己的trainable就像是每家每户的房门锁。两把锁都得打开,变量才会被更新。

2.3 两者怎么联动

在Keras中,当你设置model.trainable = False时,Keras会递归地遍历模型里的所有层,把每一层的trainable都改成False。反过来,如果设置成True,也会把每一层的trainable改成True。这个“递归”特性,就是前面说的顺序坑的根源。如果你在某个层上单独设了trainable = False,之后又对整个模型来一句model.trainable = True,那么你之前设的那个层会被“顺走”了。

变量上的trainable则相对稳定,它不会因为层被整体设为可训练而自动变回True。但要注意,BN的滑动平均变量(moving_meanmoving_variance)从诞生起就是trainable=False。它们虽然不会被优化器更新,却在训练时会被层内部逻辑当作“副产物”来更新。这一点,恰恰是很多开发者容易忽略的。

三、顺序一错,一切白搭

3.1 一个肉眼可见的“事故现场”

咱们用一个简单到不能再简单的例子,演示一下顺序错误会带来什么后果。

# 技术栈:Python + TensorFlow/Keras
import tensorflow as tf

model = tf.keras.Sequential([
    tf.keras.layers.Dense(4, activation='relu', input_shape=(8,)),
    tf.keras.layers.BatchNormalization(),
    tf.keras.layers.Dense(2, activation='softmax')
])

# 目标:把第一个Dense层冻结住
model.layers[0].trainable = False
print('冻结后,第一层的trainable:', model.layers[0].trainable)  # False

# 但接下来不小心写了这么一行:
model.trainable = True
# 这行会遍历所有层,把每层的trainable都设置为True
print('再次设置后,第一层的trainable:', model.layers[0].trainable)  # True

你看,就多了一行model.trainable = True,之前辛辛苦苦冻结的第一层又解开了。

3.2 为什么顺序会覆盖?

model.trainable是一个会被“广播”的开关。当Keras检测到模型级别的这个属性被赋值时,它不会说“哎呀,我只改那些没被单独设置的层吧”,它很耿直,直接挨个把所有层的trainable全部覆盖成同一个值。所以,如果你先单独冻结一些层,再给模型整体设置trainable,那模型级别的设置会把你之前所有的单独设置全部推翻。

反过来,正确的顺序是先让模型整体“闭关”,再选择性给某些层“开门”。这样,开门的层会保留可训练状态,其余层依然是冻结的。

3.3 正确姿势

正确的写法应该是:

# 技术栈:Python + TensorFlow/Keras
import tensorflow as tf

model = tf.keras.Sequential([
    tf.keras.layers.Dense(4, activation='relu', input_shape=(8,)),
    tf.keras.layers.BatchNormalization(),
    tf.keras.layers.Dense(2, activation='softmax')
])

# 先整体冻结
model.trainable = False

# 再局部解冻,注意顺序:先整体再局部
model.layers[0].trainable = True

# 检查一下
for layer in model.layers:
    print(layer.name, layer.trainable)

这样,第一个Dense层保持可训练,其他层保持冻结。

四、BN统计量的“叛逆期”:trainable=False 不等于全锁死

4.1 BN层的内部小账本

BN层有两个特殊的变量:moving_meanmoving_variance。它们的trainable属性天生是False,所以优化器不会去更新它们。但是,在训练模式下,BN层会偷偷做这么几件事:

  • 计算当前批次的均值和方差;
  • 用当前批次的统计量对数据进行归一化;
  • 把当前批次的统计量按一定比例(由momentum控制)合并进moving_meanmoving_variance

这个合并动作,不受层级trainable=False的管控。也就是说,即使你明确说了“这个BN层不可训练”,它依然会在训练模式下更新那两个账本。这就像你告诉门卫“别让这个人进小区”,但这个人却从旁边的小门溜进去,把小区里的公告栏给改了。

4.2 现场验证:冻结了还在变

咱们做个实验,用事实说话。

# 技术栈:Python + TensorFlow/Keras
import numpy as np
import tensorflow as tf

model = tf.keras.Sequential([
    tf.keras.layers.Dense(4, input_shape=(8,)),
    tf.keras.layers.BatchNormalization(momentum=0.9),
    tf.keras.layers.Dense(1)
])

# 先build一下,好让变量存在
model.build((None, 8))

# 只冻结BN层
model.layers[1].trainable = False

print('训练前 moving_mean:', model.layers[1].moving_mean.numpy())

# 造一批均值大约为10的数据
x = np.random.normal(10.0, 0.1, size=(32, 8)).astype(np.float32)
y = np.random.normal(0.0, 1.0, size=(32, 1)).astype(np.float32)

model.compile(optimizer='adam', loss='mse')
model.fit(x, y, epochs=1, verbose=0)

print('训练后 moving_mean:', model.layers[1].moving_mean.numpy())

看见没?明明把BN层设成了trainable=False,但跑完一个fit之后,moving_mean0.0变成了大约1.0。这说明它内部的统计量还是被更新了。如果你用的批量很小,移动均值还会被带偏,最终推理结果自然不准。

4.3 真正的冻结:让层走推理模式

要让BN层彻底“躺平”,得让它在训练期间也按推理模式工作,也就是在调用它时强制传入training=Falsefit里我们控制不了这个参数,但我们可以自己写一个简单的训练循环。

# 技术栈:Python + TensorFlow/Keras
import numpy as np
import tensorflow as tf

model = tf.keras.Sequential([
    tf.keras.layers.Dense(4, input_shape=(8,)),
    tf.keras.layers.BatchNormalization(momentum=0.9),
    tf.keras.layers.Dense(1)
])

model.build((None, 8))

# 冻结BN层
model.layers[1].trainable = False

# 数据
x = np.random.normal(10.0, 0.1, size=(32, 8)).astype(np.float32)
y = np.random.normal(0.0, 1.0, size=(32, 1)).astype(np.float32)

optimizer = tf.keras.optimizers.Adam()
loss_fn = tf.keras.losses.MeanSquaredError()

@tf.function
def train_step(x_batch, y_batch):
    with tf.GradientTape() as tape:
        # 关键:强制推理模式,BN层就不会更新统计量了
        predictions = model(x_batch, training=False)
        loss = loss_fn(y_batch, predictions)
    grads = tape.gradient(loss, model.trainable_variables)
    optimizer.apply_gradients(zip(grads, model.trainable_variables))
    return loss

train_step(x, y)

print('训练后 moving_mean 仍为:', model.layers[1].moving_mean.numpy())

因为传了training=False,BN层直接使用已有的移动平均,不再更新账本。这才是真正的“彻底冻结”。

五、精细控制的综合示例

理解了上面的机制,我们就能写一个真正精细控制的微调脚本了。下面以ResNet50为例,演示如何正确冻结大部分层,同时让深层特征参与训练,并且保证BN统计量不会被意外覆盖。

# 技术栈:Python + TensorFlow/Keras
import tensorflow as tf

# 加载预训练的ResNet50,它内部含有很多BN层
base_model = tf.keras.applications.ResNet50(
    weights='imagenet',
    include_top=False,
    input_shape=(224, 224, 3)
)

# 第一步:先整体冻结(注意顺序,先整体后局部)
base_model.trainable = False

# 第二步:让最后10个层中的非BN层恢复可训练
for layer in base_model.layers[-10:]:
    if isinstance(layer, tf.keras.layers.BatchNormalization):
        # BN层保持完全冻结,不参与任何训练时的统计量更新
        layer.trainable = False
    else:
        layer.trainable = True

# 第三步:构建完整模型,显式让基础模型以推理模式运行
inputs = tf.keras.Input(shape=(224, 224, 3))
x = base_model(inputs, training=False)  # 关键:整个基础模型都走推理模式
x = tf.keras.layers.GlobalAveragePooling2D()(x)
x = tf.keras.layers.Dense(128, activation='relu')(x)
outputs = tf.keras.layers.Dense(10, activation='softmax')(x)
model = tf.keras.Model(inputs, outputs)

# 第四步:检查冻结状态
print("=== 基础模型最后10层的trainable ===")
for layer in base_model.layers[-10:]:
    print(f"{layer.name:30s} trainable={layer.trainable}")

print("\n=== 可训练变量数量 ===")
print(len(model.trainable_variables))

这里有一个很关键的点:我们在调用base_model(inputs, training=False)时,强制整个基础模型走推理模式。这样一来,即使那些被解冻的卷积层有权重更新,BN层也不会偷偷更新移动统计量。这种方式特别适合批量比较小、或者你不想让BN的统计量被某个小批量数据带偏的场景。

如果你想在“训练某些深层卷积层”的同时,让BN也能正常更新移动统计量,那就把training参数改成True,但那样的话,你得保证自己有足够大的批量,并且能接受BN统计量在微调过程中发生变化。大多数小规模数据集上的微调,我更建议用上面的“全推理模式”写法。

六、应用场景、优缺点和注意事项

6.1 应用场景

这种精细化控制其实很常见。比如你在做医学图像分类,自己的数据集只有几千张,想用ImageNet上预训练的模型做迁移学习。这时候你通常希望保持底层特征不变,只调整高层特征。再比如目标检测,经常把主干网络的前一大部分冻结,只微调最后几个残差块和检测头。还有风格迁移、图像分割等等,都会用到类似的策略。

当模型里有BN层时,上面提到的“冻结顺序”和“BN统计量”问题就会跳出来。所以,但凡涉及以下情况,都要格外小心:

  • 用小批量训练;
  • 冻结了很多层,但忘了冻结BN;
  • 先设了一些层不可训练,后来有设置了整个模型的可训练性;
  • 想保持预训练模型BN里的“历史经验”不变,却不知道它已经被悄悄改过。

6.2 优点与缺点

这种精细控制的最大优点是稳。冻结大部分层可以防止过拟合,因为可训练的参数量大大减少。同时,保住BN的移动统计量,能让模型的数值行为更接近预训练时的状态,推理结果更可靠。另一个优点是省资源,冻结层之后,梯度计算量会下降不少,显存开销也小。

缺点是灵活性变差。如果我们把BN层彻底锁死,模型就无法适应目标数据集中不同的数据分布。特别是一些分布差异很大的任务,比如把自然图像换成医疗图像,BN统计量其实需要重新调整,一味冻结反而限制了下限。此外,手动控制层顺序、写自定义训练循环,代码变多,也更容易出错。

6.3 注意事项

  • 一定要记住“先整体冻结,再局部解冻”的顺序,别反过来。
  • 想彻底冻结BN统计量,就必须让BN层在训练时也走training=False的路径,光设trainable=False不够。
  • 不要只相信layer.trainable的布尔值,要去检查模型的trainable_variables列表,看看哪些变量真的会参与更新。
  • 如果使用fit,你没法直接给BN层传training=False,可以考虑把模型包一层,或者在自定义train_step里面做控制。
  • 保存权重时,最好把moving_meanmoving_variance一起保存。Keras的model.save默认会保存,但如果只手动保存可训练变量,就会漏掉它们,导致加载后的模型行为完全不一样。

七、总结

微调预训练模型,本质上是在“保留旧本领”和“学习新技能”之间找平衡。冻结层是常用手段,但它的坑点在于“层可训练性”和“变量可训练性”不是一回事,而BN统计量又有自己的更新逻辑。我们得牢牢记住三件事:

第一,给模型整体设置trainable会覆盖所有层的trainable,所以先整体冻结、再局部解冻才是安全顺序。

第二,layer.trainable = False只阻止优化器更新该层的可训练变量,并不能阻止BN层在训练模式下更新moving_meanmoving_variance

第三,要实现精细控制,最好在调用模型时显式指定training参数,配合自定义训练循环,让一切尽在掌握。

搞明白这几点,微调时就能少踩很多坑,模型预测也就不会变成“瞎猜”了。