如何用tidymodels对SSLR包半监督算法进行网格搜索调参?
超参数网格搜索解决方案(适配SSLR转导测试场景)
针对你用tidymodels对SSLR包半监督算法做超参数网格搜索的需求,结合转导测试的逻辑,这里提供两种可行方案,其中手动遍历的方式更直接适配SSLR的无标注数据处理逻辑:
方案一:手动遍历网格+重复划分(推荐)
这种方式不需要额外封装模型接口,直接贴合你描述的拟合/预测流程,实现起来更直观:
步骤1:加载依赖包并准备数据
library(tidymodels) library(SSLR) library(caret) # 确保目标变量为因子类型(分类任务) dat <- dat |> mutate(target = factor(target))
步骤2:定义超参数网格
为C和Cstar各设置4个取值:
param_grid <- grid_regular( C() %>% range_set(c(0.01, 10)), # 可根据需求调整取值范围 Cstar() %>% range_set(c(0.01, 10)), levels = 4 )
步骤3:生成重复的标注/无标注划分
模拟重复2折交叉验证,生成5次重复的50%无标注数据划分:
n_repeats <- 5 # 重复次数可调整 resample_list <- map(1:n_repeats, function(i) { # 每次划分50%数据作为无标注 unlabeled_idx <- caret::createDataPartition(dat$target, p = .5, list = FALSE) labeled_idx <- setdiff(1:nrow(dat), unlabeled_idx) list( labeled_idx = labeled_idx, unlabeled_idx = unlabeled_idx, full_data = dat ) })
步骤4:遍历网格与划分,执行拟合与评估
# 展开参数网格与重采样组合 grid_results <- expand_grid( resample = resample_list, params = param_grid ) |> mutate( # 拟合模型:将无标注数据的target设为NA fitted_model = pmap(list(resample, params), function(res, params) { fit_data <- res$full_data fit_data$target[res$unlabeled_idx] <- NA sslr_model <- LinearTSVMSSLR(C = params$C, Cstar = params$Cstar) fit(sslr_model, fit_data) }), # 预测无标注数据 predictions = pmap(list(fitted_model, resample), function(mod, res) { predict(mod, res$full_data[res$unlabeled_idx, ]) }), # 计算准确率 accuracy = pmap_dbl(list(predictions, resample), function(preds, res) { true_labels <- res$full_data$target[res$unlabeled_idx] mean(preds == true_labels) }) )
步骤5:汇总结果并选最优参数
# 按参数组合汇总平均准确率与标准差 summary_results <- grid_results |> group_by(C, Cstar) |> summarize( mean_acc = mean(accuracy), sd_acc = sd(accuracy), .groups = "drop" ) # 按平均准确率排序,选出最优参数 best_params <- summary_results |> arrange(desc(mean_acc)) |> slice(1)
方案二:适配tidymodels原生框架(需自定义模型接口)
如果你希望严格遵循tidymodels的工作流,需要先为SSLR模型创建parsnip自定义接口,再修改重采样逻辑:
步骤1:自定义parsnip模型规范
# 定义模型规范 lineartsvm_sslr <- function(mode = "classification", engine = "SSLR", C = NULL, Cstar = NULL) { parsnip::new_model_spec( "lineartsvm_sslr", args = list(C = rlang::enquo(C), Cstar = rlang::enquo(Cstar)), mode = mode, method = list( fit = "lineartsvm_sslr_fit", predict = "lineartsvm_sslr_predict" ), engine = engine ) } # 自定义拟合函数:处理含NA的目标列 lineartsvm_sslr_fit <- function(x, y, args, ...) { fit_data <- cbind(x, target = y) sslr_model <- SSLR::LinearTSVMSSLR(C = args$C, Cstar = args$Cstar) fit(sslr_model, fit_data) } # 自定义预测函数:返回tidymodels格式的结果 lineartsvm_sslr_predict <- function(object, new_data, type = "class", ...) { preds <- predict(object, new_data) tibble::tibble(.pred_class = preds) }
步骤2:创建工作流与自定义重采样
# 创建工作流 wf <- workflow() |> add_model(lineartsvm_sslr(mode = "classification", C = tune(), Cstar = tune())) |> add_formula(target ~ .) # 生成自定义重采样:每个折包含标注/无标注索引 resamples <- vfold_cv(dat, v = 2, repeats = 5) |> mutate( split = map(split, function(spl) { full_dat <- spl$data unlabeled_idx <- testing(spl) |> row.names() |> as.integer() # 修改split,保留无标注索引属性 new_split <- make_splits(indices = training(spl) |> row.names() |> as.integer(), data = full_dat) attr(new_split, "unlabeled_idx") <- unlabeled_idx new_split }) )
步骤3:自定义评估逻辑并执行网格搜索
由于tidymodels默认仅用训练集拟合,需要自定义控制逻辑来处理无标注数据,这里可以用tune::tune_grid()结合自定义的fit_resamples()变种,或者直接用方案一的手动方式更高效。
内容的提问来源于stack exchange,提问作者Jacques Wainer
相关产品推荐
相关产品推荐

