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

Keras LSTM不平衡多分类模型全预测单一类别问题排查求助

Core Issues Causing Your LSTM to Only Predict Class 3

Let’s dive straight into the fundamental problems with your setup that are leading to this all-class-3 prediction behavior—no hyperparameter tuning needed to fix these:

1. Critical Input Feature Dimension Mismatch

Your training data X_train has a shape of (400, 20, 17) (meaning 17 features per timestep), but your model defines nb_features = 18 and uses input_shape=(timestamps, nb_features). This mismatch is a showstopper:

  • Keras will either silently pad your input with zeros to reach 18 features (corrupting your actual data) or throw an error (depending on your setup). Either way, the model can’t learn meaningful patterns from misaligned input.
  • Fix: Update nb_features = 17 to match your real data dimensions.

2. Misused TimeDistributed(Dense(1)) Wastes Temporal Information

You set return_sequences=True on your LSTM, which outputs a full sequence of (batch_size, 20, 50) (50 hidden units per timestep). But then you apply TimeDistributed(Dense(1)), which crunches each timestep’s 50-dimensional output down to 1 dimension. After flattening, you’re left with only 20 features per sample—this discards nearly all the temporal context the LSTM extracted.

  • Fix: Remove the TimeDistributed and Flatten layers by disabling return_sequences=True (use the LSTM’s final timestep output directly):
    model = Sequential()
    model.add(LSTM(units=hidden_neurons, input_shape=(timestamps, nb_features), dropout=0.2, recurrent_dropout=0.2))
    model.add(Dropout(0.2))
    model.add(Dense(units=nb_classes, activation='softmax'))
    
    If you want to keep sequential outputs, replace TimeDistributed(Dense(1)) with another LSTM layer instead.

3. Disabled Shuffling Creates Severe Class Bias

You set shuffle=False in model.fit(). If your training data is ordered by class (e.g., all class 0 first, then class 1, ..., finally class 3), the model will learn each class sequentially and overwrite its prior knowledge—ending up heavily biased toward the last class it trained on (class 3). Even if your data isn’t strictly ordered, not shuffling means the model sees non-representative batches, leading to poor generalization.

  • Fix: Set shuffle=True to randomize training data order every epoch.

4. Unnecessary Rounding in Metric Calculation

In your Metrics callback, you’re rounding softmax outputs before taking argmax:

val_predict = np.argmax((np.asarray(self.model.predict(self.validation_data[0]))).round(), axis=1)

Softmax outputs are probabilities—rounding can turn valid low-confidence predictions into all zeros, leading to incorrect argmax results. This isn’t the main issue here, but it’s a bug that will skew your metrics.

  • Fix: Remove the .round() call:
    val_predict = np.argmax(np.asarray(self.model.predict(self.validation_data[0])), axis=1)
    

Fix these four issues first, and you’ll see your model start to distinguish between classes without needing any hyperparameter tweaks.

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.11 09:29:43