如何调整Keras LSTM VAE代码以输出形状为(24)的softmax结果
问题分析
你当前输出形状不符合的核心原因:解码器LSTM层设置了return_sequences=True,会输出形状为(batch_size, timesteps, inter_dim)的三维张量,后续的Dense层仅对最后一维做变换,因此最终output的形状为(batch_size, 96, 24),和你期望的(batch_size, 24)不符。
同时还有两处逻辑问题需要修正:
- 你已经通过
add_loss方法添加了自定义VAE损失,无需再在compile中指定loss参数,两个损失会叠加导致训练异常 - 你最终输出用了softmax激活,原损失中的
binary_crossentropy适配二分类/多标签场景,要替换为分类交叉熵更适配多分类场景
调整方案
如果你不需要保留96步序列的重构输出,只需要最终输出24维softmax结果,可以按以下方式修改代码:
import tensorflow as tf from tensorflow.keras.layers import Input, LSTM, Dense, Lambda, RepeatVector, GlobalAveragePooling1D from tensorflow.keras import backend as K from tensorflow.keras.models import Model # encoder latent_dim = 24 inter_dim = 32 timesteps, features = 96, 24 def sampling(args): z_mean, z_log_sigma = args batch_size = tf.shape(z_mean)[0] epsilon = K.random_normal(shape=(batch_size, latent_dim), mean=0., stddev=1.) return z_mean + z_log_sigma * epsilon # 输入层 input_x = Input(shape= (timesteps, features)) # 编码器LSTM h = LSTM(inter_dim)(input_x) # 隐变量层 z_mean = Dense(latent_dim)(h) z_log_sigma = Dense(latent_dim)(h) z = Lambda(sampling)([z_mean, z_log_sigma]) # 解码器部分:保留重构逻辑用于损失计算 decoder1 = Dense(inter_dim, activation='relu')(z) decoder1 = RepeatVector(timesteps)(decoder1) decoder1 = LSTM(inter_dim, return_sequences=True)(decoder1) recon_out = Dense(features)(decoder1) # 重构输出,仅用于计算损失 # 最终输出部分:先通过池化消去时间维度,再接softmax pooled = GlobalAveragePooling1D()(decoder1) output = Dense(24, activation='softmax')(pooled) # 损失函数调整 def vae_loss2(input_x, recon_out, z_log_sigma, z_mean): # 重构损失 recon = K.sum(K.categorical_crossentropy(input_x, recon_out)) # KL散度 kl = 0.5 * K.sum(K.exp(z_log_sigma) + K.square(z_mean) - 1. - z_log_sigma) return recon + kl m = Model(input_x, output) m.add_loss(vae_loss2(input_x, recon_out, z_log_sigma, z_mean)) # 已自定义损失,compile不需要再传loss参数 m.compile(optimizer='adam', metrics=['accuracy'])
可选简化方案
如果你不需要保留序列重构的逻辑,只想用VAE提取特征后直接输出24维分类结果,可以直接去掉解码器的RepeatVector和LSTM层,简化代码:
# 简化版解码器(无序列重构) decoder_h = Dense(inter_dim, activation='relu')(z) output = Dense(24, activation='softmax')(decoder_h)
运行m.summary()即可看到最终输出层的形状为(None, 24),符合需求。
内容的提问来源于stack exchange,提问作者NoSleep
相关产品推荐
相关产品推荐

