tidymodels中workflow_set遇列长度不匹配错误求助
使用tidymodels构建workflow_set时的列长度不匹配错误排查
问题场景
在机器学习任务中,尝试用tidymodels的workflow_set批量测试多个模型,构建完成后查看recipe_models对象时出现错误,提示列长度不一致。
重现代码
预处理Recipe与模型规格定义
base_recipe = recipe(high_traffic ~ ., data = recipes) %>% # 标记recipe列为ID,不作为预测变量 update_role(recipe, new_role = "id") %>% # 为分类预测变量生成哑变量 step_dummy(all_nominal_predictors()) %>% # 对数值预测变量应用Yeo-Johnson变换 step_YeoJohnson(all_numeric_predictors()) # 普通逻辑回归 logreg_spec = logistic_reg() %>% set_engine("glm") %>% set_mode("classification") # 正则化逻辑回归 regularized_spec = logistic_reg(penalty = tune(), mixture = tune()) %>% set_engine("glm") %>% set_mode("classification") # KNN模型 knn_spec = nearest_neighbor(neighbors = tune(), weight_func = tune()) %>% set_engine("kknn") %>% set_mode("classification") # XGBoost模型 xgb_spec = boost_tree(learn_rate = tune(), trees = tune()) %>% set_engine("xgboost") %>% set_mode("classification") # 神经网络模型 hidden_units <- floor(0.67 * (ncol(recipes_train) - 1)) + 1 neural_network = mlp(epochs = 1000, hidden_units = hidden_units, dropout = tune(), learn_rate = 0.001) %>% set_engine("keras") %>% set_mode("classification")
构建workflow_set
recipe_models = workflow_set( preproc = list(simple = base_recipe), models = list( logreg = logreg_spec, regularized_logreg = regularized_spec, knn = knn_spec, xgb = xgb_spec, nn = neural_network ), cross = TRUE )
错误信息
ERROR while rich displaying an object: Error in names(res) <- prefix: 'names' attribute [1] must be the same length as the vector [0]
数据集前50行
recipe,calories,carbohydrate,sugar,protein,category,servings,high_traffic 001,NA,NA,NA,NA,Pork,6,High 002,35.48,38.56,0.66,0.92,Potato,4,High 003,914.28,42.68,3.09,2.88,Breakfast,1,NA 004,97.03,30.56,38.63,0.02,Beverages,4,High 005,27.05,1.85,0.8,0.53,Beverages,4,NA 006,691.15,3.46,1.65,53.93,One Dish Meal,2,High 007,183.94,47.95,9.75,46.71,Chicken Breast,4,NA 008,299.14,3.17,0.4,32.4,Lunch/Snacks,4,NA 009,538.52,3.78,3.37,3.79,Pork,6,High 010,248.28,48.54,3.99,113.85,Chicken,2,NA 011,170.12,17.63,4.1,0.91,Beverages,1,NA 012,155.8,8.27,9.78,11.55,Breakfast,6,NA 013,274.63,23.49,1.56,2.57,Potato,4,High 014,25.23,11.51,10.32,9.57,Vegetable,4,High 015,217.14,6.69,10,15.17,Meat,4,High 016,316.45,2.65,4.68,79.71,Meat,6,High 017,454.27,1.87,2.95,61.07,Meat,2,High 018,1695.82,0.1,0.39,33.17,Meat,1,High 019,1090.75,4.65,0.69,3.49,Meat,6,High 020,127.55,27.55,1.51,8.91,Chicken,2,NA 021,9.26,17.44,8.16,10.81,Potato,6,High 022,40.53,87.91,104.91,11.93,Dessert,4,NA 023,82.73,3.17,7.95,26.04,Breakfast,4,NA 024,NA,NA,NA,NA,Meat,2,NA 025,1161.49,1.53,8.88,12.57,Breakfast,1,High 026,56.29,22.35,11.38,34.79,One Dish Meal,4,High 027,411.16,51.7,27.78,70.3,Pork,2,High 028,574.75,13.12,1.84,13.85,Potato,4,High 029,595.39,62.67,2.64,4.96,Potato,2,High 030,164.76,33.58,17.87,220.14,One Dish Meal,2,High 031,215.98,52.66,6.25,32.32,Pork,2,High 032,617.11,23.1,32.83,45.89,Breakfast,6,NA 033,347.06,9.5,5.92,82.58,Chicken Breast,4,NA 034,497.17,1.47,1.51,2.97,Lunch/Snacks,6,High 035,575.63,20.71,0.2,6.24,Breakfast,6,High 036,796.89,29.1,9.63,2.28,Lunch/Snacks,2,NA 037,1321.78,70.07,7.75,19.51,Breakfast,1,NA 038,44.55,99.82,2.62,15.57,Breakfast,4,NA 039,264.62,1.5,18.44,32.62,Chicken Breast,4,NA 040,44.81,4.62,0.4,5.9,Vegetable,4,High 041,621.54,14.16,10.7,39.69,Chicken Breast,6,High 042,290.1,4.43,1.05,40.64,Breakfast,4,NA 043,576.89,4.79,20.92,4.29,One Dish Meal,2,NA 044,262.12,17.46,0.33,87.05,Chicken Breast,4,NA 045,64.29,16.95,0.77,11.2,Breakfast,1,NA 046,83.39,13.06,1.62,3.44,Vegetable,6,High 047,69.01,39.17,39.54,0.17,Beverages,4,NA 048,43.91,48.16,4.58,7.92,Breakfast,4,NA 049,NA,NA,NA,NA,Chicken Breast,4,NA 050,1724.25,45.52,0.07,49.37,Breakfast,1,High
数据清洗代码
recipes = recipes %>% # 清洗servings列 mutate(servings = ifelse(servings == "4 as a snack", "4", servings)) %>% mutate(servings = ifelse(servings == "6 as a snack", "6", servings)) %>% mutate(servings = factor(servings)) %>% # 转换high_traffic列:NA转为0,High转为1,再转为因子 mutate(high_traffic = ifelse(is.na(high_traffic), 0, 1)) %>% mutate(high_traffic = factor(high_traffic)) %>% # 将category转为因子 mutate(category = factor(category)) %>% # 按category分组填充数值列的缺失值(用组内中位数) group_by(category) %>% mutate(across( .cols = where(is.numeric), .fns = ~ifelse(is.na(.x), median(.x, na.rm = TRUE), .x), .names = "{col}" )) %>% ungroup()
问题原因与解决方案
核心原因1:正则化逻辑回归引擎错误
regularized_spec中使用了set_engine("glm"),但基础glm引擎不支持penalty和mixture正则化参数,这会导致模型规格无效,进而在workflow_set内部处理时引发列不匹配错误。
修正代码:
regularized_spec = logistic_reg(penalty = tune(), mixture = tune()) %>% set_engine("glmnet") %>% # 替换为支持正则化的glmnet引擎 set_mode("classification")
核心原因2:未定义recipes_train对象
神经网络模型中计算hidden_units时依赖recipes_train,但代码中未提供该数据集的定义(未做训练/测试集划分),若recipes_train不存在,会导致hidden_units计算异常,影响workflow的构建。
修正代码:
先划分数据集:
library(rsample) set.seed(123) recipe_split = initial_split(recipes, prop = 0.7) recipes_train = training(recipe_split) recipes_test = testing(recipe_split)
潜在原因3:分组填充缺失值可能产生NaN
部分category组的数值列可能全为NA,导致median(.x, na.rm = TRUE)返回NaN,进而数据中存在无效值,干扰预处理流程。
检查与修正:
检查数值列是否存在NaN:
recipes %>% select(where(is.numeric)) %>% summarise(across(everything(), ~sum(is.na(.x))))
若存在NaN,可改用全局中位数填充,或删除异常分组:
# 改用全局中位数填充 recipes = recipes %>% mutate(across( .cols = where(is.numeric), .fns = ~ifelse(is.na(.x), median(.x, na.rm = TRUE), .x), .names = "{col}" ))
内容的提问来源于stack exchange,提问作者Giulio Mario Martena
相关产品推荐
相关产品推荐

