Keras预训练ResNet50提取跳连接出现计算图断开问题如何解决
Keras预训练ResNet提取中间跳连接层解决方案
报错原因
你遇到的Graph disconnected错误核心原因是:你定义的convlayer_Q张量没有作为输入传入ResNet50模型,两条计算路径完全没有关联,从ResNet50取的中间层输出和前面的输入层无法形成完整的计算图。
正确实现方法
你需要将预处理后的张量传入ResNet50模型打通计算路径,推荐通过层名而非硬编码层序号提取跳连接,避免不同Keras版本层序变化带来的问题,参考代码如下:
import tensorflow as tf from tensorflow.keras.layers import Input, Conv2D from tensorflow.keras.applications import ResNet50 input_shape = (480,854,4) inputlayer_Q = Input(shape=input_shape, name="inputlayer_Q") # 4通道输入转3通道适配ResNet50输入要求 convlayer_Q = Conv2D(filters=3, kernel_size=(3,3), padding='same')(inputlayer_Q) # 加载预训练ResNet50 model_Q = ResNet50( include_top=False, weights='imagenet' ) # 打通计算路径:将预处理后的张量传入ResNet50 resnet_output = model_Q(convlayer_Q) # 按层名提取跳连接输出,ResNet50对应层名可打印model_Q.layers确认 res2_skip = model_Q.get_layer('conv2_block3_out').output res3_skip = model_Q.get_layer('conv3_block4_out').output res4_skip = model_Q.get_layer('conv4_block6_out').output # 构建完整编码器模型,包含原始输入和所有需要的输出 encoder_Q = tf.keras.Model( inputs=inputlayer_Q, outputs=[res2_skip, res3_skip, res4_skip, resnet_output] )
注意事项
- 孪生网络权值共享可以直接复用同一个
encoder_Q实例处理另一路输入即可 - 如需冻结预训练权重,设置
model_Q.trainable = False即可 - 若要修改内置模型结构,可通过上述提取中间层的方式重新组装自定义计算图,无需修改官方源码
内容的提问来源于stack exchange,提问作者Ahmed Hamdi
相关产品推荐
相关产品推荐

