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

使用Keras Seq2Seq模型时遇维度不匹配错误求助

错误原因及修复方案

核心错误点

  • 编码器输入形状不匹配:你的编码器输入数据形状是(592,58),但代码中encoder_inputs定义为(None,592),完全搞反了序列长度和特征数的维度。LSTM要求输入为3D张量(batch_size, timesteps, features),需先将一维序列数据reshape为3D,再修正Input层定义。
  • 解码器Input层维度错误:代码中decoder_inputs = keras.Input(shape=(None,52,9))定义了4D输入(批量维度+3个维度),但LSTM仅接受3D输入。你的解码器输入数据是(592,52,9)(3D),Input层只需指定(timesteps, features)即可。
  • 解码器输出与标签维度不匹配:你设置了decoder_lstm(return_sequences=True),会返回形状为(batch_size, 52, 256)的序列输出,但你的标签是(592,51*9)的固定长度向量,两者维度无法对齐。

修复后的完整代码

import keras
from keras.layers import Input, LSTM, Dense, Flatten
from keras.models import Model

# 1. 处理编码器输入数据,将(592,58) reshape为3D张量
all_seqtf = all_seqtf.reshape((592, 58, 1))  # 每个样本58个时间步,特征数为1

# 编码器部分
encoder_inputs = keras.Input(shape=(58, 1))
encoder_lstm = keras.layers.LSTM(units=256, return_state=True)
encoder_outputs, state_h, state_c = encoder_lstm(encoder_inputs)
encoder_states = [state_h, state_c]

# 解码器部分(适配固定长度输出标签)
decoder_inputs = keras.Input(shape=(52, 9))
# 若输出固定长度向量,设置return_sequences=False
decoder_lstm = keras.layers.LSTM(units=256, return_sequences=False, return_state=True)
decoder_outputs, _, _ = decoder_lstm(decoder_inputs, initial_state=encoder_states)
decoder_dense = Dense(51*9, activation='relu')
decoder_outputs = decoder_dense(decoder_outputs)

# 构建并训练模型
model = Model([encoder_inputs, decoder_inputs], decoder_outputs)
model.compile(optimizer="rmsprop", loss="msle")
model.summary()
model.fit(x=[all_seqtf, all_anitf], y=all_ani2tf, batch_size=64, epochs=100, validation_split=0.2)

序列输出场景适配(若需逐时间步预测)

如果你的任务实际需要解码器返回序列输出,需将标签reshape为(592, 51, 9),并调整解码器结构:

# 处理标签为序列格式
all_ani2tf = all_ani2tf.reshape((592, 51, 9))

# 解码器部分修改
decoder_lstm = keras.layers.LSTM(units=256, return_sequences=True, return_state=True)
decoder_outputs, _, _ = decoder_lstm(decoder_inputs, initial_state=encoder_states)
decoder_dense = Dense(9, activation='relu')
decoder_outputs = decoder_dense(decoder_outputs)
# 展平序列以匹配标签的扁平化形状
decoder_outputs = Flatten()(decoder_outputs)

内容的提问来源于stack exchange,提问作者Mr.daq

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.22 19:43:03