You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

解决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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.22 10:42:05