基于VGG-16迁移学习模型:TensorBoard显示孤立冗余层问题问询
这个现象的核心在于TensorBoard扫描的是整个TensorFlow计算图,而Keras的Model对象只是计算图的一个子集视图,具体原因可以拆成这几点:
计算图残留节点
当你执行vgg_input_model = tf.keras.applications.VGG16(...)时,TensorFlow会在当前默认计算图中实例化VGG16的所有层——哪怕你设置了include_top=False,也只是跳过了顶层的全连接层,block4、block5这些卷积层依然会被创建。后续你用tf.keras.Model截取到block3_pool的子图,只是定义了一个从输入到该层输出的计算路径,但那些没被包含的层并没有被从计算图中删除,它们变成了没有被当前模型引用的孤立节点。Keras模型与计算图的差异
model.summary()、model.get_layer()和plot_model都是基于Keras模型的内部结构来工作的,只会展示模型明确包含的层;而TensorBoard是直接读取整个TensorFlow计算图的所有节点,不管这些节点是否被当前模型的计算路径用到,所以能看到那些孤立层。早期TF版本的特性
你使用的是TF 2.1的nightly版本,这个阶段的TF 2.x在计算图管理上还不够成熟,即时执行模式下容易残留这类未被引用的节点。后续的TF版本对计算图清理做了优化,这类现象会少很多。
这些孤立层不会影响模型的训练和推理,但如果觉得TensorBoard的图太乱,可以用下面的方法清理计算图:
方法一:用临时计算图隔离无关层
先在临时计算图中加载完整VGG16,截取子模型的配置后,在默认计算图中重建只包含需要层的模型:
# 在临时计算图中加载完整VGG16并获取子模型配置 with tf.Graph().as_default(): vgg_input_model = tf.keras.applications.VGG16(weights='imagenet', include_top=False, input_tensor=x) final_vgg_layer = vgg_input_model.get_layer("block3_pool") sub_model_config = tf.keras.Model(inputs=vgg_input_model.inputs, outputs=final_vgg_layer.output).get_config() # 在默认计算图中重建子模型并加载权重 input_model = tf.keras.Model.from_config(sub_model_config) for layer in input_model.layers: layer.set_weights(vgg_input_model.get_layer(layer.name).get_weights()) input_model.trainable = True
方法二:重置默认计算图
在加载VGG16前重置默认计算图,避免之前的计算图残留影响:
tf.compat.v1.reset_default_graph() # 之后再构建你的模型 vgg_input_model = tf.keras.applications.VGG16(weights='imagenet', include_top=False, input_tensor=x) final_vgg_layer = vgg_input_model.get_layer("block3_pool") input_model = tf.keras.Model(inputs=vgg_input_model.inputs, outputs=final_vgg_layer.output) input_model.trainable = True
注意:在TF 2.x的即时执行模式下,可能需要配合tf.compat.v1.disable_eager_execution()使用,或者在函数式API中更谨慎地管理计算图。
方法三:克隆子模型
用tf.keras.models.clone_model复制子模型,它会在当前计算图中重新创建所需的层,不会包含原模型的无关层:
vgg_input_model = tf.keras.applications.VGG16(weights='imagenet', include_top=False, input_tensor=x) final_vgg_layer = vgg_input_model.get_layer("block3_pool") temp_model = tf.keras.Model(inputs=vgg_input_model.inputs, outputs=final_vgg_layer.output) # 克隆模型并加载权重 input_model = tf.keras.models.clone_model(temp_model) input_model.set_weights(temp_model.get_weights()) input_model.trainable = True
内容的提问来源于stack exchange,提问作者FlashDD

