TensorFlow加载模型时出现Graph disconnected错误求解决方案
问题分析与解决方案
你遇到的ValueError: Graph disconnected错误,本质是模型输入张量与ResNet50层的连接在构建特征提取器时出现断裂。原代码中直接将自带输入层的ResNet50加入Sequential模型,导致输入张量的引用关系混乱,Keras无法追踪完整的计算图。
以下是两种可行的解决方法:
方案一:显式添加Input层到Sequential模型
先创建不带输入参数的ResNet50,再在Sequential开头显式加入Input层,确保整个模型的输入路径清晰可追踪。
import tensorflow as tf num_classes = 10 input_shape = (32, 32, 3) # 加载ResNet50,不指定input_shape/input_tensor,避免自带输入层 base_model = tf.keras.applications.ResNet50(weights='imagenet', include_top=False) for layer in base_model.layers: layer.trainable = False # 显式添加Input层作为模型起点,统一输入张量 model = tf.keras.Sequential([ tf.keras.layers.Input(shape=input_shape), base_model, tf.keras.layers.GlobalAveragePooling2D(), tf.keras.layers.Dense(num_classes, activation='softmax') ]) # 用model.input作为特征提取器的输入,确保计算图连通 extractor = tf.keras.Model(inputs=model.input, outputs=[layer.output for layer in model.layers])
方案二:使用函数式API构建模型(更推荐)
函数式API能明确追踪输入与各层的连接关系,避免Sequential模型的隐式输入层冲突问题,逻辑更清晰。
import tensorflow as tf num_classes = 10 input_shape = (32, 32, 3) # 显式定义输入张量 inputs = tf.keras.Input(shape=input_shape) # 加载ResNet50并绑定输入张量,确保输入路径唯一 base_model = tf.keras.applications.ResNet50(weights='imagenet', include_top=False, input_tensor=inputs) base_model.trainable = False # 构建后续网络层,明确张量流向 x = base_model.output x = tf.keras.layers.GlobalAveragePooling2D()(x) outputs = tf.keras.layers.Dense(num_classes, activation='softmax')(x) # 创建完整模型 model = tf.keras.Model(inputs=inputs, outputs=outputs) # 构建特征提取器:可按需选择输出层 # 若仅需ResNet50、GAP、全连接层的输出 extractor = tf.keras.Model(inputs=model.input, outputs=[base_model.output, model.layers[-2].output, model.layers[-1].output]) # 若需所有层(含ResNet50内部层)的输出,可取消下方注释 # extractor = tf.keras.Model(inputs=model.input, outputs=[layer.output for layer in model.layers])
内容的提问来源于stack exchange,提问作者Vincent Ma
相关产品推荐
相关产品推荐

