训练可变长度时间序列LSTM时损失为NaN的问题排查与解决
LSTM训练损失持续为NaN的问题排查与解决方案
核心问题分析
输出与标签形状不匹配
你的所有LSTM层都设置了return_sequences=True,最终模型输出形状为(611, 1200, 1)(每个样本的每个时间步都生成一个预测值),但你将标签yreshape为(611*1200,),维度完全不匹配。再加上填充NaN的时间步对应的标签大概率是无效值,直接导致损失计算出现NaN。而你的任务是用最近5分钟数据预测下一个数值(单步预测),不需要每个时间步都输出结果,只需每个样本对应一个预测值即可。Masking层作用未覆盖输出端
Masking层仅会忽略输入中的NaN时间步,但不会自动过滤输出端对应填充步的预测结果。如果标签中包含这些无效步的NaN值,MSE计算时NaN参与运算会直接返回NaN。学习率过高
你设置的learning_rate=1e-2远高于Adam优化器默认的1e-3,过高的学习率容易引发梯度爆炸,导致参数更新后出现NaN。输入数据可能存在隐性NaN
若有效时间步的特征数据中本身存在NaN(非填充的缺失值),即使Masking层跳过了填充的NaN,这些真实缺失值也会在LSTM计算中传播,最终导致损失为NaN。
修正方案
1. 调整模型结构适配单步预测
将最后一层LSTM的return_sequences改为False,让模型每个样本输出一个预测值,对应标签形状(611,):
def lstmFit(y, X, n_hidden=1, n_neurons=30, learning_rate=1e-3): lstm = Sequential() lstm.add(Masking(mask_value=np.nan, input_shape=(None, X.shape[2]))) # 隐藏层(除最后一层)保持return_sequences=True for layer in range(n_hidden - 1): lstm.add(LSTM(n_neurons, activation="tanh", recurrent_activation = "sigmoid", return_sequences=True)) # 最后一层LSTM不返回序列,对应单步预测输出 lstm.add(LSTM(n_neurons, activation="tanh", recurrent_activation = "sigmoid", return_sequences=False)) lstm.add(Dense(1)) # 使用默认级别的学习率,避免梯度爆炸 optimizer = Adam(learning_rate=learning_rate) lstm.compile(loss="mse", optimizer=optimizer) early_stopping = EarlyStopping(monitor='loss', patience=10, verbose=1, restore_best_weights=True) # 标签保持原形状(611,)即可,无需reshape lstm.fit(X, y, epochs=100, callbacks=[early_stopping]) y_train_fit = lstm.predict(X) return lstm, y_train_fit
2. 清理输入与标签数据
- 检查输入数组
X的有效时间步(非填充NaN部分),若存在真实缺失值,用均值、中位数或线性插值法填充。 - 确认标签
y为每个样本对应一个预测值,形状为(611,),且无NaN值。
3. 验证Masking层有效性
可以通过查看中间层输出,确认Masking层是否正确跳过了填充的NaN时间步:
from tensorflow.keras.models import Model # 提取第一层LSTM的输出 intermediate_model = Model(inputs=lstm.input, outputs=lstm.layers[1].output) intermediate_output = intermediate_model.predict(X[:1]) # 观察输出是否对应输入NaN的位置有被屏蔽的状态(通常为零向量) print(intermediate_output)
4. 进一步稳定训练
如果调整后仍出现NaN,尝试:
- 进一步降低学习率至1e-4;
- 给LSTM层添加初始化器,避免参数初始值异常:
lstm.add(LSTM(n_neurons, activation="tanh", recurrent_activation = "sigmoid", return_sequences=False, kernel_initializer='glorot_uniform', recurrent_initializer='orthogonal'))
内容的提问来源于stack exchange,提问作者user26458368
相关产品推荐
相关产品推荐

