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

Keras条件变分自编码器出现Graph disconnected错误如何解决

错误原因
  • 你的编码器、解码器计算链路的输入是直接用外部数据张量X、cond拼接得到的,和你最后定义的X_n、label_n两个Keras输入层没有任何关联
  • 声明模型时传入的输入[X_n, label_n]到输出h_p之间没有可连通的计算路径,Keras无法构建完整的计算图,因此抛出图断开错误
修复代码

把整个计算流的起点替换为你定义的两个输入层即可,修复后可运行的代码如下:

import tensorflow as tf
from tensorflow.keras.layers import *
from tensorflow.keras.models import Model

# 请保留你原来的sample_z函数实现
def sample_z(args):
    mu, log_sigma = args
    batch = tf.shape(mu)[0]
    dim = tf.shape(mu)[1]
    epsilon = tf.random.normal(shape=(batch, dim))
    return mu + tf.exp(log_sigma / 2) * epsilon

n_x = 20 # X的特征列数
n_y = 2 # 标签的列数
n_z = 5

# 输入层作为整个计算流的起点
X_n = Input(shape=(n_x,))
label_n = Input(shape=(n_y,))

# 编码器输入拼接
inputs = concatenate([X_n, label_n], axis=-1)
in_layer = Reshape((inputs.shape[1], 1))(inputs)

h_q =  Bidirectional(LSTM(16,activation='tanh', return_sequences=True))(in_layer)
h_q = BatchNormalization()(h_q)
h_q = Bidirectional(LSTM(32, activation='tanh', return_sequences=True))(h_q)
h_q = BatchNormalization()(h_q)
h_q = Bidirectional(LSTM(64, activation='tanh', return_sequences=True))(h_q)
h_q = BatchNormalization()(h_q)
h_q = Bidirectional(LSTM(128, activation='tanh'))(h_q)
h_q = BatchNormalization()(h_q)

mu = Dense(n_z, activation='linear')(h_q)
log_sigma = Dense(n_z, activation='linear')(h_q)

# 重参数化
z = Lambda(sample_z, output_shape = (n_z, ))([mu, log_sigma])

# 解码器输入拼接隐变量和条件标签
z_cond = concatenate([z, label_n], axis=-1)

# 解码器部分
h_p = RepeatVector(22)(z_cond) # 如果要还原20维输出可调整为20
h_p = BatchNormalization()(h_p)
h_p = Bidirectional(LSTM(128, activation='tanh', return_sequences=True))(h_p)
h_p = BatchNormalization()(h_p)
h_p = Bidirectional(LSTM(64, activation='tanh', return_sequences=True))(h_p)
h_p = BatchNormalization()(h_p)
h_p = Bidirectional(LSTM(32, activation='tanh', return_sequences=True))(h_p)
h_p = BatchNormalization()(h_p)
h_p = Bidirectional(LSTM(16, activation='tanh', return_sequences=True))(h_p)
h_p = BatchNormalization()(h_p)
h_p = Flatten()(TimeDistributed(Dense(1,activation='sigmoid'))(h_p))

# 此时输入输出链路完全连通,可正常构建模型
model = Model([X_n, label_n], outputs=h_p)
补充说明

如果需要同时输出mu和log_sigma用于计算VAE的损失,只需要修改Model的输出参数即可:
model = Model([X_n, label_n], outputs=[h_p, mu, log_sigma])

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.27 20:27:03