如何在R的caret包中利用最优超参数获取10折交叉验证各折预测结果
获取caret中SVM最优超参数对应的交叉验证预测结果
我完全懂你的需求——你用caret跑了带3次重复的10折交叉验证SVM模型,现在fit.svm$pred返回的是所有超参数组合的预测结果,但你只想要最优超参数对应的那些,甚至还想算出每个折在3次重复后的平均预测值对吧?下面一步步来搞定:
第一步:锁定最优超参数
你已经用fit.svm$bestTune拿到了最优的参数组合(对于svmRadial来说,一般是sigma和C这俩参数),我们先把它存下来方便后续筛选:
best_params <- fit.svm$bestTune
第二步:筛选最优参数对应的预测结果
fit.svm$pred是一个包含所有调优参数预测结果的数据框,我们只需要把和最优参数匹配的行挑出来就行。这里提供两种方式:
用基础R筛选:
best_preds <- fit.svm$pred[fit.svm$pred$sigma == best_params$sigma & fit.svm$pred$C == best_params$C, ]
用dplyr筛选(更简洁易读):
如果你装了dplyr包,用管道操作会更顺手:
library(dplyr) best_preds <- fit.svm$pred %>% filter(sigma == best_params$sigma, C == best_params$C)
现在best_preds里就只有最优超参数对应的所有交叉验证预测结果了,里面的列包括真实值obs、预测值pred、折数fold、重复次数resample等。
第三步:计算各折重复后的均值预测
如果你想要每个折在3次重复后的平均预测值,可以按fold和真实值obs分组计算(确保每个观测的预测均值对应正确):
fold_mean_preds <- best_preds %>% group_by(fold, obs) %>% summarise(mean_pred = mean(pred), .groups = "drop")
这样得到的fold_mean_preds里,每一行就是某个折里某个观测在3次重复中的平均预测值。
另外补充一下:如果你想直接看最优模型在交叉验证中的整体表现,也可以查看fit.svm$results,这里面是每个超参数组合的CV评估指标(比如你指定的RMSE),最优参数对应的那一行就是你要的最优模型的CV结果。
内容的提问来源于stack exchange,提问作者UseR10085
相关产品推荐
相关产品推荐

