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

使用Caret训练nnet时,如何获取训练与CV集预测值绘制误差图?

嘿,刚好用Caret做过类似的调参可视化,给你捋清楚怎么拿到训练集和CV保留集的预测值,还能画出你想要的误差对比图!

第一步:配置训练控制参数,开启预测值保存

首先要在trainControl里设置savePredictions = "all",这样Caret会把每一轮交叉验证中保留集的预测值都保存下来。同时指定交叉验证的方法(比如5折CV):

library(caret)

# 配置CV控制参数
ctrl <- trainControl(
  method = "cv",          # 交叉验证方法
  number = 5,             # 5折交叉验证
  savePredictions = "all",# 保存所有CV迭代的预测值
  verboseIter = TRUE      # 可选,实时查看调参进度
)
第二步:定义调参网格

把你要优化的weight_decay(对应nnet的decay参数)和隐藏层大小(对应nnet的size参数)的候选值列出来,用expand.grid生成所有参数组合:

# 定义参数网格
tune_grid <- expand.grid(
  size = c(5, 10, 15, 20),  # 隐藏层大小候选值
  decay = c(0.001, 0.01, 0.1)  # weight decay候选值
)
第三步:训练模型并获取CV保留集预测值

用train函数训练模型,传入刚才的控制参数和调参网格。训练完成后,模型对象里的$pred就是所有CV保留集的预测数据框:

# 训练nnet模型
nnet_model <- train(
  y ~ .,                     # 你的目标变量和特征的公式
  data = your_train_data,    # 替换成你的训练数据集
  method = "nnet",
  trControl = ctrl,
  tuneGrid = tune_grid,
  maxit = 1000,              # 增加迭代次数确保模型收敛
  trace = FALSE              # 关闭训练过程的冗余输出
)

# 查看CV保留集的预测值
head(nnet_model$pred)

这个pred数据框里包含了:每个样本作为保留集时的真实值(obs)、预测值(pred)、对应的参数组合(size/decay)、所属的CV折数(Resample),非常方便后续计算CV误差。

第四步:获取训练集的预测值

Caret默认不会保存每个参数组合下整个训练集的预测值,所以我们需要循环遍历每个参数组合,单独训练模型并预测训练集:

# 循环获取每个参数组合的训练集预测值
train_pred_list <- lapply(1:nrow(tune_grid), function(i) {
  # 取出当前参数组合
  current_params <- tune_grid[i, ]
  # 用整个训练集训练模型(不做CV)
  temp_model <- train(
    y ~ .,
    data = your_train_data,
    method = "nnet",
    trControl = trainControl(method = "none"),  # 关闭CV,直接训练全量数据
    tuneGrid = current_params,
    maxit = 1000,
    trace = FALSE
  )
  # 整理预测结果
  data.frame(
    size = current_params$size,
    decay = current_params$decay,
    obs = your_train_data$y,
    pred = predict(temp_model, your_train_data)
  )
})

# 合并成一个数据框
train_preds <- do.call(rbind, train_pred_list)
第五步:计算误差并绘制对比图

现在我们有了训练集和CV保留集的预测值,接下来可以计算误差(比如回归用RMSE,分类用Accuracy),然后画图对比:

library(dplyr)
library(ggplot2)

# 计算CV误差(这里以回归问题的RMSE为例)
cv_errors <- nnet_model$pred %>%
  group_by(size, decay) %>%
  summarise(
    cv_rmse = sqrt(mean((obs - pred)^2)),
    .groups = "drop"
  )

# 计算训练集误差
train_errors <- train_preds %>%
  group_by(size, decay) %>%
  summarise(
    train_rmse = sqrt(mean((obs - pred)^2)),
    .groups = "drop"
  )

# 合并误差数据并画图
error_compare <- merge(train_errors, cv_errors, by = c("size", "decay"))

ggplot(error_compare, aes(x = size, color = factor(decay))) +
  geom_line(aes(y = train_rmse), linewidth = 1) +
  geom_point(aes(y = train_rmse), size = 2) +
  geom_line(aes(y = cv_rmse), linewidth = 1, linetype = "dashed") +
  geom_point(aes(y = cv_rmse), size = 2) +
  labs(
    x = "隐藏层大小",
    y = "RMSE",
    color = "Weight Decay",
    title = "训练误差 vs 交叉验证误差"
  ) +
  theme_minimal()

这样你就能清晰看到每个参数组合下,训练误差和CV误差的变化趋势,帮你判断是否过拟合或者欠拟合啦!

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.19 07:53:03