K近邻(KNN)分类预测绘图报错排查及解决方案咨询
KNN分类模型决策边界绘图报错修复方案
问题说明
尝试用K近邻(KNN)方法绘制分类预测的决策边界,但运行plot(kmeans_mod, xTrain)时出现如下报错:
Error in if (!(plotType %in% c("level", "scatter", "line"))) stop("plotType must be either level, scatter or line") : the condition has length > 1
期望得到包含特征空间分类边界+样本散点的可视化图(两类样本点分布,中间带有分类分界线)。
报错原因
caret包的train()返回的模型对象,其默认plot()方法仅用于绘制模型性能(比如不同k值的准确率变化),不支持直接传入训练数据生成决策边界图。你传入的xTrain被错误识别为plotType参数,导致参数类型不匹配,触发报错。
解决方法
要生成目标决策边界图,需手动生成特征空间的网格数据,用模型预测网格点类别后绘制分类区域,再叠加原始样本点。具体步骤如下:
1. 生成网格数据
基于训练数据的两个特征范围,生成密集网格点覆盖整个特征空间:
# 获取特征的取值范围 x1_range <- range(xTrain[,1]) x2_range <- range(xTrain[,2]) # 生成100*100的网格点 grid_x1 <- seq(from = x1_range[1], to = x1_range[2], length.out = 100) grid_x2 <- seq(from = x2_range[1], to = x2_range[2], length.out = 100) grid <- expand.grid(X1 = grid_x1, X2 = grid_x2)
2. 预测网格点类别
用训练好的KNN模型预测所有网格点的分类结果:
grid_pred <- predict(kmeans_mod, newdata = grid)
3. 绘制决策边界与样本点
使用ggplot2包完成可视化:
library(ggplot2) ggplot() + # 绘制分类区域背景 geom_tile(data = grid, aes(x = X1, y = X2, fill = grid_pred), alpha = 0.3) + # 叠加训练样本点 geom_point(data = xTrain, aes(x = xTrain[,1], y = xTrain[,2], color = as.factor(yTrain)), size = 2) + # 设置标签与主题 labs(x = colnames(xTrain)[1], y = colnames(xTrain)[2], fill = "预测类别", color = "真实类别") + theme_minimal()
完整可运行代码
set.seed(20220719) library(caret) library(ggplot2) # 加载数据集(替换为你的数据集读取代码) # classification <- read.csv("你的数据集路径") # 划分训练集与测试集 ii = createDataPartition(classification[,3], p = .75, list = F) xTrain = classification[ii, 1:2] yTrain = classification[ii, 3] xTest = classification[-ii, 1:2] yTest = classification[-ii, 3] # 设置训练控制参数 opts = trainControl(method = 'repeatedcv', number = 10, repeats = 5) # 训练KNN模型,寻找最优k knn_mod = train(x = xTrain, y = as.factor(yTrain), method ='knn', trControl = opts, tuneGrid = data.frame(k = seq(3, 10))) # 测试模型性能 yTestPred = predict(knn_mod, newdata = xTest) confusionMatrix(as.factor(yTestPred), as.factor(yTest)) # 生成决策边界图 x1_range <- range(xTrain[,1]) x2_range <- range(xTrain[,2]) grid_x1 <- seq(from = x1_range[1], to = x1_range[2], length.out = 100) grid_x2 <- seq(from = x2_range[1], to = x2_range[2], length.out = 100) grid <- expand.grid(X1 = grid_x1, X2 = grid_x2) grid_pred <- predict(knn_mod, newdata = grid) ggplot() + geom_tile(data = grid, aes(x = X1, y = X2, fill = grid_pred), alpha = 0.3) + geom_point(data = xTrain, aes(x = xTrain[,1], y = xTrain[,2], color = as.factor(yTrain)), size = 2) + labs(x = colnames(xTrain)[1], y = colnames(xTrain)[2], fill = "预测类别", color = "真实类别") + theme_minimal()
注意事项
- 若未安装
ggplot2,先运行install.packages("ggplot2")完成安装 - 需确保
classification数据框已正确加载,替换代码中的数据集读取部分 - 若要叠加测试集样本点,可添加
geom_point(data = xTest, aes(x = xTest[,1], y = xTest[,2], shape = as.factor(yTest)), size = 2)来区分训练/测试样本
内容的提问来源于stack exchange,提问作者Joe
相关产品推荐
相关产品推荐

