如何在tidymodels工作流中查看K近邻的最近邻样本?
如何在tidymodels中追溯kknn的最近邻行
你使用tidymodels构建了kknn分类模型,现在想要追溯预测时被识别的3个最近邻对应的原数据行,偏好使用tidyverse风格的方法。
步骤1:提取拟合好的kknn模型对象
从已拟合的workflow中提取底层的kknn模型对象:
library(tidymodels) library(tidyverse) # 你的原有建模代码 knn_rec <- recipe(Species ~ ., data = iris) knn_lookup <- workflow() %>% add_model(nearest_neighbor(neighbors = 3) %>% set_engine("kknn") %>% set_mode("classification")) %>% add_recipe(knn_rec) %>% fit(data = iris) # 提取模型 model_fit <- extract_fit_parsnip(knn_lookup) kknn_obj <- model_fit$fit
步骤2:获取最近邻的行索引
调用kknn的预测方法,直接获取目标样本对应的邻居索引:
# 对目标样本进行预测并提取邻居信息 pred_result <- predict(kknn_obj, newdata = iris[1,1:4], k = 3) # 取出邻居的行索引 neighbor_indices <- pred_result$neighbors
步骤3:提取并整理最近邻数据
用tidyverse的函数从原数据中筛选出邻居行,并添加排名信息:
nearest_neighbors <- iris %>% slice(neighbor_indices) %>% mutate(neighbor_rank = row_number()) # 排名1为最近的邻居 # 查看结果 nearest_neighbors
进阶:整合进tidymodels流程
如果需要在workflow的预测输出中直接包含邻居信息,可以自定义函数实现:
# 自定义函数:返回包含预测结果和邻居数据的整洁数据框 get_pred_with_neighbors <- function(workflow_obj, new_data, k) { model_fit <- extract_fit_parsnip(workflow_obj) kknn_obj <- model_fit$fit pred_result <- predict(kknn_obj, newdata = new_data, k = k) bind_cols( # 基础预测结果 predict(workflow_obj, new_data = new_data), # 类别概率 predict(workflow_obj, new_data = new_data, type = "prob"), # 邻居数据(宽格式展示) iris %>% slice(pred_result$neighbors) %>% mutate(neighbor_rank = row_number()) %>% pivot_wider( names_from = neighbor_rank, values_from = everything(), names_prefix = "neighbor_" ) ) } # 使用函数获取整合后的结果 get_pred_with_neighbors(knn_lookup, iris[1,1:4], 3)
内容的提问来源于stack exchange,提问作者GreenManXY
相关产品推荐
相关产品推荐

