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

如何获取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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.06 11:10:42