在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
相关产品推荐
相关产品推荐

