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

如何修正AVL平衡搜索树字典,使其仅在叶子节点存储值

基于AVL树的字典:仅在叶子节点存储值的修正方案

学校布置了一项任务:实现一个基于AVL平衡搜索树的字典,要求数据仅存储在叶子节点中。我编写了一个值存储在普通节点中的AVL平衡搜索树字典类,请问如何正确修正该代码,使其仅在叶子节点存储值?以下是我当前的代码:

class Node:
    def __init__(self, key, value=None, left=None, right=None, inv=False):
        if inv:
            left, right = right, left
        self.key = key
        self.value = value
        self.left = left
        self.right = right
        self.exist = 1
        if left is not None and right is not None:
            self.height = 1 + max(self.left.height, self.right.height)
        elif left is not None:
            self.height = 1 + self.left.height
        elif right is not None:
            self.height = 1 + self.right.height
        else:
            self.height = 1


class Dict:
    def __init__(self):
        self.root = None
        self.len = 0

    def __len__(self):
        return self.len

    def __getitem__(self, key):
        if not self.root:
            raise KeyError()
        else:
            return self.get(self.root, key)

    def __contains__(self, key):
        try:
            self.__getitem__(self.root, key)
            return True
        except:
            return False

    def __setitem__(self, key, value):
        if not self.root:
            self.root = Node(key, value)
        else:
            self.root = self.put(self.root, key, value)
        self.len += 1

    def __delitem__(self, key):
        if key not in self:
            raise KeyError()
        else:
            self.delete(self.root, key)
        self.len -= 1

    def delete(self, tree, key):
        if key == tree.key:
            tree.exist = 0
            return
        return self.delete(self.children(tree, key < tree.key)[1], key)

    def height(self, tree): return 0 if tree is None else tree.height

    def children(self, tree, inv): return (tree.right, tree.left) if inv else (tree.left, tree.right)

    def reassoc(self, tree, inv):
        l, r = self.children(tree, inv)
        rl, rr = self.children(r, inv)
        return Node(r.key, r.value, Node(tree.key, tree.value, l, rl, inv), rr, inv)

    def avl(self, tree, inv):
        l, r = self.children(tree, inv)
        if self.height(r) - self.height(l) < 2:
            return tree
        rl, rr = self.children(r, inv)
        if self.height(rl) - self.height(rr) == 1:
            r = self.reassoc(r, not inv)
        return self.reassoc(Node(tree.key, tree.value, l, r, inv), inv)

    def put(self, tree, key, value):
        if tree is None:
            return Node(key, value, None, None)
        if tree.key == key:
            self.len -= 1
            if tree.exist==0:self.len+=1
            return Node(key, value, tree.left, tree.right)
        inv = key < tree.key
        left, right = self.children(tree, inv)
        return self.avl(Node(tree.key, tree.value, left,
                        self.put(right, key, value), inv), inv)

    def get(self, tree, key):
        if tree is None:
            raise KeyError()
        if key == tree.key and tree.exist == 0:
            raise KeyError()
        elif key == tree.key and tree.exist != 0:
            return tree.value
        return self.get(self.children(tree, key < tree.key)[1], key)

def print_tree(tree, indent = 0):
  if tree == None: print()
  print('   '*indent + str(tree.key) + ' -> ' + str(tree.value))
  if tree.left: print_tree(tree.left, indent + 2)
  if tree.right: print_tree(tree.right, indent + 2)
t = Dict()
for v, k in enumerate([5,7,2,1,3,6,2,7]):
  t.__setitem__(k, v)
  print_tree(t.root)
  print()

核心修正要点

要实现仅叶子节点存储值,需明确区分内部节点(仅存key用于导航,必有左右子节点,无有效value)和叶子节点(存key和value,无左右子节点),核心调整如下:

  • Node类扩展:新增is_leaf属性明确节点类型,重新实现高度计算逻辑,叶子节点高度固定为1,内部节点高度基于子节点高度。
  • 插入逻辑:插入时需遍历到叶子节点位置,若目标key已存在则更新叶子值;若不存在则拆分路径上的内部节点,创建新的内部节点和叶子节点,同时维护AVL平衡。
  • 查询逻辑:必须遍历到叶子节点才能返回有效value,内部节点仅用于导航,不存储数据。
  • 删除逻辑:删除叶子节点后,需合并其父节点(若父节点仅剩单个子节点),并重新平衡AVL树,同步更新字典长度。
  • 平衡维护:适配新的节点结构,确保平衡调整时不破坏内部节点与叶子节点的结构规则。

修正后的完整代码

