如何用R的Tidymodels解决多分类模型的类别不平衡问题?
解决极端长尾多分类的类别不平衡问题
你遇到的是极端长尾多分类场景——1330行数据对应337个类别,且绝大多数类别仅5个样本,常规分层拆分、重采样方法完全不适用,核心原因:
initial_split的strata参数要求每个类别有足够样本分配到训练/测试集,你多数类别仅5个样本,0.7比例拆分后训练集仅3个、测试集仅2个,达不到分层的最小样本要求;- SMOTE/ROSE这类生成式重采样需要小类别有足够样本基础来生成可信合成数据,5个样本根本撑不起有效合成;下采样则会把大类砍到5个样本量级,直接废掉训练数据。
以下是针对性的解决方案:
1. 先调整数据拆分逻辑,保证测试集覆盖所有类别
放弃自动分层拆分,手动实现每个类别至少保留1个样本到测试集的拆分逻辑,避免测试集出现训练集未见过的类别:
set.seed(123) # 按Category分组,每个类别抽1个样本到测试集 test_samples <- model_data %>% group_by(Category) %>% slice_sample(n = 1) %>% ungroup() # 剩余样本作为训练集 training_set <- anti_join(model_data, test_samples, by = names(model_data)) test_set <- test_samples # 验证拆分后训练集的类别样本分布 table(training_set$Category) %>% sort(decreasing = TRUE)
2. 合并小类别(最有效的预处理步骤)
337个类别里绝大多数只有5个样本,模型根本学不到这些小类的特征模式,建议将样本量低于阈值(比如10)的类别合并为“其他”类,大幅减少类别数量同时保证小类样本的聚合:
# 统计每个类别的样本量 category_counts <- table(model_data$Category) %>% as.data.frame() colnames(category_counts) <- c("Category", "Count") # 合并样本量<10的类别为"Other" threshold <- 10 model_data_merged <- model_data %>% mutate(Category = ifelse(Category %in% category_counts$Category[category_counts$Count >= threshold], as.character(Category), "Other")) %>% mutate(Category = factor(Category)) # 查看合并后的类别分布 table(model_data_merged$Category) %>% sort(decreasing = TRUE)
合并后类别数量会大幅压缩(比如你原有9个类别样本量≥6,剩余328个合并为Other,总类别数仅10个),后续的分层拆分、模型训练都会变得可行。
3. 给模型加类别权重,让模型关注小类
如果无法合并类别,就在训练时给小类别设置更高的权重,以tidymodels的XGBoost模型为例:
library(tidymodels) library(xgboost) # 计算类别权重:总样本数/(类别数*类别样本数),让小类权重更高 class_weights <- 1 / table(training_set$Category) class_weights <- class_weights / sum(class_weights) * length(class_weights) # 定义带权重的XGBoost模型 xgb_spec <- boost_tree(tree_depth = tune(), learn_rate = tune()) %>% set_engine("xgboost", scale_pos_weight = class_weights) %>% set_mode("classification") # 构建工作流 model_wf <- workflow() %>% add_recipe(model_recipe) %>% add_model(xgb_spec) # 合并类别后可使用分层交叉验证 cv_folds <- vfold_cv(training_set, strata = Category, v = 5)
4. 换用对不平衡数据友好的模型
比如朴素贝叶斯模型,对极端不平衡的多分类场景适应性更强:
nb_spec <- naive_Bayes() %>% set_engine("klaR") %>% set_mode("classification") model_wf <- workflow() %>% add_recipe(model_recipe) %>% add_model(nb_spec) # 训练并评估 nb_fit <- fit(model_wf, data = training_set) nb_pred <- predict(nb_fit, test_set, type = "prob") %>% bind_cols(test_set %>% select(Category)) # 用宏F1评估(避免大类主导准确率) accuracy(nb_pred, truth = Category, estimate = .pred_class) f_meas(nb_pred, truth = Category, estimate = .pred_class, beta = 1, estimator = "macro")
关键注意事项
- 绝对不要用准确率作为主要评估指标,改用宏F1、加权F1或对数损失,这些指标能真实反映模型对小类的性能;
- 如果必须保留所有337个类别,直接放弃重采样思路,专注于加权模型和合适的评估指标——重采样在这种极端长尾场景下完全无效。
内容的提问来源于stack exchange,提问作者Harry Kalsted
相关产品推荐
相关产品推荐

