如何获取BeamTree叶子节点引用以删除低累积和节点?
BeamTree剪枝:保留前N个高累积和叶子节点
问题描述
我实现了一个BeamTree树形结构,需要计算每个叶子节点的累积和,最终保留累积和排名前N的叶子节点。目前已成功计算累积和并获取叶子节点内容,但无法通过引用删除低累积和节点,请问有解决方法吗?
原实现代码
import numpy as np class BeamTree(dict): def __init__(self, *args, **kwargs): super(BeamTree, self).__init__(*args, **kwargs) self.__dict__ = self self.nodes = [] def add(self, a, aa, q): self.nodes.append( BeamTree({ 'a': a, 'aa': aa, 'qq': q, 'qs':0 }) ) return self.nodes[-1] def qsum(self, q=0): if len(self.nodes) == 0 : return [] leafs = [] for node in self.nodes: node['qs'] = q + node['qq'] leafs.extend( node.qsum(node['qs']) ) if len(node.nodes) == 0 : leafs.append(node) if len(leafs) > 0 : return leafs return [] def generate(self, branch=3, depth=3): if depth < 1 : return for b in range(branch) : sym = 's' + str(np.random.randint(100)) aix = np.random.randint(100) q = np.random.rand() node = self.add(sym, aix, q) node.generate(branch, depth-1)
测试示例
In [212]: b=BeamTree(); b.generate(2,2) In [213]: l=b.qsum(0) In [214]: b Out[214]: {'nodes': [{'a': 's80', 'aa': 56, 'qq': 0.673, 'qs': 0.673, 'nodes': [{'a': 's8', 'aa': 16, 'qq': 0.115, 'qs': 0.788, 'nodes': []}, {'a': 's64', 'aa': 10, 'qq': 0.599, 'qs': 1.272, 'nodes': []}]}, {'a': 's67', 'aa': 0, 'qq': 0.900, 'qs': 0.900, 'nodes': [{'a': 's69', 'aa': 23, 'qq': 0.801, 'qs': 1.700, 'nodes': []}, {'a': 's8', 'aa': 41, 'qq': 0.826, 'qs': 1.726, 'nodes': []}]}]} In [215]: l Out[215]: [{'a': 's8', 'aa': 16, 'qq': 0.115, 'qs': 0.788, 'nodes': []}, {'a': 's64', 'aa': 10, 'qq': 0.599, 'qs': 1.272, 'nodes': []}, {'a': 's69', 'aa': 23, 'qq': 0.801, 'qs': 1.700, 'nodes': []}, {'a': 's8', 'aa': 41, 'qq': 0.826, 'qs': 1.726, 'nodes': []}] In [216]: del l[0] In [217]: l Out[217]: [{'a': 's64', 'aa': 10, 'qq': 0.599, 'qs': 1.272, 'nodes': []}, {'a': 's69', 'aa': 23, 'qq': 0.801, 'qs': 1.700, 'nodes': []}, {'a': 's8', 'aa': 41, 'qq': 0.826, 'qs': 1.726, 'nodes': []}] In [218]: b Out[218]: {'nodes': [{'a': 's80', 'aa': 56, 'qq': 0.673, 'qs': 0.673, 'nodes': [{'a': 's8', 'aa': 16, 'qq': 0.115, 'qs': 0.788, 'nodes': []}, {'a': 's64', 'aa': 10, 'qq': 0.599, 'qs': 1.272, 'nodes': []}]}, {'a': 's67', 'aa': 0, 'qq': 0.900, 'qs': 0.900, 'nodes': [{'a': 's69', 'aa': 23, 'qq': 0.801, 'qs': 1.700, 'nodes': []}, {'a': 's8', 'aa': 41, 'qq': 0.826, 'qs': 1.726, 'nodes': []}]}]}
解决方案
问题核心:删除叶子列表中的元素仅修改列表本身,不会影响原树的节点结构。要真正剪枝,需要找到叶子的父节点,从父节点的nodes列表中移除该叶子,同时清理无有效子节点的中间节点。
1. 给节点添加父节点引用
修改BeamTree的__init__和add方法,让每个节点能追溯到父节点:
def __init__(self, *args, **kwargs): super(BeamTree, self).__init__(*args, **kwargs) self.__dict__ = self self.nodes = [] self.parent = None # 根节点父节点为None def add(self, a, aa, q): node = BeamTree({ 'a': a, 'aa': aa, 'qq': q, 'qs':0 }) node.parent = self # 绑定父节点 self.nodes.append(node) return node
2. 筛选前N个目标叶子节点
按累积和qs降序排序,取前N个并转为集合方便判断:
N = 3 # 示例保留前3个 sorted_leafs = sorted(leafs, key=lambda x: x['qs'], reverse=True) top_leafs = set(sorted_leafs[:N])
3. 实现剪枝方法
给BeamTree类添加递归剪枝方法,删除非目标叶子及空父节点:
def prune_non_top_leaves(self, top_leafs): to_remove = [] # 先递归处理子节点 for node in self.nodes: if len(node.nodes) > 0: node.prune_non_top_leaves(top_leafs) # 如果是叶子且不在目标集合,标记删除 if len(node.nodes) == 0 and node not in top_leafs: to_remove.append(node) # 移除标记的节点 for node in to_remove: self.nodes.remove(node) # 若当前节点无剩余子节点且不是根节点,通知父节点删除自己 if len(self.nodes) == 0 and self.parent is not None: self.parent.nodes.remove(self)
完整修改后的代码
import numpy as np class BeamTree(dict): def __init__(self, *args, **kwargs): super(BeamTree, self).__init__(*args, **kwargs) self.__dict__ = self self.nodes = [] self.parent = None def add(self, a, aa, q): node = BeamTree({ 'a': a, 'aa': aa, 'qq': q, 'qs':0 }) node.parent = self self.nodes.append(node) return node def qsum(self, q=0): if len(self.nodes) == 0 : return [] leafs = [] for node in self.nodes: node['qs'] = q + node['qq'] leafs.extend( node.qsum(node['qs']) ) if len(node.nodes) == 0 : leafs.append(node) if len(leafs) > 0 : return leafs return [] def generate(self, branch=3, depth=3): if depth < 1 : return for b in range(branch) : sym = 's' + str(np.random.randint(100)) aix = np.random.randint(100) q = np.random.rand() node = self.add(sym, aix, q) node.generate(branch, depth-1) def prune_non_top_leaves(self, top_leafs): to_remove = [] for node in self.nodes: if len(node.nodes) > 0: node.prune_non_top_leaves(top_leafs) if len(node.nodes) == 0 and node not in top_leafs: to_remove.append(node) for node in to_remove: self.nodes.remove(node) if len(self.nodes) == 0 and self.parent is not None: self.parent.nodes.remove(self)
测试剪枝效果
b = BeamTree() b.generate(2,2) leafs = b.qsum(0) # 保留前3个最高累积和的叶子 N = 3 sorted_leafs = sorted(leafs, key=lambda x: x['qs'], reverse=True) top_leafs = set(sorted_leafs[:N]) # 执行剪枝 b.prune_non_top_leaves(top_leafs) # 查看剪枝后的树结构 print(b)
执行后,原树将只保留目标叶子节点及其完整路径上的祖先节点,低累积和的叶子及其无有效子节点的父节点会被移除。
内容的提问来源于stack exchange,提问作者sten
相关产品推荐
相关产品推荐

