如何使用quanteda实现交叉验证?针对类别不平衡数据优化F1值
使用 quanteda + caret 实现类别不平衡文本的交叉验证(优化F1值)
刚好我之前处理过类似的场景,quanteda和caret其实可以通过自定义模型接口完美兼容,而且针对类别不平衡的交叉验证也有成熟的方案,下面一步步给你讲清楚:
首先先重现你的示例场景,方便后续对照:
示例数据与样本内训练代码
示例数据集
library(tidyverse) library(quanteda) library(caret) dtrain <- data_frame(text = c("Chinese Beijing Chinese", "Chinese Chinese Shanghai", "Chinese Macao", "Tokyo Japan Chinese"), doc_id = 1:4, class = c("Y", "Y", "Y", "N")) dtrain <- dtrain %>% mutate(class = as.factor(class)) # 查看数据结构 dtrain # A tibble: 4 x 3 # text doc_id class # <chr> <int> <fct> #1 Chinese Beijing Chinese 1 Y #2 Chinese Chinese Shanghai 2 Y #3 Chinese Macao 3 Y #4 Tokyo Japan Chinese 4 N
你的样本内训练代码
trainingset <- dfm(corpus(dtrain, docid_field = 'doc_id', text_field = 'text')) nb_test <- textmodel_nb(trainingset, docvars(trainingset, "class"), prior = "docfreq") myprediction <- predict(nb_test, trainingset)$nb.predicted confusionMatrix(table(dtrain$class, myprediction), mode = 'prec_recall') # 输出结果 # Confusion Matrix and Statistics # # myprediction # N Y # N 1 0 # Y 0 3 # # Accuracy : 1 # 95% CI : (0.3976, 1) # No Information Rate : 0.75 # P-Value [Acc > NIR] : 0.3164 # # Kappa : 1 # # Mcnemar's Test P-Value : NA # # Precision : 1.00 # Recall : 1.00 # F1 : 1.00
核心解决方案:自定义caret兼容的quanteda模型
因为caret默认没有集成quanteda的模型,我们可以通过自定义训练/预测函数,把quanteda的文本处理和朴素贝叶斯模型接入caret的交叉验证框架,同时针对类别不平衡优化F1值。
步骤1:编写文本预处理工具函数
这个函数负责把文本向量转换成dfm,并且保证测试集和训练集的特征一致:
text_to_dfm <- function(text_vector, train_dfm = NULL) { corpus_obj <- corpus(text_vector) dfm_obj <- dfm(corpus_obj) # 测试集匹配训练集的特征,避免特征不一致 if (!is.null(train_dfm)) { dfm_obj <- dfm_match(dfm_obj, features = featnames(train_dfm)) } return(dfm_obj) }
步骤2:定义caret兼容的训练和预测函数
我们需要为textmodel_nb写适配caret的训练和预测逻辑:
# 训练函数:输入dfm矩阵和类别变量,返回quanteda的NB模型 train_nb_quanteda <- function(x, y, weights = NULL, ...) { textmodel_nb(x, y, ...) } # 预测函数:输入模型和新文本,返回预测类别 predict_nb_quanteda <- function(modelFit, newdata, preProc = NULL, submodels = NULL) { # 把新文本转换成dfm,并匹配训练集的特征 new_dfm <- text_to_dfm(newdata, train_dfm = modelFit$x) predict(modelFit, new_dfm)$nb.predicted }
步骤3:把自定义模型注册到caret
用setModelInfo把我们的模型添加到caret的模型列表中:
setModelInfo( model = "quantedaNB", list( label = "Quanteda Naive Bayes", library = "quanteda", type = c("Classification"), parameters = data.frame(parameter = "prior", class = "character", label = "Prior Type"), grid = function(x, y, len = NULL, search = "grid") { # 可以测试不同的先验类型 data.frame(prior = c("docfreq", "uniform")) }, fit = train_nb_quanteda, predict = predict_nb_quanteda, predictors = function(x) featnames(x$x), levels = function(x) x$y$levels, sort = function(x) x[order(x$prior), ] ) )
步骤4:配置交叉验证策略(针对类别不平衡)
因为你的数据类别高度不平衡,我们用分层交叉验证保证每折的类别比例和原始数据一致,同时选择F1作为优化指标(比准确率更适合不平衡数据集):
train_control <- trainControl( method = "cv", number = 5, # 5折交叉验证 stratified = TRUE, # 分层保持类别比例 summaryFunction = prSummary, # 用精确率/召回率/F1作为评估指标 classProbs = FALSE ) # 设定优化目标为F1值 metric <- "F1"
步骤5:用caret训练带交叉验证的模型
这里直接传入文本列即可,自定义函数会自动处理文本到dfm的转换:
nb_caret <- train( x = dtrain$text, y = dtrain$class, method = "quantedaNB", trControl = train_control, metric = metric, tuneGrid = data.frame(prior = "docfreq") # 可以换成"uniform"测试效果 ) # 查看交叉验证结果 print(nb_caret)
额外优化:针对严重不平衡的SMOTE过采样
如果你的数据不平衡程度极高(比如少数类占比<10%),可以在交叉验证的每折中加入SMOTE过采样,进一步平衡数据集:
library(DMwR) # 需要先安装这个包 train_control_smote <- trainControl( method = "cv", number = 5, stratified = TRUE, summaryFunction = prSummary, sampling = "smote" # 对每折训练集做SMOTE过采样 ) # 用SMOTE训练模型 nb_caret_smote <- train( x = dtrain$text, y = dtrain$class, method = "quantedaNB", trControl = train_control_smote, metric = metric, tuneGrid = data.frame(prior = "docfreq") )
方案优势
- 完美兼容quanteda的文本处理能力和caret的交叉验证框架
- 分层交叉验证避免了不平衡数据集的验证偏差
- 以F1为优化指标,更贴合不平衡分类的实际需求
- 可选的SMOTE过采样进一步缓解类别不平衡问题
内容的提问来源于stack exchange,提问作者ℕʘʘḆḽḘ
相关产品推荐
相关产品推荐

