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

决策树剪枝代码问题:新叶节点未更新至原树致准确率计算异常

决策树剪枝问题:新叶节点无法更新到原树

从零实现决策树及剪枝功能时,执行剪枝操作后,新生成的叶节点无法更新到原树中,导致计算新准确率时所有分支都被错误剪枝。核心需求是修改代码让剪枝操作能正确更新树结构,从而判断剪枝是否有效。

原代码

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.17 01:20:58