You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用Keras函数式API融合两个UNet模型报错:Graph disconnected

Keras函数式API连接双UNet模型时出现图断开错误

问题现象

尝试用Keras函数式API连接两个UNet模型的decoder_stage4b_relu层输出,拼接后通过卷积层得到最终输出,但运行时触发以下错误:

ValueError: Graph disconnected: cannot obtain value for tensor KerasTensor(type_spec=TensorSpec(shape=(None, 512, 512, 3), dtype=tf.float32, name='data'), name='data', description="created by layer 'data'") at layer "bn_data". The following previous layers were accessed without issue: []

错误原因

创建model_a和model_b时指定了input_shape参数,这会让每个UNet模型自动创建独立的输入层。后续调用model_a(data_input)仅做了一次前向传播,但并未将data_input与模型的内部层关联到同一个计算图中。直接通过model_a.get_layer(...).output获取的张量,仍然绑定在原模型的输入层上,和新定义的data_input没有形成连通的计算流,导致Keras无法追踪输入到输出的完整路径。

修复方案

需要让两个UNet模型的层与自定义的data_input形成连通的计算图,核心是将模型作为层使用,让data_input作为输入流经模型的层,获取关联的中间层输出。

修复后代码(方法1:通过子模型获取中间层)

import tensorflow as tf
from tensorflow import keras
import segmentation_models as sm

# 自定义参数(需根据实际情况替换)
BACKBONE1 = 'resnet34'
BACKBONE2 = 'mobilenetv2'
n_classes = 1
activation = 'relu'

data_input = keras.Input(shape=(512,512,3))

# 处理模型A:创建子模型获取中间层输出
model_a = sm.Unet(BACKBONE1, encoder_weights='imagenet', classes=n_classes, activation=activation)
# 基于原模型的输入和目标中间层创建子模型
model_a_mid = keras.Model(inputs=model_a.input, outputs=model_a.get_layer('decoder_stage4b_relu').output)
# 用自定义输入data_input传入子模型,得到关联的中间层输出
model_a_mid_output = model_a_mid(data_input)

# 同理处理模型B
model_b = sm.Unet(BACKBONE2, encoder_weights='imagenet', classes=n_classes, activation=activation)
model_b_mid = keras.Model(inputs=model_b.input, outputs=model_b.get_layer('decoder_stage4b_relu').output)
model_b_mid_output = model_b_mid(data_input)

# 拼接中间层输出并构建最终输出
concat = tf.keras.layers.concatenate([model_a_mid_output, model_b_mid_output], axis=3)
data_output = keras.layers.Conv2D(3, 2, padding="same", activation="sigmoid")(concat)

# 构建完整的集成模型
ensemble_model = keras.Model(inputs=data_input, outputs=data_output, name="ensemble_model")
ensemble_model.summary()

修复后代码(方法2:遍历模型层获取中间层)

import tensorflow as tf
from tensorflow import keras
import segmentation_models as sm

# 自定义参数(需根据实际情况替换)
BACKBONE1 = 'resnet34'
BACKBONE2 = 'mobilenetv2'
n_classes = 1
activation = 'relu'

data_input = keras.Input(shape=(512,512,3))

# 处理模型A:遍历层获取中间输出
model_a = sm.Unet(BACKBONE1, encoder_weights='imagenet', classes=n_classes, activation=activation, input_shape=None)
x_a = data_input
for layer in model_a.layers:
    x_a = layer(x_a)
    if layer.name == 'decoder_stage4b_relu':
        model_a_mid_output = x_a
        break

# 同理处理模型B
model_b = sm.Unet(BACKBONE2, encoder_weights='imagenet', classes=n_classes, activation=activation, input_shape=None)
x_b = data_input
for layer in model_b.layers:
    x_b = layer(x_b)
    if layer.name == 'decoder_stage4b_relu':
        model_b_mid_output = x_b
        break

# 拼接中间层输出并构建最终输出
concat = tf.keras.layers.concatenate([model_a_mid_output, model_b_mid_output], axis=3)
data_output = keras.layers.Conv2D(3, 2, padding="same", activation="sigmoid")(concat)

# 构建完整的集成模型
ensemble_model = keras.Model(inputs=data_input, outputs=data_output, name="ensemble_model")
ensemble_model.summary()

修复原理

两种方法都确保了model_a_mid_output和model_b_mid_output是从data_input经过模型层计算得到的张量,Keras可以完整追踪从输入到输出的所有计算节点,计算图不再断开,因此能正常构建并编译模型。

内容的提问来源于stack exchange,提问作者dspeanut

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.25 23:48:21