升级至Keras 3.0.0后TimeDistributed层报错,如何修正该LSTM架构?
Keras 3中TimeDistributed层与LSTM配合的正确写法
问题背景
原有基于Keras 2.15的LSTM+TimeDistributed代码可正常运行,但升级到Keras 3.0.0后出现如下错误:
ValueError: Exception encountered when calling TimeDistributed.call(). Invalid dtype: <class 'NoneType'> Arguments received by TimeDistributed.call(): • inputs=tf.Tensor(shape=(1, None, 5), dtype=float32) • training=True • mask=None
原代码如下:
from numpy import array from keras.models import Sequential from keras.layers import Dense from keras.layers import TimeDistributed from keras.layers import LSTM # prepare sequence length = 5 seq = array([i/float(length) for i in range(length)]) X = seq.reshape(1, length, 1) y = seq.reshape(1, length, 1) # define LSTM configuration n_neurons = length n_batch = 1 n_epoch = 1000 # create LSTM model = Sequential() model.add(LSTM(n_neurons, input_shape=(length, 1), return_sequences=True)) model.add(TimeDistributed(Dense(1))) model.compile(loss='mean_squared_error', optimizer='adam') print(model.summary()) # train LSTM model.fit(X, y, epochs=n_epoch, batch_size=n_batch, verbose=2) # evaluate result = model.predict(X, batch_size=n_batch, verbose=0) for value in result[0,:,0]: print('%.1f' % value)
解决方法
Keras 3中,无需再显式使用TimeDistributed层,因为Dense等层会自动对序列输入的每个时间步应用运算。直接移除TimeDistributed包装,保留Dense(1)即可。
修改后的完整代码:
from numpy import array from keras.models import Sequential from keras.layers import Dense from keras.layers import LSTM # prepare sequence length = 5 seq = array([i/float(length) for i in range(length)]) X = seq.reshape(1, length, 1) y = seq.reshape(1, length, 1) # define LSTM configuration n_neurons = length n_batch = 1 n_epoch = 1000 # create LSTM model = Sequential() model.add(LSTM(n_neurons, input_shape=(length, 1), return_sequences=True)) # 直接使用Dense,无需TimeDistributed包装 model.add(Dense(1)) model.compile(loss='mean_squared_error', optimizer='adam') print(model.summary()) # train LSTM model.fit(X, y, epochs=n_epoch, batch_size=n_batch, verbose=2) # evaluate result = model.predict(X, batch_size=n_batch, verbose=0) for value in result[0,:,0]: print('%.1f' % value)
原理说明
Keras 3对序列输入的处理逻辑做了优化:当输入形状为(batch_size, timesteps, features)时,Dense层会默认在最后一个维度(features)上进行运算,自动遍历所有时间步,效果完全等价于Keras 2中TimeDistributed(Dense(...))的作用。因此原代码中显式的TimeDistributed包装反而会导致类型错误。
内容的提问来源于stack exchange,提问作者pieterbons
相关产品推荐
相关产品推荐

