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

如何编写集成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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.18 01:37:17