如何修正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
相关产品推荐
相关产品推荐

