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

如何调整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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.26 11:24:00