交叉验证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
相关产品推荐
相关产品推荐

