In the following tutorial Transfer learning and fine-tuning by TensorFlow it is explained that that when unfreezing a model that contains BatchNormalization (BN) layers, these should be kept in inference mode by passing training=False when calling the base model.
[…]
Important notes about
BatchNormalizationlayerMany image models contain
BatchNormalizationlayers. That layer is a special case on every imaginable count. Here are a few things to keep in mind.
BatchNormalizationcontains 2 non-trainable weights that get updated during training. These are the variables tracking the mean and variance of the inputs.- When you set
bn_layer.trainable = False, theBatchNormalizationlayer will run in inference mode, and will not update its mean & variance statistics. This is not the case for other layers in general, as weight trainability & inference/training modes are two orthogonal concepts. But the two are tied in the case of theBatchNormalizationlayer.- When you unfreeze a model that contains
BatchNormalizationlayers in order to do fine-tuning, you should keep theBatchNormalizationlayers in inference mode by passingtraining=Falsewhen calling the base model. Otherwise the updates applied to the non-trainable weights will suddenly destroy what the model has learned.[…]
In the examples they pass training=False when calling the base model, but later they set base_model.trainable=True, which for my understanding is the opposite of inference mode, because the BN layers will be set to trainable as well.
For my understanding there would have to be 0 trainable_weights and 4 non_trainable_weights for inference mode, which would be identical to when setting the bn_layer.trainable=False, which they stated would be the case for running the bn_layer in inference mode.
I checked the number of trainable_weights and number of non_trainable_weights and they are both 2.
I am confused by the tutorial, how can I really be sure BN layer are in inference mode when doing fine tuning on a model?
Does setting training=False on the model overwrite the behavior of bn_layer.trainable=True? So that even if the trainable_weights get listed with 2 these would not get updated during training (fine tuning)?
Update:
Here I found some further information: BatchNormalization layer - on keras.io.
[...]
About setting
layer.trainable = Falseon aBatchNormalizationlayer:The meaning of setting
layer.trainable = Falseis to freeze the layer, i.e. its internal state will not change during training: its trainable weights will not be updated duringfit()ortrain_on_batch(), and its state updates will not be run.Usually, this does not necessarily mean that the layer is run in inference mode (which is normally controlled by the
trainingargument that can be passed when calling a layer). "Frozen state" and "inference mode" are two separate concepts.However, in the case of the
BatchNormalizationlayer, settingtrainable = Falseon the layer means that the layer will be subsequently run in inference mode (meaning that it will use the moving mean and the moving variance to normalize the current batch, rather than using the mean and variance of the current batch).This behavior has been introduced in TensorFlow 2.0, in order to enable layer.trainable = False to produce the most commonly expected behavior in the convnet fine-tuning use case.
Note that: - Setting
trainableon an model containing other layers will recursively set thetrainablevalue of all inner layers. - If the value of thetrainableattribute is changed after callingcompile()on a model, the new value doesn't take effect for this model untilcompile()is called again.
Question:
- In case I want to fine tune the whole model, so I am going to unfreeze the
base_model.trainable = True, would I have to manually set the BN layers tobn_layer.trainable = Falsein order to keep them in inference mode? - What does happen when with the call of the
base_modelpassingtraining=Falseand additionally settingbase_model.trainable=True? Do layers likeBatchNormalizationandDropoutstay in inference mode?