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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.15 00:02:33