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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.25 19:09:33