如何基于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
相关产品推荐
相关产品推荐

