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

交叉验证tidyclust workflowset指标收集与K值可视化问题

问题与解决方案

问题描述

想用经过交叉验证的tidyclust workflowset绘制K值(聚类数)与平方和比率(Sum of Square Ratio)的关系图,但使用collect_metrics()函数后,找不到能和平均误差配对绘图的独立K值列。尝试从.config字段提取K值,但不确定该方法是否可靠。

尝试代码

tune_results <- wf_set %>%
  collect_metrics() %>%
  filter(.metric == "sse_ratio")

tune_results %>%
  ggplot(aes(x = as.numeric(stringr::str_sub(.config, -2, -1)), y = mean, color = wflow_id)) +
  geom_point() +
  geom_line() +
  theme_minimal() +
  ggtitle("Plot of WSS/TSS ratio by Cluster Number") +
  ylab("mean WSS/TSS ratio, over 10 folds") +
  xlab("Number of clusters") +
  scale_x_continuous(breaks = 1:10)

可复现示例

if (!requireNamespace("pacman", quietly = TRUE)) {
  message("Installing pacman...")
  install.packages("pacman")
}

# 安装包
pacman::p_load(tidyverse, tidymodels, tidyclust, janitor, ClusterR, knitr, moments, visdat, skimr, DescTools)

mtcars <- mtcars %>%
  mutate(
    `am` = factor(`am`, labels = c(`0` = "auto", `1` = "man")),
    `vs` = factor(`vs`, labels = c(`0` = "V-shaped", `1` = "straight")),
    `cyl` = factor(`cyl`),
    `gear` = factor(`gear`),
    `carb` = factor(`carb`)
  )

# 设置3折交叉验证
mtcars_cv <- vfold_cv(mtcars, v = 3)

# 设置随机种子保证可复现
set.seed(123)

# 模型定义
kmeans_spec <- k_means(num_clusters = tune())

# 预处理配方
rec1 <- recipe(~., data = mtcars) %>%
  step_dummy(all_nominal_predictors()) %>%
  step_zv(all_predictors()) %>%
  step_normalize(all_numeric_predictors())

rec2 <- recipe(~., data = mtcars) %>%
  step_novel(all_nominal()) %>%
  step_dummy(all_nominal()) %>%
  step_zv(all_predictors()) %>%
  step_normalize(all_predictors()) %>%
  step_pca(all_predictors(), num_comp = 2)

rec3 <- recipe(~ ., data = mtcars) %>%
  step_dummy(all_nominal_predictors()) %>%
  step_zv(all_predictors()) %>% 
  step_normalize(all_numeric_predictors()) %>% 
  step_center(all_numeric())

# 生成聚类数网格
clust_num_grid <- grid_regular(num_clusters(),
  levels = 10
)

# 创建工作流集合
wf_set <- workflow_set(
  preproc = list(rec1, rec2, rec3),
  models = list(kmeans_spec)
)

# 调优函数
tune_cluster_wf <- function(id) {
  tune_cluster(
    extract_workflow(wf_set, id),
    resamples = mtcars_cv,
    grid = clust_num_grid,
    metrics = cluster_metric_set(sse_within_total, sse_total, sse_ratio),
    control = tune::control_grid(save_pred = TRUE, extract = identity)
  )
}

# 运行调优
wf_set$result <- map(wf_set$wflow_id, tune_cluster_wf)

# 提取指标并绘图
tune_results <- wf_set %>%
  collect_metrics() %>%
  filter(.metric == "sse_ratio")

tune_results %>%
  ggplot(aes(x = as.numeric(stringr::str_sub(.config, -2, -1)), y = mean, color = wflow_id)) +
  geom_point() +
  geom_line() +
  theme_minimal() +
  ggtitle("Plot of WSS/TSS ratio by Cluster Number") +
  ylab("mean WSS/TSS ratio, over 10 folds") +
  xlab("Number of clusters") +
  scale_x_continuous(breaks = 1:10)

解决方案

方法1:使用标准输出的参数列(推荐)

collect_metrics()应该直接返回调优的参数列num_clusters,无需从.config提取。修改绘图代码如下:

tune_results <- wf_set %>%
  collect_metrics() %>%
  filter(.metric == "sse_ratio")

tune_results %>%
  ggplot(aes(x = num_clusters, y = mean, color = wflow_id)) +
  geom_point() +
  geom_line() +
  theme_minimal() +
  ggtitle("WSS/TSS比率随聚类数变化图") +
  ylab("3折交叉验证下的平均WSS/TSS比率") +
  xlab("聚类数") +
  scale_x_continuous(breaks = 1:10)

如果环境中未自动带出num_clusters列,可先展开result列再提取指标:

tune_results <- wf_set %>%
  unnest(result) %>%
  collect_metrics() %>%
  filter(.metric == "sse_ratio")

方法2:稳健提取.config中的K值

若必须从.config提取,建议用正则匹配末尾数字,避免因格式变化出错:

tune_results <- wf_set %>%
  collect_metrics() %>%
  filter(.metric == "sse_ratio") %>%
  mutate(num_clusters = as.numeric(stringr::str_extract(.config, "\\d+$")))

tune_results %>%
  ggplot(aes(x = num_clusters, y = mean, color = wflow_id)) +
  geom_point() +
  geom_line() +
  theme_minimal() +
  ggtitle("WSS/TSS比率随聚类数变化图") +
  ylab("3折交叉验证下的平均WSS/TSS比率") +
  xlab("聚类数") +
  scale_x_continuous(breaks = 1:10)

内容的提问来源于stack exchange,提问作者TheGoat

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.27 16:44:57