R语言如何获取随机森林中每棵树的分类概率
问题原因
randomForest包原生predict方法本身未提供返回单棵树分类概率的内置逻辑,你尝试的参数组合失效属于包本身的设计限制:在predict.all=TRUE的代码分支中,开发者硬编码了返回单树的类别响应结果,未对type="prob"参数做适配,因此无论是否指定概率输出类型,$individual字段只会返回分类标签,不会返回概率值。
可行实现方案
方案1:不依赖第三方包,基于现有randomForest模型解析计算
不需要重新训练模型,通过提取树结构+终端节点匹配的方式即可得到所有单树的分类概率,代码如下:
# 1. 获取每个样本在每棵树上落到的终端节点ID node_mapping <- attr(predict(rf_cl, newdata, nodes = TRUE), "nodes") class_levels <- levels(rf_cl$y) n_class <- length(class_levels) n_sample <- nrow(newdata) n_tree <- rf_cl$ntree # 2. 初始化结果存储数组:维度为 [样本数, 类别数, 树序号] single_tree_prob <- array( data = NA_real_, dim = c(n_sample, n_class, n_tree), dimnames = list(rownames(newdata), class_levels, 1:n_tree) ) # 3. 逐棵树计算终端节点概率,匹配到对应样本 for (i in 1:n_tree) { # 提取第i棵树的完整结构 tree_info <- getTree(rf_cl, k = i, labelVar = TRUE) # 筛选终端节点,计算每个终端节点的类别概率(节点内类别投票数/节点总样本数) terminal_mask <- tree_info$status == -1 terminal_nodes <- tree_info[terminal_mask, ] vote_cols <- paste0("votes.", class_levels) terminal_prob <- terminal_nodes[, vote_cols] / rowSums(terminal_nodes[, vote_cols]) # 终端节点ID存储在左女儿字段,作为匹配键 rownames(terminal_prob) <- terminal_nodes$`left daughter` # 按样本对应的节点ID匹配概率值 sample_node_id <- node_mapping[, i] match_loc <- match(sample_node_id, rownames(terminal_prob)) single_tree_prob[,,i] <- as.matrix(terminal_prob[match_loc, ]) }
调用方式:如果需要获取第k棵树的分类概率,直接取single_tree_prob[,,k]即可,返回结果行对应样本,列对应类别,值为该样本在对应类别上的单树预测概率,和单棵树独立调用predict(type="prob")的结果完全一致。
方案2:换用ranger包训练模型(更简便)
如果不需要拘泥于randomForest包,可以换用效率更高的ranger包实现随机森林,该包原生支持返回单棵树的分类概率,不需要手动解析树结构:
library(ranger) # 训练时需要指定probability=TRUE开启概率输出,keep.inbag=TRUE保留树结构用于单树预测 rf_model <- ranger( formula = 你的因变量 ~ ., data = 你的训练集, probability = TRUE, keep.inbag = TRUE ) # 预测时开启predict.all即可拿到所有单树结果 pred_res <- predict(rf_model, data = newdata, predict.all = TRUE) # pred_res$predictions为三维数组:维度[样本数, 类别数, 树数量],直接就是需要的单树概率 single_tree_prob_ranger <- pred_res$predictions
内容的提问来源于stack exchange,提问作者Whitney
相关产品推荐
相关产品推荐

