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

使用class_weight='balanced'的随机森林可视化树时value非整数问题求助

解决Random Forest设置class_weight='balanced'后决策树节点value显示小数的问题

当给Random Forest设置class_weight='balanced'时,Scikit-learn会自动为每个类别计算权重(公式:权重 = 总样本数 / (类别数量 × 该类样本数)),决策树节点的value字段存储的是加权后的样本数,因此会显示小数,而非原始的整数样本计数。

解决思路

要恢复显示原始整数样本数,只需将每个节点的value值除以对应类别的权重,再取整即可(因为加权值 = 原始样本数 × 权重,所以原始样本数 = 加权值 / 权重)。

完整实现代码

import graphviz
from sklearn import tree
from sklearn.utils.class_weight import compute_class_weight
import numpy as np
import re

# 假设你的训练标签为y,训练好的Random Forest模型为model
trees = model.estimators_

# 1. 计算与模型训练时一致的类别权重
classes = np.unique(y)
class_weights = compute_class_weight('balanced', classes=classes, y=y)
weight_map = {cls: w for cls, w in zip(classes, class_weights)}

# 2. 定义修正节点value显示的函数
def fix_node_value_text(dot_str, weight_map):
    # 匹配决策树节点中的value=[x, y]格式
    value_pattern = r'value=\[([\d.]+), ([\d.]+)\]'
    def replace_func(match):
        # 提取加权后的数值并转换为原始样本数
        weighted_vals = [float(match.group(1)), float(match.group(2))]
        original_vals = [round(val / weight_map[cls]) for val, cls in zip(weighted_vals, classes)]
        return f'value=[{original_vals[0]}, {original_vals[1]}]'
    # 替换所有匹配的value字段
    return re.sub(value_pattern, replace_func, dot_str)

# 3. 生成并修正dot数据,可视化决策树
dot_data = tree.export_graphviz(
    trees[0], 
    out_file=None, 
    filled=True, 
    rounded=True, 
    special_characters=True
    # 可添加feature_names=X_rf.columns参数显示特征名称
)
fixed_dot_data = fix_node_value_text(dot_data, weight_map)
graph = graphviz.Source(fixed_dot_data)
graph

注意事项

  • 如果你的类别标签不是0和1(比如自定义字符串标签),代码会自动适配,因为classes变量会获取数据中的实际类别列表
  • 使用round()取整是因为加权计算可能存在微小浮点误差,结果会非常接近整数
  • 该方法仅修改可视化的显示内容,完全不影响模型本身的性能和预测逻辑

内容的提问来源于stack exchange,提问作者Edo Wexler

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.21 12:37:26