R语言kNN报错:调用外部函数时存在NA/NaN/Inf值(参数6)
解决kNN算法中的NA/NaN报错问题
报错原因分析
你的报错和警告来自两个核心问题:
- 非数值型特征无法被kNN处理:kNN基于距离计算,要求所有特征必须是数值型。你的
train.Sz和test.Sz包含大量字符型分类变量(比如Miesiąc、Terminal、Towar等),直接传入knn()函数时会自动尝试转换,转换失败生成NA,触发警告和最终错误。 - k值设置不合理:
k=202的数值可能远大于训练样本数量,导致算法无法找到足够的邻居。
分步解决方案
1. 对分类变量进行数值编码
需要将字符型/因子型分类变量转换为数值格式,推荐使用独热编码(适合无顺序的分类变量),用model.matrix实现:
# 将数据中的字符列转为因子,方便后续编码 data1[] <- lapply(data1, function(x) if(is.character(x)) factor(x) else x) # 生成独热编码,去掉截距项避免多重共线性 encoded_data <- model.matrix(~ . -1, data = data1) # 重新划分训练集和测试集 dat.d <- sample(1:nrow(encoded_data), size = nrow(encoded_data)*0.7, replace = FALSE) train.Sz <- encoded_data[dat.d, ] test.Sz <- encoded_data[-dat.d, ] # 提取标签(对应原数据的Status.szkody列) train.Sz_l <- data1$Status.szkody[dat.d] test.Sz_l <- data1$Status.szkody[-dat.d]
2. 检查并清理NA值
编码后可能残留少量NA,需要移除含NA的样本:
# 清理训练集NA train.Sz <- na.omit(train.Sz) train.Sz_l <- train.Sz_l[rownames(train.Sz)] # 清理测试集NA test.Sz <- na.omit(test.Sz) test.Sz_l <- test.Sz_l[rownames(test.Sz)]
3. 合理设置k值
k值必须小于训练样本数,通常取训练样本数的平方根并设为奇数(避免平局):
# 计算合理k值 k_val <- floor(sqrt(nrow(train.Sz))) k_val <- if(k_val %% 2 == 0) k_val + 1 else k_val # 运行kNN library(class) knn_result <- knn(train = train.Sz, test = test.Sz, cl = train.Sz_l, k = k_val)
完整修改后的代码
# 读取数据 data <- read.csv('~/Desktop/test1.csv', sep = ";") data1 <- subset(data,select=c(4,5,6,7,8,12,15,16)) # 处理分类变量:字符转因子+独热编码 data1[] <- lapply(data1, function(x) if(is.character(x)) factor(x) else x) encoded_data <- model.matrix(~ . -1, data = data1) # 划分数据集 dat.d <- sample(1:nrow(encoded_data), size = nrow(encoded_data)*0.7, replace = FALSE) train.Sz <- encoded_data[dat.d, ] test.Sz <- encoded_data[-dat.d, ] # 提取标签 train.Sz_l <- data1$Status.szkody[dat.d] test.Sz_l <- data1$Status.szkody[-dat.d] # 清理NA值 train.Sz <- na.omit(train.Sz) train.Sz_l <- train.Sz_l[rownames(train.Sz)] test.Sz <- na.omit(test.Sz) test.Sz_l <- test.Sz_l[rownames(test.Sz)] # 计算k值并运行kNN library(class) k_val <- floor(sqrt(nrow(train.Sz))) k_val <- if(k_val %% 2 == 0) k_val + 1 else k_val knn_result <- knn(train = train.Sz, test = test.Sz, cl = train.Sz_l, k = k_val) # 查看分类结果混淆矩阵 table(knn_result, test.Sz_l)
内容的提问来源于stack exchange,提问作者mzwk
相关产品推荐
相关产品推荐

