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

递归构建决策树时self未更新,find_best_split调用异常求助

问题修复方案

你的核心问题是在递归函数中错误使用了self(原始根节点实例),而非传入的当前node参数,导致所有分裂操作都基于根节点的DataFrame,而非当前递归节点的子集数据。以下是具体修复步骤:

1. 替换递归中的self为当前node实例

在grow_tree_recursive函数的特征循环中,将self.find_best_split(feature)改为node.find_best_split(feature),确保使用当前节点的数据集计算最优分裂。

同时,所有节点属性赋值(如splitFeature、giniGain、node_left、node_right)都应指向当前node,而非self:

# 原错误代码
temp_feature_value,temp_gini,temp_left_node,temp_right_node,temp_feature_name = self.find_best_split(feature)
self.splitFeature = ""
self.giniGain = max_gini
self.node_left = grow_tree_recursive(Node(x_left,y_left),depth+1)

# 修复后
temp_feature_value,temp_gini,temp_left_node,temp_right_node,temp_feature_name = node.find_best_split(feature)
node.splitFeature = temp_feature_name
node.giniGain = max_gini
node.node_left = grow_tree_recursive(Node(x_left,y_left),depth+1)

2. 修正数据集子集生成逻辑

building_x函数原本使用根节点的self.x和self.y,但递归中的分裂索引是基于当前节点的数据集的,因此需要修改函数参数,传入当前节点的x和y:

# 原错误代码
def building_x(lst_idx):
    if (len(lst_idx)) == 0:
        return None
    else:
        df1 = self.x.iloc[lst_idx]
        df2 = self.y.iloc[lst_idx]
        return df1, df2

# 修复后
def building_x(x, y, lst_idx):
    if len(lst_idx) == 0:
        return None
    else:
        df1 = x.iloc[lst_idx]
        df2 = y.iloc[lst_idx]
        return df1, df2

# 调用时传入当前节点的x和y
x_left,y_left = building_x(node.x, node.y, left_split_data)
x_right,y_right = building_x(node.x, node.y, right_split_data)

3. 让递归函数返回构建完成的节点

grow_tree_recursive需要返回处理后的节点,否则node_left和node_right会被赋值为None:

def grow_tree_recursive(node,depth):
    # ... 原有逻辑 ...
    node.node_left = grow_tree_recursive(Node(x_left,y_left),depth+1)
    node.node_right = grow_tree_recursive(Node(x_right,y_right),depth+1)
    return node  # 新增:返回当前节点

完整修正代码

class Node:
    def __init__(self, x, y, min_leaf=5, class_weight=None, max_depth=5, depth=0):
        self.x = x
        self.y = y
        self.node_right = None
        self.node_left = None
        self.min_leaf = min_leaf
        self.class_weight = class_weight
        self.max_depth = max_depth
        self.depth = depth
        self.numOfSamples = len(y)
        self.target_labels_count = [y.tolist().count(0), y.tolist().count(1)]
        self.splitFeature = ""
        self.featureList = x.columns
        self.giniGain = 0

    def grow_tree(self):
        def building_x(x, y, lst_idx):
            if len(lst_idx) == 0:
                return None
            else:
                df1 = x.iloc[lst_idx]
                df2 = y.iloc[lst_idx]
                return df1, df2

        def grow_tree_recursive(node, depth):
            if depth == node.max_depth or node.numOfSamples < 2 * node.min_leaf:
                return node
            
            max_gini = 0
            left_split_data = None
            right_split_data = None
            best_feature = ""

            for feature in list(node.featureList):
                temp_feature_value, temp_gini, temp_left_node, temp_right_node, temp_feature_name = node.find_best_split(feature)
                if max_gini < temp_gini:
                    max_gini = temp_gini
                    left_split_data = temp_left_node
                    right_split_data = temp_right_node
                    best_feature = temp_feature_name

            if max_gini == 0:
                return node
            
            node.giniGain = max_gini
            node.splitFeature = best_feature

            x_left, y_left = building_x(node.x, node.y, left_split_data)
            x_right, y_right = building_x(node.x, node.y, right_split_data)

            node.node_left = grow_tree_recursive(Node(x_left, y_left, node.min_leaf, node.class_weight, node.max_depth, depth+1), depth+1)
            node.node_right = grow_tree_recursive(Node(x_right, y_right, node.min_leaf, node.class_weight, node.max_depth, depth+1), depth+1)
            
            return node

        return grow_tree_recursive(self, 0)

    def find_best_split(self, var_idx):
        dict1 = {}

        for i in self.x[var_idx]:
            if i not in dict1:
                dict1[i] = (0, [], [])
        
        for key in dict1:
            lhs = []
            rhs = []
            for i in range(len(self.x[var_idx])):
                if self.x[var_idx].iloc[i] <= key:
                    lhs.append(i)
                else:
                    rhs.append(i)
            if len(lhs) >= self.min_leaf and len(rhs) >= self.min_leaf:
                dict1[key] = (self.get_gini_gain(lhs, rhs), lhs, rhs)
            else:
                dict1[key] = (0, [], [])
        
        max_gini = max(dict1.items(), key=lambda x: x[1][0])
        return max_gini[0], max_gini[1][0], max_gini[1][1], max_gini[1][2], var_idx

    def get_gini_gain(self, lhs, rhs):
        total = len(lhs) + len(rhs)
        gini_parent = self.calculate_gini(self.y)
        gini_left = self.calculate_gini(self.y.iloc[lhs])
        gini_right = self.calculate_gini(self.y.iloc[rhs])
        weighted_gini = (len(lhs)/total)*gini_left + (len(rhs)/total)*gini_right
        return gini_parent - weighted_gini

    def calculate_gini(self, y):
        counts = y.value_counts(normalize=True)
        return 1 - sum(counts**2)

额外说明

  • 补充了get_gini_gain和calculate_gini方法(原代码中缺失,否则无法运行)
  • 增加了递归终止条件:达到最大深度或样本数不足2*min_leaf时停止分裂
  • 修正了find_best_split中访问特征值的方式:用self.x[var_idx].iloc[i]替代self.x[var_idx][i],避免因DataFrame索引非连续导致的错误

内容的提问来源于stack exchange,提问作者Shalev Levi Sagzan

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.24 05:07:04