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

Tidymodels中能否用step_dummy()实现多分类的一对多二分类?

解决tidymodels中多分类转一对多(One-vs-Rest)二分类的问题

问题根源

你当前的配方写法存在两个核心问题:

  1. 初始配方公式~ value1 + value2未包含结果变量prevalent_topic,导致step_dummy无法找到该变量生成哑编码;
  2. step_dummy的设计初衷是处理预测变量的因子哑编码,而非将结果变量转为one-vs-rest的二分类标签。

可行解决方案

方案1:使用discrim包的one_vs_all()自动处理(推荐)

tidymodels生态中的discrim包提供了one_vs_all()函数,可以快速将二分类模型包装为one-vs-rest的多分类模型,无需手动生成哑变量:

library(tidymodels)
library(discrim)

# 1. 创建基础配方(标准化预测变量)
base_rec <- recipe(prevalent_topic ~ value1 + value2, data = dfFT_train) %>%
  step_normalize(value1, value2)

# 2. 用one_vs_all包装逻辑回归模型
ovr_model <- one_vs_all(
  logistic_reg() %>% 
    set_engine("glm") %>%  # 也可以用"glmnet"等其他引擎
    set_mode("classification")
)

# 3. 构建workflow并拟合
ovr_wf <- workflow() %>%
  add_recipe(base_rec) %>%
  add_model(ovr_model)

ovr_fit <- fit(ovr_wf, data = dfFT_train)

拟合完成后,你可以直接用这个模型进行预测,它会自动输出每个样本属于各个主题的概率,或直接预测类别。

方案2:手动循环生成每个二分类任务

如果你希望手动控制每个one-vs-rest的二分类模型,可以针对每个主题单独构建workflow:

library(tidymodels)

# 获取所有主题类别
topics <- levels(dfFT_train$prevalent_topic)

# 循环创建每个主题的二分类workflow
ovr_workflows <- map(topics, function(target_topic) {
  # 构建配方:创建二分类结果变量(是否为当前主题)
  rec <- recipe(prevalent_topic ~ value1 + value2, data = dfFT_train) %>%
    step_normalize(value1, value2) %>%
    # 生成二分类标签,skip=TRUE避免在测试集重新计算
    step_mutate(outcome = as.factor(prevalent_topic == target_topic), skip = TRUE) %>%
    # 更新变量角色:将新生成的outcome设为结果,原prevalent_topic转为预测变量(避免干扰)
    update_role(outcome, new_role = "outcome") %>%
    update_role(prevalent_topic, new_role = "predictor")
  
  # 定义二分类模型
  model <- logistic_reg() %>%
    set_engine("glm") %>%
    set_mode("classification")
  
  # 组合workflow
  workflow() %>%
    add_recipe(rec) %>%
    add_model(model)
})

# 为workflow列表命名,方便后续调用
names(ovr_workflows) <- topics

# 拟合所有模型
ovr_fits <- map(ovr_workflows, fit, data = dfFT_train)

之后你可以单独调用每个模型进行预测,再整合结果(比如取概率最高的类别作为最终预测)。

方案3:手动生成哑变量后构建模型(不推荐)

如果一定要用哑变量作为结果变量,需要先在数据中生成所有哑变量,再针对每个哑变量构建workflow:

library(tidymodels)
library(tibble)

# 先在训练集中生成结果变量的哑编码
dfFT_train_dummies <- dfFT_train %>%
  # 生成哑变量(-1去掉截距项)
  mutate(dummy_vars = list(model.matrix(~prevalent_topic - 1, data = .) %>% as_tibble())) %>%
  unnest(dummy_vars)

# 获取所有哑变量列名
dummy_cols <- str_subset(names(dfFT_train_dummies), "prevalent_topic_")

# 为每个哑变量构建workflow
ovr_workflows_dummy <- map(dummy_cols, function(response_col) {
  rec <- recipe(reformulate(c("value1", "value2"), response = response_col), 
                data = dfFT_train_dummies) %>%
    step_normalize(value1, value2)
  
  model <- logistic_reg() %>%
    set_engine("glm") %>%
    set_mode("classification")
  
  workflow() %>% add_recipe(rec) %>% add_model(model)
})

names(ovr_workflows_dummy) <- dummy_cols

# 拟合模型
ovr_fits_dummy <- map(ovr_workflows_dummy, fit, data = dfFT_train_dummies)

这种方式需要额外处理数据,不如前两种方法简洁,仅适合特殊场景需求。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.13 03:50:01