Tidymodels中能否用step_dummy()实现多分类的一对多二分类?
解决tidymodels中多分类转一对多(One-vs-Rest)二分类的问题
问题根源
你当前的配方写法存在两个核心问题:
- 初始配方公式
~ value1 + value2未包含结果变量prevalent_topic,导致step_dummy无法找到该变量生成哑编码; 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
相关产品推荐
相关产品推荐

