解决BART模型输出中因子水平名称混淆的问题
解决BART模型中因子水平变量名混乱的问题
问题描述
在二分类任务中训练BART模型时,预测变量Type.1是包含18个水平的因子,原始水平如下:
levels(poke$Type.1) [1] "Bug" "Dark" "Dragon" "Electric" "Fairy" "Fighting" [7] "Fire" "Flying" "Ghost" "Grass" "Ground" "Ice" [13] "Normal" "Poison" "Psychic" "Rock" "Steel" "Water"
但模型输出的varcount.mean数组(用于判断变量重要性)中,因子水平被自动命名为Type.11、Type.12这类难以解读的名称:
bart1$varcount.mean Type.11 Type.12 Type.13 Type.14 Type.15 Type.16 1.930 1.825 1.782 1.804 1.983 1.864 Type.17 Type.18 Type.19 Type.110 Type.111 Type.112 1.913 0.000 1.950 1.977 1.983 1.987 Type.113 Type.114 Type.115 Type.116 Type.117 Type.118 2.105 2.004 1.871 2.102 1.906 2.342 Total HP Attack Defense Sp..Atk Sp..Def 4.400 2.266 2.415 2.175 2.508 2.652 Speed Generation 2.711 2.281
需要无需手动逐行重命名,通过预处理或便捷函数恢复原始因子水平名称。
解决方案
方案1:预处理时手动生成哑变量(推荐)
BART默认会自动将因子转为哑变量,但命名规则不友好。可以提前把Type.1转为带原始水平名称的哑变量,这样模型输出的变量名会直接对应原始水平:
library(BART) library(tidyr) library(dplyr) # 读取数据并预处理目标变量 pokeB <- read.csv("~/Downloads/Pokemon.csv", header=T) pokeB$Legend <- as.integer(ifelse(pokeB$Legendary=="True", 1, 0)) # 将Type.1转为带原始水平名称的哑变量 poke_dummy <- pokeB %>% select(Type.1, Total, HP, Attack, Defense, Sp..Atk, Sp..Def, Speed, Generation, Legend) %>% mutate(Type.1 = factor(Type.1, levels = levels(.$Type.1))) %>% pivot_wider(names_from = Type.1, values_from = Type.1, names_prefix = "Type.1_", values_fn = ~ifelse(!is.na(.), 1, 0)) # 划分训练集和测试集 set.seed(1) train <- sample(1:nrow(poke_dummy), nrow(poke_dummy)/2) x <- poke_dummy %>% select(-Legend) y <- poke_dummy$Legend xtrain <- x[train,] ytrain <- y[train] xtest <- x[-train,] ytest <- y[-train] # 训练模型 bart1 <- mc.gbart(xtrain, ytrain, x.test=xtest, type='pbart', mc.cores=4)
此时varcount.mean中的变量名会变成Type.1_Bug、Type.1_Dark等,直接对应原始水平,无需后续修改。
方案2:事后批量重命名变量名
如果已经完成模型训练,可以通过字符串匹配和替换,自动将Type.11类名称映射回原始水平:
# 获取原始因子水平 type_levels <- levels(poke$Type.1) # 获取varcount.mean的当前变量名 var_names <- names(bart1$varcount.mean) # 批量替换变量名 new_var_names <- sapply(var_names, function(name) { if (grepl("^Type\\.1\\d+$", name)) { # 提取数字部分:Type.11提取1,Type.110提取10 num <- as.integer(sub("Type\\.1(\\d+)", "\\1", name)) # 映射到原始水平并生成新名称 paste0("Type.1_", type_levels[num]) } else { # 非Type.1变量保持原名 name } }) # 重命名varcount.mean names(bart1$varcount.mean) <- new_var_names # 重新生成排序后的变量重要性数据框 ord1 <- order(bart1$varcount.mean, decreasing=T) vars1 <- as.data.frame(bart1$varcount.mean[ord1])
这段代码会自动匹配所有Type.1开头的变量名,根据数字对应到原始因子水平,完成批量重命名。
注意事项
- BART的
mc.gbart函数没有直接保留因子水平名称的参数,因此预处理生成哑变量是更稳妥的方式。 - 方案2的替换逻辑依赖于因子水平的顺序与自动编码的数字顺序一致,需确保训练前
poke$Type.1的水平顺序未被修改。
内容的提问来源于stack exchange,提问作者JLosc
相关产品推荐
相关产品推荐

