如何编写集成predict()、table()和round()的R语言caret模型评估函数
问题描述
我正在用R语言的caret包train()函数生成预测模型,已完成模型训练,代码如下:
# Define the sampling methods and models sampling_methods <- c("none", "boot", "LGOCV", "cv", "repeatedcv") models <- c("LogitBoost", "gbm", "rf") # Function to fit the caret model WE_Proportion_fit_caret_model <- function(sampling_method, model) { # Set the seed for the model set.seed(sample(1:100, 1)) # Create the caret model control ctrl <- trainControl(method = sampling_method) # Fit the caret model WE_Proportion_fit <- train(Lethal ~ ., WE_SP_Train_Final, method = model, trControl = ctrl) # Return the fitted model return(WE_Proportion_fit) } # Iterate over sampling methods and models using purrr WE_Proportion_results <- crossing(sampling_method = sampling_methods, model = models) %>% mutate(model_fit = map2(sampling_method, model, WE_Proportion_fit_caret_model)) # Print the results print(WE_Proportion_results$model_fit)
目前需要对每个模型对象重复执行以下代码完成评估:
WE_boot_RF_pred <- predict(model_boot_RF, WE_Test_Year, type = "raw") WE_boot_RF_tab <- table(Actual = WE_Test_Year$Lethal, Predicted = WE_boot_RF_pred) WE_boot_RF_CEM <- round(compute.eval.metrics(WE_boot_RF_tab) * 100, 2) WE_boot_RF_CEM
请问有没有集成predict()、table()和round()的现成函数?或者该如何编写自定义函数来批量处理这些模型对象?
解决方案
没有现成的集成函数,但你可以自定义函数打包评估逻辑,再结合purrr实现批量处理。
1. 编写自定义评估函数
将预测、混淆矩阵构建、指标计算和格式化逻辑封装成一个函数:
evaluate_model <- function(fitted_model, test_data) { # 生成预测结果 pred <- predict(fitted_model, test_data, type = "raw") # 构建实际值vs预测值的混淆矩阵 conf_tab <- table(Actual = test_data$Lethal, Predicted = pred) # 计算评估指标并转为百分比格式(保留两位小数) eval_metrics <- round(compute.eval.metrics(conf_tab) * 100, 2) # 返回格式化后的指标 return(eval_metrics) }
2. 批量处理所有模型
利用dplyr和purrr的组合,直接在现有结果数据框中添加评估结果列:
library(dplyr) library(purrr) # 批量评估所有模型,新增评估指标列 WE_Proportion_results <- WE_Proportion_results %>% mutate(eval_metrics = map(model_fit, evaluate_model, test_data = WE_Test_Year))
3. 查看与整理结果
你可以直接查看原始结果,也可以将嵌套的指标展开为结构化数据框:
# 查看所有模型的评估指标 print(WE_Proportion_results$eval_metrics) # 展开为易读的数据框格式 library(tidyr) WE_Proportion_results %>% unnest(eval_metrics)
补充说明
- 确保
compute.eval.metrics()函数能正常处理混淆矩阵(如果是自定义函数或来自特定包,需提前加载对应依赖) - 可根据需求扩展自定义函数,比如添加参数控制预测类型、小数位数等
内容的提问来源于stack exchange,提问作者Nicholas Bradley
相关产品推荐
相关产品推荐

