如何移除tf.keras预训练模型输入层,替换为自定义输入适配单通道灰度图
问题原因
你遇到的计算图断开报错,本质是Keras Functional模型的计算图在实例化时就已经固定,你尝试的effB0.layers[0] = effB0.layers[0](X)只是修改了模型层列表的元素值,并不会更新模型内部的节点连接关系。EfficientNetB0实例化时生成的原生3通道输入节点仍然和后续层绑定,和你定义的单通道输入没有形成连通路径,因此抛出错误。
可行解决方案
你需要手动重构EfficientNet的层连接路径,完全跳过原生输入节点,以你自定义的单通道转3通道的张量作为起始输入,过程中可以任意修改、裁剪中间层,完全符合你的自定义需求,参考代码如下:
import tensorflow as tf from tensorflow.keras.layers import Input, Conv2D from tensorflow.keras.models import Model def build_generator(input_shape=(256,256,1), cut_layer_name=None): # 加载EfficientNetB0结构与预训练权重(不需要预训练则weights设为None) effB0 = tf.keras.applications.EfficientNetB0( input_shape=(256,256,3), include_top=False, weights='imagenet' ) # 定义单通道输入 inputs = Input(shape=input_shape, name="model_input") initializer = tf.random_normal_initializer(0., 0.02) # 单通道转3通道 x = Conv2D( filters=3, kernel_size=1, strides=1, padding='same', kernel_initializer=initializer, activation='relu', name='first_conv' )(inputs) # 跳过原生输入层,从第二层开始依次传递张量 for layer in effB0.layers[1:]: # 如需编辑中间层,在此处根据层名判断做自定义修改即可,示例: # if layer.name == "target_layer_name": # x = 自定义操作(x) # else: x = layer(x) # 如需裁剪模型,指定裁剪层名后遇到该层即可停止遍历 if cut_layer_name is not None and layer.name == cut_layer_name: break # 构建连通的新模型 model = Model(inputs=inputs, outputs=x) return model generator = build_generator()
该方案不同于你不想使用的X = effB0(X)黑盒调用写法,你可以在遍历层的过程中自由插入自定义逻辑、修改层参数、裁剪模型,完全满足编辑中间层的需求,同时不会破坏预训练权重的加载与使用。
内容的提问来源于stack exchange,提问作者Madara
相关产品推荐
相关产品推荐

