Keras中VAE训练报错:tf__vae_loss缺失z_log_var和z_mean参数
解决VAE训练时的TypeError:缺失z_log_var和z_mean参数
错误原因
Keras自定义损失函数默认仅接收y_true(输入样本)和y_pred(模型输出)两个参数。你的vae_loss函数额外声明了z_log_var和z_mean参数,但这两个是编码器的中间层输出张量,无法通过常规的loss参数传递给训练流程,导致调用fit时触发参数缺失错误。
解决方案
将损失函数修改为仅接收y_true和y_pred两个参数,直接引用全局作用域中的z_mean和z_log_var张量(因为它们在模型构建阶段已定义)。同时修复sampling函数中固定batch_size的问题,避免维度不匹配。
关键修改点
- 调整损失函数签名:移除
z_log_var和z_mean参数,直接使用全局变量 - 修复采样函数维度:用
K.shape(z_mean)[0]替代固定的batch_size,适配动态batch大小 - 恢复重构损失的维度缩放:重新添加
* original_dim,保证损失尺度合理
修正后的完整代码
from keras.layers import Input, Dense, Lambda from keras.models import Model from keras import backend as K from keras import losses from keras.datasets import mnist import numpy as np batch_size = 100 original_dim = 28*28 latent_dim = 2 intermediate_dim = 256 nb_epoch = 5 epsilon_std = 1.0 x = Input(shape=(original_dim,), name="input") h = Dense(intermediate_dim, activation='relu', name="encoding")(x) z_mean = Dense(latent_dim, name="mean")(h) z_log_var = Dense(latent_dim, name="log-variance")(h) def sampling(args): z_mean, z_log_var = args # 用动态batch大小替代固定值,适配训练时的不同batch尺寸 batch_size = K.shape(z_mean)[0] epsilon = K.random_normal(shape=(batch_size, latent_dim), mean=0., stddev=epsilon_std) return z_mean + K.exp(z_log_var / 2) * epsilon z = Lambda(sampling, output_shape=(latent_dim,))([z_mean, z_log_var]) encoder = Model(x, [z_mean, z_log_var, z], name="encoder") input_decoder = Input(shape=(latent_dim,), name="decoder_input") decoder_h = Dense(intermediate_dim, activation='relu', name="decoder_h")(input_decoder) x_decoded = Dense(original_dim, activation='sigmoid', name="flat_decoded")(decoder_h) decoder = Model(input_decoder, x_decoded, name="decoder") output_combined = decoder(encoder(x)[2]) vae = Model(x, output_combined) vae.summary() # 修改损失函数,仅保留y_true和y_pred两个参数,直接引用全局的z_log_var和z_mean def vae_loss(y_true, y_pred): reconstruction_loss = losses.binary_crossentropy(y_true, y_pred) * original_dim kl_loss = -0.5 * K.sum(1 + z_log_var - K.square(z_mean) - K.exp(z_log_var), axis=-1) return K.mean(reconstruction_loss + kl_loss) vae.compile(optimizer='rmsprop', loss=vae_loss) (x_train, y_train), (x_test, y_test) = mnist.load_data() x_train = x_train.astype('float32') / 255. x_test = x_test.astype('float32') / 255. x_train = x_train.reshape((len(x_train), np.prod(x_train.shape[1:]))) x_test = x_test.reshape((len(x_test), np.prod(x_test.shape[1:]))) vae.fit(x_train, x_train, shuffle=True, epochs=nb_epoch, batch_size=batch_size, validation_data=(x_test, x_test), verbose=1)
内容的提问来源于stack exchange,提问作者jiraiya1729
相关产品推荐
相关产品推荐

