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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.16 09:25:03