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

使用R neuralnet()训练的神经网络拟合佳但测试预测极差的排查求助

Hey Patrick, let's dig into this neuralnet issue you're facing—this is a super common problem with a few key fixes we can work through step by step.

Core Problem Breakdown

First off, what you're seeing is a mix of severe overfitting and potentially training/test data distribution mismatch:

  • Full-data training works great because the model is memorizing every noise point and idiosyncrasy in the dataset, not learning generalizable patterns.
  • Poor test set performance after splitting means either your training/test sets don't represent the same underlying population, or the model is overfitting so hard it can't adapt to unseen data.
  • The weird early stopping result (lowest test MSE near 0 epochs) is a huge red flag—this suggests your model isn't learning anything useful for the test set, which almost always ties back to data prep or split issues.
Step-by-Step Troubleshooting

1. Verify Training/Test Set Distribution Consistency

This is the #1 culprit for this kind of issue. Run these checks to make sure your splits aren't skewed:

# For continuous features: Compare means, std devs, and quantiles
continuous_feats <- names(your_data)[sapply(your_data, is.numeric) & names(your_data) != "target"]
distribution_check <- sapply(continuous_feats, function(col) {
  data.frame(
    Train_Mean = mean(training_set[[col]]),
    Test_Mean = mean(test_set[[col]]),
    Train_SD = sd(training_set[[col]]),
    Test_SD = sd(test_set[[col]]),
    Train_P50 = median(training_set[[col]]),
    Test_P50 = median(test_set[[col]])
  )
})
print(t(distribution_check))

# For categorical features: Compare relative frequencies
categorical_feats <- names(your_data)[sapply(your_data, is.factor) | sapply(your_data, is.character)]
for (feat in categorical_feats) {
  cat("\n---", feat, "---\n")
  print(prop.table(table(training_set[[feat]])))
  print(prop.table(table(test_set[[feat]])))
}

If any feature has a >10% difference in mean/median or category frequencies between sets, your split is bad. Switch to stratified sampling (use caret::createDataPartition instead of random splits) especially if your target variable is imbalanced.

2. Fix Data Preprocessing Consistency

9 times out of 10, people mess up preprocessing by scaling training and test sets separately. You must fit scalers/transformers only on the training set, then apply them to the test set:

# Example with standardization (center + scale)
library(caret)
scaler <- preProcess(training_set, method = c("center", "scale"))
training_scaled <- predict(scaler, training_set)
test_scaled <- predict(scaler, test_set) # Critical: Use training-set scaler!

If you scaled each set independently, their feature distributions will be shifted, and your model will never generalize.

3. Reduce Model Complexity to Fight Overfitting

neuralnet's default settings might be too complex for your 3600-row dataset. Try these tweaks:

  • Cut hidden neurons: Start small (e.g., hidden = c(3) instead of the default) — 3600 observations don't support a huge network.
  • Add regularization: Use the weight_decay parameter (if your neuralnet version supports it) or increase the threshold to stop training earlier.
  • Simplify activation functions: Switch from tanh to logistic or even linear activation for simpler patterns.

4. Fix Your Early Stopping Logic

Your current early stopping setup might be using the test set directly (which is a big no-no) or tracking the wrong metric. Here's a proper implementation with a separate validation set:

# Split training set into training + validation (80/20 split)
train_val_split <- createDataPartition(training_scaled$target, p = 0.8, list = FALSE)
train_sub <- training_scaled[train_val_split, ]
val_sub <- training_scaled[-train_val_split, ]

best_mse <- Inf
best_model <- NULL
max_epochs <- 50
patience <- 3
patience_counter <- 0

for (epoch in 1:max_epochs) {
  # Train a model for one epoch (adjust threshold/stepmax as needed)
  model <- neuralnet(target ~ ., data = train_sub, hidden = c(3), threshold = 0.01, stepmax = 1e5)
  # Predict on validation set
  val_pred <- predict(model, val_sub)
  val_mse <- mean((val_pred - val_sub$target)^2)
  
  # Track best model
  if (val_mse < best_mse) {
    best_mse <- val_mse
    best_model <- model
    patience_counter <- 0
  } else {
    patience_counter <- patience_counter + 1
    # Stop if validation MSE doesn't improve for `patience` epochs
    if (patience_counter >= patience) {
      cat("Early stopping at epoch", epoch, "\n")
      break
    }
  }
}

# Evaluate best model on test set
test_pred <- predict(best_model, test_scaled)
test_mse <- mean((test_pred - test_scaled$target)^2)
cat("Final Test MSE:", test_mse, "\n")

Never use the test set for early stopping—this leads to data leakage and overestimates of model performance.

5. Double-Check Target Variable Integrity

Quick sanity check: Make sure your test set's target variable isn't corrupted (e.g., missing values, wrong encoding):

# Check for missing values
sum(is.na(test_set$target))
# Check target distribution matches training set
prop.table(table(training_set$target))
prop.table(table(test_set$target))
Extra Tips
  • If your data is time-series, never use random splits—split by time (e.g., first 2880 rows for training, last 720 for test) to avoid leakage.
  • Use cross-validation (caret::train with method = "neuralnet") to get a more reliable estimate of generalization performance.
  • Plot training/validation MSE over epochs to visualize overfitting:
library(ggplot2)
# Assuming you tracked MSEs in lists during training
mse_tracking <- data.frame(
  Epoch = 1:length(train_mse_list),
  Training_MSE = unlist(train_mse_list),
  Validation_MSE = unlist(val_mse_list)
)
ggplot(mse_tracking, aes(x = Epoch)) +
  geom_line(aes(y = Training_MSE, color = "Training")) +
  geom_line(aes(y = Validation_MSE, color = "Validation")) +
  labs(title = "MSE vs. Training Epochs", color = "Dataset") +
  theme_minimal()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.27 07:33:14