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

使用Keras函数式API时,如何从双向LSTM提取状态c和h?

解决Keras双向LSTM提取状态的问题

嗨,这个问题我之前也碰到过,本质是双向LSTM的返回状态结构和单向LSTM不一样导致的!

当你给单向LSTM设置return_state=True时,它会返回三个值:(output, state_h, state_c),但双向LSTM因为同时有前向和后向两个独立的LSTM层,所以返回的状态会多出来两个变量,完整的返回结构是:
(encoder_outputs, forward_h, forward_c, backward_h, backward_c)

你原来的代码只解包三个变量,自然会触发异常。下面给你两种常见的处理方式:

方式一:拼接前向和后向的状态(推荐,保留双向信息)

这种方式会把两个方向的隐藏状态、细胞状态分别拼接,得到包含双向完整信息的状态,再传给解码器使用:

from keras.layers import Concatenate

encoder_inputs = Input(shape=(None, num_encoder_tokens))
# 定义双向LSTM并开启return_state
encoder = Bidirectional(LSTM(latent_dim, return_state=True))
# 完整解包所有返回值
encoder_outputs, forward_h, forward_c, backward_h, backward_c = encoder(encoder_inputs)

# 拼接前向和后向的隐藏状态、细胞状态
state_h = Concatenate()([forward_h, backward_h])
state_c = Concatenate()([forward_c, backward_c])

注意:拼接后的状态维度是2*latent_dim,如果你的解码器还是用原来的latent_dim作为LSTM的units参数,可以通过Dense层把拼接后的状态映射回对应维度:

from keras.layers import Dense

state_h = Dense(latent_dim)(Concatenate()([forward_h, backward_h]))
state_c = Dense(latent_dim)(Concatenate()([forward_c, backward_c]))

方式二:只使用单个方向的状态(简化,丢失部分信息)

如果你不需要保留双向的全部信息,也可以只取前向或者后向的状态,和单向LSTM的逻辑保持一致:

encoder_inputs = Input(shape=(None, num_encoder_tokens))
encoder = Bidirectional(LSTM(latent_dim, return_state=True))
encoder_outputs, forward_h, forward_c, backward_h, backward_c = encoder(encoder_inputs)

# 只使用前向的状态(也可以换成backward_h和backward_c)
state_h = forward_h
state_c = forward_c

这样调整之后,你的seq2seq翻译模型就能正常运行啦!

内容的提问来源于stack exchange,提问作者Mauro Gentile

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.21 03:37:05