使用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
相关产品推荐
相关产品推荐

