如何确保训练集与测试集的model.matrix列数一致?求更优方法
这确实是建模时特别容易踩的坑——训练集和测试集的因子水平不一致,直接用model.matrix就会导致列数不匹配!你的思路完全正确:先从训练集提取因子水平,再统一应用到测试集,不过确实有更简洁、成熟的实现方式,我给你分享几个常用的方案:
方法1:用tidymodels的recipes包(推荐)
这是现在R语言建模生态里最规范的预处理工具,能把训练集的所有预处理规则(包括因子编码、缺失值填充、标准化等)保存下来,一键应用到测试集,完全避免列不匹配的问题:
library(tidyverse) library(recipes) # 构造原始数据 x_ent <- tibble(x1 = c(1, 2, 3, 4), x2 = c('a', 'b', 'a', 'c')) %>% mutate_if(is.character, as.factor) x_pr <- tibble(x1 = c(5, 6, 7), x2 = c('a', 'b', 'a')) %>% mutate_if(is.character, as.factor) # 创建预处理配方:对所有分类变量做哑编码,保留截距 rec <- recipe(~ ., data = x_ent) %>% step_intercept() %>% # 对应model.matrix的默认截距项 step_dummy(all_nominal_predictors(), one_hot = FALSE) # 按训练集因子水平生成哑变量 # 基于训练集训练配方(提取规则) rec_prepped <- prep(rec, training = x_ent) # 分别应用到训练集和测试集生成矩阵 x_mat_ent <- bake(rec_prepped, new_data = x_ent) %>% as.matrix() x_mat_pr <- bake(rec_prepped, new_data = x_pr) %>% as.matrix() # 检查维度,两者列数完全一致 dim(x_mat_ent) # 4 4 dim(x_mat_pr) # 3 4
方法2:用caret包的dummyVars
这是传统建模流程里常用的工具,专门用于生成哑变量,同样能基于训练集锁定编码规则:
library(tidyverse) library(caret) # 构造原始数据 x_ent <- tibble(x1 = c(1, 2, 3, 4), x2 = c('a', 'b', 'a', 'c')) %>% mutate_if(is.character, as.factor) x_pr <- tibble(x1 = c(5, 6, 7), x2 = c('a', 'b', 'a')) %>% mutate_if(is.character, as.factor) # 基于训练集创建哑变量生成规则 dv <- dummyVars(~ ., data = x_ent) # 生成训练集和测试集的矩阵 x_mat_ent <- predict(dv, newdata = x_ent) x_mat_pr <- predict(dv, newdata = x_pr) # 检查维度 dim(x_mat_ent) # 4 3 dim(x_mat_pr) # 3 3
方法3:优化你自己的实现
你的核心逻辑没问题,但原函数有个小bug——没有返回encode_test函数,导致调用后会得到NULL。这里给你优化得更简洁:
library(tidyverse) create_encoder <- function(train_data){ # 提取训练集中的因子变量及其水平 cat_vars <- train_data %>% select_if(is.factor) %>% names() label_dic <- map(cat_vars, ~levels(train_data[[.]])) %>% set_names(cat_vars) # 返回编码测试集的函数 function(test_data){ test_data %>% modify_at(cat_vars, ~factor(., levels = label_dic[[cur_column()]])) } } # 使用示例 x_ent <- tibble(x1 = c(1, 2, 3, 4), x2 = c('a', 'b', 'a', 'c')) %>% mutate_if(is.character, as.factor) x_pr <- tibble(x1 = c(5, 6, 7), x2 = c('a', 'b', 'a')) %>% mutate_if(is.character, as.factor) encode_test <- create_encoder(x_ent) x_pr_encoded <- encode_test(x_pr) # 生成矩阵,列数完全一致 x_mat_ent <- model.matrix(~., x_ent) x_mat_pr <- model.matrix(~., x_pr_encoded) dim(x_mat_ent) # 4 4 dim(x_mat_pr) # 3 4
总结
如果只是处理因子编码,自己写函数足够灵活;但如果后续还要做更多预处理(比如缺失值、标准化),优先用recipes或caret,它们能一站式解决所有预处理一致性问题。
内容的提问来源于stack exchange,提问作者Alfie González
相关产品推荐
相关产品推荐

