决策树剪枝代码问题:新叶节点未更新至原树致准确率计算异常
决策树剪枝问题:新叶节点无法更新到原树
从零实现决策树及剪枝功能时,执行剪枝操作后,新生成的叶节点无法更新到原树中,导致计算新准确率时所有分支都被错误剪枝。核心需求是修改代码让剪枝操作能正确更新树结构,从而判断剪枝是否有效。
原代码
class TreeNode: def __init__(self, feature, split, depth, left = None, right = None): """ self.feature = the feature the node splits upon self.split = the value the of the feature the node splits upon self.left = the left child self.right = the right child self.depth = depth of the tree at this point """ self.feature = feature self.split = split self.left = left self.right = right self.depth = depth def getLeft(self): return self.left def getRight(self): return self.right def getFeature(self): return self.feature def getSplit(self): return self.split def getDepth(self): return self.depth def eval(self, sample): if sample[self.feature] < self.split: return self.left.eval(sample) else: return self.right.eval(sample) class LeafNode: def __init__(self, roomNumber, users, depth): """ self.roomNumer = the room number that the leaf is assigned self.depth = depth of the leaf in the tree """ self.roomNumber = roomNumber self.depth = depth self.users = users def getRoomNumber(self): return self.roomNumber def getDepth(self): return self.depth def getUsers(self): return self.users def eval(self, sample): self.users += 1 return self.getRoomNumber() def pruneTree(original_tree, validation, node): if node is None: return None if isinstance(node, LeafNode): return node node.left = pruneTree(original_tree, validation, node.left) node.right = pruneTree(original_tree, validation, node.right) if isinstance(node.left, LeafNode) and isinstance(node.right, LeafNode): current_accuracy = evaluate(validation, original_tree) leftRoom, leftPopulation = node.left.getRoomNumber(), node.left.getUsers() rightRoom, rightPopulation = node.right.getRoomNumber(), node.right.getUsers() previous_feature, previous_split, previous_depth, previous_left, previous_right = node.getFeature(), node.getSplit(), node.getDepth(), node.getLeft(), node.getRight() newRoom = -1 newPopulation = leftPopulation + rightPopulation if rightPopulation >= leftPopulation: newRoom = rightRoom else: newRoom = leftRoom node = LeafNode(roomNumber = newRoom, users=newPopulation, depth = previous_depth) new_accuracy = evaluate(validation, original_tree) if new_accuracy < current_accuracy: node = TreeNode(split = previous_split, feature=previous_feature, depth=previous_depth) node.left = previous_left node.right = previous_right return node def evaluate(test_db, trained_tree): num_correct = 0 for data in test_db: sample = data[:-1] prediction = trained_tree.eval(sample) if prediction == data[-1]: num_correct += 1 return num_correct/len(test_db) pruned_tree = pruneTree(tree, validation, tree)
问题根源
- 引用传递失效:函数内直接给
node赋值新节点,只是修改局部变量,未同步到原树的父节点引用,导致剪枝后的节点无法融入树结构。 - 准确率计算错误:剪枝后仍用未修改的
original_tree计算准确率,对比完全无效。 - 节点状态污染:
LeafNode.eval中修改self.users会在评估时改变节点状态,干扰后续剪枝判断。
修改后的代码及逻辑
1. 修正节点类(避免状态污染)
class LeafNode: def __init__(self, roomNumber, users, depth): self.roomNumber = roomNumber self.depth = depth self.users = users def getRoomNumber(self): return self.roomNumber def getDepth(self): return self.depth def getUsers(self): return self.users def eval(self, sample): # 移除评估时的用户数累加,保持节点状态稳定 return self.getRoomNumber()
2. 重构剪枝函数(递归构建新树+局部样本评估)
def pruneTree(node, validation): # 递归处理叶子节点 if isinstance(node, LeafNode): return node # 先剪枝左右子树 left_pruned = pruneTree(node.left, validation) right_pruned = pruneTree(node.right, validation) # 创建当前节点的剪枝后版本(保留内部节点) current_node = TreeNode(node.feature, node.split, node.depth, left_pruned, right_pruned) # 仅当左右子树都是叶子时尝试剪枝 if isinstance(left_pruned, LeafNode) and isinstance(right_pruned, LeafNode): # 获取当前节点覆盖的验证样本(仅评估当前分支的影响) def get_node_samples(target_node, samples): node_samples = [] for sample in samples: current = target_node path_found = False while isinstance(current, TreeNode): if current == target_node: node_samples.append(sample) path_found = True break if sample[current.feature] < current.split: current = current.left else: current = current.right if path_found: continue return node_samples node_samples = get_node_samples(current_node, validation) if not node_samples: # 无样本覆盖,直接剪枝 return LeafNode( roomNumber=right_pruned.getRoomNumber() if right_pruned.getUsers() >= left_pruned.getUsers() else left_pruned.getRoomNumber(), users=right_pruned.getUsers() + left_pruned.getUsers(), depth=current_node.depth ) # 计算剪枝前的正确预测数 correct_original = 0 for sample in node_samples: pred = current_node.eval(sample[:-1]) if pred == sample[-1]: correct_original += 1 # 计算剪枝后的正确预测数 new_room = right_pruned.getRoomNumber() if right_pruned.getUsers() >= left_pruned.getUsers() else left_pruned.getRoomNumber() correct_pruned = sum(1 for sample in node_samples if sample[-1] == new_room) # 剪枝后准确率不下降则保留剪枝结果 if correct_pruned >= correct_original: return LeafNode( roomNumber=new_room, users=right_pruned.getUsers() + left_pruned.getUsers(), depth=current_node.depth ) else: return current_node else: return current_node def evaluate(test_db, trained_tree): num_correct = 0 for data in test_db: sample = data[:-1] prediction = trained_tree.eval(sample) if prediction == data[-1]: num_correct += 1 return num_correct/len(test_db) # 调用方式 pruned_tree = pruneTree(tree, validation) final_accuracy = evaluate(validation, pruned_tree)
修改核心点
- 递归构建新树:剪枝函数返回剪枝后的新节点,父节点通过接收返回值更新自身的
left/right,确保树结构正确更新。 - 局部样本评估:仅针对当前节点覆盖的验证样本对比剪枝前后的正确率,避免修改原树,同时提升效率。
- 状态稳定:移除
LeafNode.eval中的状态修改,保证评估过程中节点数据不变。
内容的提问来源于stack exchange,提问作者isuf
相关产品推荐
相关产品推荐

