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

如何基于e1071包SVM的10折交叉验证结果生成混淆矩阵?

如何基于10折交叉验证的SVM模型生成混淆矩阵?

你提到用e1071包的svm()做10折交叉验证后想生成混淆矩阵——这里需要注意:默认带cross=10参数训练的SVM模型并不会自动保存每折的预测标签,所以没法直接从现有模型对象里提取数据来做矩阵。我们需要手动实现交叉验证流程,收集每个样本的真实标签和预测标签,再生成混淆矩阵。

下面是具体的步骤和代码:

1. 准备工作:加载必要的包

首先确保你安装并加载了e1071和caret包(caret能帮我们快速创建折叠,还提供更完善的混淆矩阵函数):

library(e1071)
library(caret)

2. 创建10折交叉验证的折叠索引

为为了保证实验可重复,先设置随机种子,然后生成10折的样本索引(按type的类别分布分层抽样,避免某折样本类别失衡):

set.seed(123) # 固定随机种子,确保结果可复现
folds <- createFolds(dtm$type, k = 10)

3. 循环执行交叉验证,收集真实与预测标签

遍历每一个折叠,分别训练模型、预测测试集,并把真实标签和预测标签存起来:

# 初始化空向量用于存储全局结果
true_labels <- c()
pred_labels <- c()

for (fold_idx in folds) {
  # 划分当前折的训练集和测试集
  train_set <- dtm[-fold_idx, ]
  test_set <- dtm[fold_idx, ]
  
  # 用你指定的参数训练SVM模型(gamma=0.5, cost=1)
  svm_fold_model <- svm(type ~ ., data = train_set, gamma = 0.5, cost = 1)
  
  # 对测试集做预测
  fold_pred <- predict(svm_fold_model, newdata = test_set)
  
  # 收集当前折的结果到全局向量
  true_labels <- c(true_labels, test_set$type)
  pred_labels <- c(pred_labels, fold_pred)
}

4. 生成混淆矩阵

现在我们有了所有样本的真实标签和预测标签,就可以生成混淆矩阵了:

用基础的table()函数生成简单混淆矩阵

适合快速查看类别匹配情况:

# 生成基础混淆矩阵
confusion_table <- table(真实标签 = true_labels, 预测标签 = pred_labels)
print(confusion_table)

用caret::confusionMatrix()生成带详细统计的混淆矩阵

这个函数会返回准确率、召回率、F1值、Kappa系数等更多实用统计指标:

# 生成带统计指标的混淆矩阵(注意第一个参数是预测值,第二个是真实值)
confusion_mat <- confusionMatrix(pred_labels, true_labels)
print(confusion_mat)

额外:查看单折的混淆矩阵

如果你想查看每一次交叉验证的单独混淆矩阵,可以在循环内部添加代码:

# 在循环内添加这部分代码,查看当前折的混淆矩阵
fold_confusion <- table(真实标签 = test_set$type, 预测标签 = fold_pred)
cat("当前折混淆矩阵:\n")
print(fold_confusion)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.15 04:01:54