如何确定Bi-LSTM模型的时间步长以计算其计算复杂度?
先直接解决你最关心的时间步长问题:
一、找到你的模型时间步长
从你的代码和模型摘要里能直接得到答案:
- 代码里
Embedding层设置了input_length=max_len,这个max_len就是输入序列的长度,也就是时间步长。 - 看模型摘要里
embedding_1的输出形状:(None, 768, 768),其中第二个维度768就是时间步长T。
原因很简单:Embedding层的输出形状遵循(batch_size, input_length, embedding_dim)的格式,这里input_length就是每个输入样本的序列长度,也就是LSTM要处理的时间步数。所以你的模型时间步长T=768。
二、计算模型的计算复杂度
接下来拆解复杂度计算,重点看你提到的Bi-LSTM部分,这是模型的计算核心:
1. Bi-LSTM的每时间步复杂度
Bi-LSTM由正向、反向两个独立的LSTM组成,所以要先算单个LSTM的复杂度再乘以2。
对于一个隐藏维度为H的LSTM,每个时间步的计算涉及4个门(输入门、遗忘门、输出门、候选细胞状态),每个门都需要做输入到隐藏层和隐藏层到隐藏层的线性变换。单个LSTM每时间步的计算复杂度为:
O(H * (E + H))
其中:
- H是LSTM隐藏单元数(你的模型里是16)
- E是输入维度(这里是Embedding层的输出维度768)
因为是双向LSTM,所以每时间步的复杂度要乘以2,即:
O(2 * H * (E + H)) = O(2 * 16 * (768 + 16)) = O(25088)
复杂度分析里通常忽略常数项,所以可以简化为O(H*(E+H))。
2. 整个序列的总复杂度(SGD下)
如果使用完整的通过时间反向传播(BPTT)(你的模型默认是这种情况),整个序列的计算复杂度是时间步长乘以每时间步复杂度:
O(T * H * (E + H)) = O(768 * 16 * 784) = O(19,267,584)
3. 其他层的复杂度补充
其他层(Embedding、BatchNorm、Dense等)的计算量相比Bi-LSTM要小很多:
- Embedding层:如果是可训练的,前向是查表操作(O(T)),反向是每个词向量的梯度更新(O(T*E));如果是冻结的,反向不需要计算。
- BatchNorm/Activation/Dropout:都是逐元素操作,复杂度为输出的元素总数,比如第一个BatchNorm是O(768*768),远小于Bi-LSTM的计算量。
- Dense层:复杂度是O(32*2),可以忽略不计。
所以整个模型的计算复杂度主导项是Bi-LSTM部分,总复杂度近似为O(T*H*(E+H))。
补充你提到的O(W)的理解
你说的“每时间步的学习计算复杂度为O(W)”,这里的W通常指模型的可训练参数总数。对于SGD优化,每一轮迭代需要更新所有参数,所以总复杂度是O(W);但如果是按时间步的BPTT,每处理一个时间步需要更新LSTM的相关参数,而Embedding层的参数更新是基于整个序列的,所以更准确的是,完整BPTT下每处理一个序列的复杂度是O(T*W_lstm + W_emb),其中W_lstm是Bi-LSTM的参数数,W_emb是Embedding层的参数数。
你的模型里:
- W_lstm=100480(对应模型摘要里Bidirectional层的Param#)
- W_emb=37147392(Embedding层的Param#)
所以每序列的总参数更新复杂度是O(768100480 + 37147392) ≈ O(77,168,640 + 37,147,392) = O(114,316,032),量级还是O(TW)(W是总参数数)。
内容的提问来源于stack exchange,提问作者PeakyBlinder

