Keras中两种LSTM构建方式参数数量差异原因探究
嘿,这个问题我之前踩过坑,咱们把它拆解开讲清楚~
先说说弃用警告的原因
Keras 2之后就不推荐用input_dim和input_length这两个参数了,官方建议用input_shape或者batch_input_shape来定义输入形状,所以m2的写法会触发弃用警告,这是正常的,赶紧换成新写法就好啦。
参数数量差异的核心原因:输入形状的定义完全搞反了!
LSTM的参数数量计算是固定公式,但前提是要搞对输入的特征维度和序列长度,咱们分别看两个模型:
模型m2:LSTM(1, input_dim=5, input_length=1)
在旧的API里,input_dim指的是每个时间步的特征数,input_length指的是序列的长度。所以m2的输入形状被Keras解析为:(batch_size, input_length, input_dim) = (None, 1, 5)——也就是每条数据是1个时间步,每个时间步有5个特征。
LSTM的参数计算公式是:4 * (输入特征数 + 隐藏单元数 + 1),这里的+1是偏置项。代入m2的数值:4*(5 + 1 + 1) = 4*7 = 28,正好对应m2输出的28个参数。
模型m3:LSTM(1, batch_input_shape=(None,5,1))
batch_input_shape的格式是(batch_size, sequence_length, feature_size),也就是第一个维度是批量大小(设为None表示不固定),第二个是序列长度,第三个是每个时间步的特征数。所以m3的输入形状是:每条数据有5个时间步,每个时间步只有1个特征。
同样用公式计算参数:4*(1 + 1 + 1) = 4*3 = 12,这就是m3的12个参数的由来。
总结一下
两种写法的参数数量差异,本质是你把输入的序列长度和特征数搞反了:
- m2是「1个时间步,每个时间步5个特征」
- m3是「5个时间步,每个时间步1个特征」
如果想让m3和m2的参数数量一致,只需要把batch_input_shape改成(None,1,5)就行,这样输入形状和m2一致,参数数也会是28,而且没有弃用警告。
内容的提问来源于stack exchange,提问作者VISQL

