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

Keras LSTM验证损失从第1轮起持续上升问题求助

Hey there! Let's tackle your stock prediction LSTM problem head-on—this is such a common (and frustrating) hurdle with time series modeling, especially for something as noisy and unpredictable as stock prices. Let’s break down the key issues in your setup and actionable fixes to try:

1. Fix Critical Data Preprocessing & Variable Naming Issues

First, let’s clean up the basics that might be silently messing with your results:

  • Variable Name Mix-Up: You’ve got a bug here: x_test = scaler.transform(x_val). You’re overwriting the x_test variable with validation data, which makes your "test" set actually the validation set. Rename this to x_val_scaled to avoid confusion later.
  • NaN Handling: Using diff() introduces NaN values at the start of your dataset. Instead of hard-coding a slice from index 1000 onward, explicitly drop or fill NaNs first:
    df = df.diff(periods=1, axis=0).dropna()
    
    This ensures you’re not working with invalid data points accidentally.
  • Scaling Consistency: When using StandardScaler, make sure you only fit it on the training data (you’re doing this right!), but double-check that your target variable y_train/y_test is scaled if needed. Predicting raw diffed close+1 values while your features are scaled can create scale mismatches that confuse the model—consider scaling the target variable too (and inverse-transforming predictions later).

2. Fix Stateful LSTM Behavior

Stateful LSTMs are powerful but easy to misconfigure—this is likely a major culprit for your stagnant validation loss:

  • Reset States Between Epochs: Stateful LSTMs retain their hidden state across batches, but you need to reset this state at the end of each epoch to avoid carrying over irrelevant information from one training cycle to the next. Add this callback:
    from keras.callbacks import LambdaCallback
    reset_states = LambdaCallback(on_epoch_end=lambda batch, logs: model.reset_states())
    
    Then add it to your fit() call: callbacks=[reduce_lr, reset_states]
  • Verify Batch-Sample Alignment: Ensure your training/validation sample counts are exact multiples of your batch size. You’ve done this with slicing, but double-check that len(dataXtrain) % batch_size == 0 and len(dataXtest) % batch_size == 0 to avoid unexpected behavior.

3. Simplify Your Model to Fight Overfitting

Your current model is way too complex for most time series tasks (4 stacked LSTMs with up to 512 units!)—this is why it’s overfitting the training data but can’t generalize:

  • Start Small: Reduce the number of layers and units drastically. Try a simpler baseline first:
    model = Sequential()
    model.add(LSTM(64, return_sequences=False, stateful=True, 
                   batch_input_shape=(batch_size, timesteps, features),
                   recurrent_dropout=0.2))  # Use recurrent dropout for LSTMs
    model.add(Dense(1, activation='linear'))
    
    Recurrent dropout is designed specifically for recurrent layers (unlike regular Dropout, which only affects input layers) and helps prevent overfitting without killing training signal.
  • Add Early Stopping: Stop training as soon as validation loss stops improving to avoid wasting time on overfitting. Add this callback:
    from keras.callbacks import EarlyStopping
    early_stop = EarlyStopping(monitor='val_loss', patience=10, restore_best_weights=True)
    
    This will automatically revert to the model weights that gave the best validation loss.

4. Re-evaluate Your Task & Features

Stock prices are notoriously hard to predict—let’s make sure you’re setting up the task for success:

  • Try Classification First: Instead of predicting exact price changes (regression), start with a simpler binary classification task: predict whether the price will go up or down. This reduces noise and makes it easier to see if your model is learning any signal at all.
  • Prune Your Features: 200 features is way too many—most are probably redundant or irrelevant. Use a feature selection method like Random Forest feature importance to pick the top 20-30 features that correlate most with your target. Less noise = easier for the model to learn meaningful patterns.

5. Adjust Training Hyperparameters

Tweak these to help your model generalize better:

  • Lower Learning Rate: Nadam’s default learning rate (0.001) might be too high for your complex model. Try starting with 0.0001 and let the ReduceLROnPlateau take it from there.
  • Smaller Batch Size: A batch size of 256 is quite large for stateful LSTMs. Try reducing it to 64 or 32—smaller batches can help the model generalize better by exposing it to more varied batch combinations.

Modified Code Snippet (Key Fixes)

Here’s how your model setup and training loop might look with the above fixes:

# Fix variable naming
x_val_scaled = scaler.transform(x_val)

# ... (rest of data preparation remains similar)

# Simplified model with recurrent dropout
model = Sequential()
model.add(LSTM(64, stateful=True, batch_input_shape=(batch_size, timesteps, features),
               recurrent_dropout=0.2))
model.add(Dense(1, activation='linear'))
model.summary()

# Add critical callbacks
reduce_lr = ReduceLROnPlateau(monitor='val_loss', factor=0.9, patience=5, min_lr=1e-6, verbose=1)
early_stop = EarlyStopping(monitor='val_loss', patience=10, restore_best_weights=True)
reset_states = LambdaCallback(on_epoch_end=lambda batch, logs: model.reset_states())

model.compile(loss='mse', optimizer='nadam', metrics=[coeff_determination,'mse','mae','mape'])

history = model.fit(dataXtrain, dataYtrain,
                    validation_data=(dataXtest, dataYtest),
                    epochs=100,
                    batch_size=batch_size,
                    shuffle=False,
                    verbose=1,
                    callbacks=[reduce_lr, early_stop, reset_states])

Final Note

It’s important to remember that stock prices follow a random walk for the most part—even a perfectly tuned model might not beat a simple baseline (like predicting "no change" every time). But fixing the technical issues above will help you rule out model misconfiguration as the cause of your validation loss plateau.

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.15 06:47:13