递归构建决策树时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
相关产品推荐
相关产品推荐