class Node:
    def __init__(self, key, value=None, left=None, right=None, is_leaf=False):
        self.key = key
        self.value = value
        self.left = left
        self.right = right
        self.is_leaf = is_leaf
        # 动态计算节点高度
        self.height = 1 if self.is_leaf else (1 + max(self.left.height, self.right.height) if (self.left and self.right) else 0)

class Dict:
    def __init__(self):
        self.root = None
        self.len = 0

    def __len__(self):
        return self.len

    def __getitem__(self, key):
        if not self.root:
            raise KeyError(key)
        return self.get(self.root, key)

    def __contains__(self, key):
        try:
            self[key]
            return True
        except KeyError:
            return False

    def __setitem__(self, key, value):
        if not self.root:
            self.root = Node(key, value, is_leaf=True)
            self.len = 1
        else:
            prev_len = self.len
            self.root = self.put(self.root, key, value)
            # 仅在新增节点时计数
            if self.len == prev_len:
                self.len += 1

    def __delitem__(self, key):
        if key not in self:
            raise KeyError(key)
        self.root = self.delete(self.root, key)
        self.len -= 1

    def height(self, tree):
        return tree.height if tree is not None else 0

    def update_height(self, tree):
        if not tree.is_leaf:
            tree.height = 1 + max(self.height(tree.left), self.height(tree.right))

    def rotate_right(self, z):
        y = z.left
        T3 = y.right

        y.right = z
        z.left = T3

        self.update_height(z)
        self.update_height(y)
        return y

    def rotate_left(self, z):
        y = z.right
        T2 = y.left

        y.left = z
        z.right = T2

        self.update_height(z)
        self.update_height(y)
        return y

    def get_balance(self, tree):
        return self.height(tree.right) - self.height(tree.left) if tree is not None else 0

    def balance(self, tree):
        if tree is None or tree.is_leaf:
            return tree

        balance = self.get_balance(tree)

        # 左左失衡
        if balance < -1 and self.get_balance(tree.left) <= 0:
            return self.rotate_right(tree)

        # 左右失衡
        if balance < -1 and self.get_balance(tree.left) > 0:
            tree.left = self.rotate_left(tree.left)
            return self.rotate_right(tree)

        # 右右失衡
        if balance > 1 and self.get_balance(tree.right) >= 0:
            return self.rotate_left(tree)

        # 右左失衡
        if balance > 1 and self.get_balance(tree.right) < 0:
            tree.right = self.rotate_right(tree.right)
            return self.rotate_left(tree)

        self.update_height(tree)
        return tree

    def put(self, tree, key, value):
        # 遍历到叶子节点处理插入/更新
        if tree.is_leaf:
            if tree.key == key:
                tree.value = value
                return tree
            # 新建内部节点和叶子节点
            new_leaf = Node(key, value, is_leaf=True)
            if key < tree.key:
                return Node(tree.key, None, new_leaf, tree, is_leaf=False)
            else:
                return Node(tree.key, None, tree, new_leaf, is_leaf=False)
        else:
            if key < tree.key:
                tree.left = self.put(tree.left, key, value)
            elif key > tree.key:
                tree.right = self.put(tree.right, key, value)
            # 内部节点key不存储值,直接返回
            return self.balance(tree)

    def get(self, tree, key):
        if tree.is_leaf:
            if tree.key == key:
                return tree.value
            raise KeyError(key)
        if key < tree.key:
            return self.get(tree.left, key)
        elif key > tree.key:
            return self.get(tree.right, key)
        # 内部节点key不对应有效数据
        raise KeyError(key)

    def delete(self, tree, key):
        if tree is None:
            return None

        if tree.is_leaf:
            return None if tree.key == key else tree

        if key < tree.key:
            tree.left = self.delete(tree.left, key)
        elif key > tree.key:
            tree.right = self.delete(tree.right, key)

        # 合并仅剩单个子节点的内部节点
        if tree.left is None:
            return tree.right
        if tree.right is None:
            return tree.left

        return self.balance(tree)

def print_tree(tree, indent=0):
    if tree is None:
        return
    node_type = "Leaf" if tree.is_leaf else "Internal"
    val_str = str(tree.value) if tree.is_leaf else "None"
    print(f"{'   '*indent}[{node_type}] {tree.key} -> {val_str}")
    if tree.left:
        print_tree(tree.left, indent + 2)
    if tree.right:
        print_tree(tree.right, indent + 2)

# 测试示例
t = Dict()
for v, k in enumerate([5,7,2,1,3,6,2,7]):
    t[k] = v
    print(f"插入 {k}:{v} 后树结构:")
    print_tree(t.root)
    print()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.04 00:50:36