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

在fchollet的seq2seq示例代码中,如何为编码器和解码器添加更多LSTM层?

在fchollet的Seq2Seq示例中添加多层LSTM的完整方案

我来帮你搞定这个多层LSTM的问题!fchollet的那个seq2seq入门示例确实经典,但加多层的时候很容易在状态传递和张量形状上踩坑,我一步步给你讲清楚怎么改,还有要避开的坑。

一、编码器部分的修改

原来的示例里编码器是单层LSTM,要改成多层,核心是中间层要返回完整序列供下一层处理,最后一层返回所有层的状态传给解码器。这里有两种实现方式:

方式1:逐层堆叠LSTM(直观易懂)

适合新手理解每一层的输入输出:

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

latent_dim = 256  # 保持和原示例一致的隐层维度
num_encoder_tokens = ...  # 你的输入词汇表大小
num_decoder_tokens = ...  # 你的输出词汇表大小

# 编码器输入
encoder_inputs = Input(shape=(None, num_encoder_tokens))

# 第一层LSTM:必须设置return_sequences=True,把所有时间步的输出传给下一层
encoder_lstm1 = LSTM(latent_dim, return_sequences=True, return_state=True)
encoder_outputs1, state_h1, state_c1 = encoder_lstm1(encoder_inputs)

# 第二层LSTM:同样需要return_sequences=True
encoder_lstm2 = LSTM(latent_dim, return_sequences=True, return_state=True)
encoder_outputs2, state_h2, state_c2 = encoder_lstm2(encoder_outputs1)

# 最后一层LSTM:不需要return_sequences,但要return_state返回最终状态
encoder_lstm3 = LSTM(latent_dim, return_state=True)
encoder_outputs, state_h3, state_c3 = encoder_lstm3(encoder_outputs2)

# 收集所有层的隐藏状态和细胞状态,传给解码器对应层
encoder_states = [state_h1, state_c1, state_h2, state_c2, state_h3, state_c3]

方式2:用StackedRNNCells(简洁高效)

Keras提供了StackedRNNCells来简化多层RNN的堆叠,自动处理状态传递:

from keras.layers import Input, LSTMCell, RNN, Dense
from keras.models import Model

# 定义3层LSTMCell
encoder_cells = [LSTMCell(latent_dim) for _ in range(3)]
encoder_rnn = RNN(encoder_cells, return_state=True)

encoder_inputs = Input(shape=(None, num_encoder_tokens))
encoder_outputs, *encoder_states = encoder_rnn(encoder_inputs)
# 此时encoder_states会自动按 [h1, c1, h2, c2, h3, c3] 的顺序排列

二、解码器训练阶段的修改

训练阶段的解码器需要接收编码器的所有层状态作为初始输入,同时处理目标序列输入:

对应逐层堆叠的编码器

# 解码器输入
decoder_inputs = Input(shape=(None, num_decoder_tokens))

# 第一层解码器LSTM:初始状态用编码器第一层的h和c
decoder_lstm1 = LSTM(latent_dim, return_sequences=True, return_state=True)
decoder_outputs1, _, _ = decoder_lstm1(decoder_inputs, initial_state=[encoder_states[0], encoder_states[1]])

# 第二层解码器LSTM:初始状态用编码器第二层的h和c
decoder_lstm2 = LSTM(latent_dim, return_sequences=True, return_state=True)
decoder_outputs2, _, _ = decoder_lstm2(decoder_outputs1, initial_state=[encoder_states[2], encoder_states[3]])

# 第三层解码器LSTM:初始状态用编码器第三层的h和c
decoder_lstm3 = LSTM(latent_dim, return_sequences=True, return_state=True)
decoder_outputs, _, _ = decoder_lstm3(decoder_outputs2, initial_state=[encoder_states[4], encoder_states[5]])

# 输出层:和原示例一致
decoder_dense = Dense(num_decoder_tokens, activation='softmax')
decoder_outputs = decoder_dense(decoder_outputs)

# 构建训练模型
model = Model([encoder_inputs, decoder_inputs], decoder_outputs)

对应StackedRNNCells的编码器

decoder_cells = [LSTMCell(latent_dim) for _ in range(3)]
decoder_rnn = RNN(decoder_cells, return_sequences=True, return_state=True)

decoder_inputs = Input(shape=(None, num_decoder_tokens))
decoder_outputs, *decoder_states = decoder_rnn(decoder_inputs, initial_state=encoder_states)

decoder_dense = Dense(num_decoder_tokens, activation='softmax')
decoder_outputs = decoder_dense(decoder_outputs)

model = Model([encoder_inputs, decoder_inputs], decoder_outputs)

三、解码器推理阶段的修改

推理阶段需要一步步生成序列,所以要保存每一层的状态并循环传递,这也是最容易出形状错误的地方:

对应逐层堆叠的编码器

# 推理用编码器:输出所有层的状态
encoder_model = Model(encoder_inputs, encoder_states)

# 解码器的状态输入:对应每一层的h和c,共6个输入
decoder_state_inputs = [
    Input(shape=(latent_dim,)), Input(shape=(latent_dim,)),
    Input(shape=(latent_dim,)), Input(shape=(latent_dim,)),
    Input(shape=(latent_dim,)), Input(shape=(latent_dim,))
]

# 解码器各层的状态传递
decoder_outputs1, state_h1, state_c1 = decoder_lstm1(
    decoder_inputs, initial_state=[decoder_state_inputs[0], decoder_state_inputs[1]]
)
decoder_outputs2, state_h2, state_c2 = decoder_lstm2(
    decoder_outputs1, initial_state=[decoder_state_inputs[2], decoder_state_inputs[3]]
)
decoder_outputs, state_h3, state_c3 = decoder_lstm3(
    decoder_outputs2, initial_state=[decoder_state_inputs[4], decoder_state_inputs[5]]
)

# 生成输出并收集新状态
decoder_outputs = decoder_dense(decoder_outputs)
decoder_states = [state_h1, state_c1, state_h2, state_c2, state_h3, state_c3]

# 构建推理解码器模型
decoder_model = Model(
    [decoder_inputs] + decoder_state_inputs,
    [decoder_outputs] + decoder_states
)

对应StackedRNNCells的编码器

encoder_model = Model(encoder_inputs, encoder_states)

# 解码器状态输入:数量和encoder_states一致
decoder_state_inputs = [Input(shape=(latent_dim,)) for _ in encoder_states]

decoder_outputs, *decoder_states = decoder_rnn(
    decoder_inputs, initial_state=decoder_state_inputs
)
decoder_outputs = decoder_dense(decoder_outputs)

decoder_model = Model(
    [decoder_inputs] + decoder_state_inputs,
    [decoder_outputs] + decoder_states
)

四、常见形状问题排查

你遇到的张量形状问题,大概率是这几个原因:

  • 中间层忘记设置return_sequences=True:如果中间层不返回完整序列,下一层LSTM会收到形状为(batch_size, latent_dim)的输入,而它需要的是(batch_size, timesteps, latent_dim),直接报错。
  • 编码器和解码器状态数量不匹配:多层LSTM每层都有h和c两个状态,3层就是6个状态,不能只传递最后一层的2个状态给解码器,否则解码器初始状态数量不够,形状不匹配。
  • 推理阶段状态输入没对应:推理时解码器的状态输入数量要和编码器输出的状态数量完全一致,每个输入的形状都是(latent_dim,),不能少也不能多。

内容的提问来源于stack exchange,提问作者S.Mandal

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.25 03:48:01