如何高效随机选取含根的任意树N节点均匀子树?
均匀选取含根的大小为N的随机子树(高效实现)
问题背景
给定一棵巨型树结构数据集,需要实现一个函数:输入树和数值N,输出包含根节点、大小恰好为N的连通子树,要求所有合法子树被选中的概率均等;同时当N=1000时,需保证运行效率(原树平均每个节点有数十个子节点)。常见算法要么无法保证均匀性,要么内存/时间开销过大。
解决方案(基于论文8.2节逻辑转化)
论文中提出的高效均匀采样方法,核心是动态规划+递归采样,结合在线计算避免预存全部DP值,平衡内存与效率:
1. 关键定义
对任意节点u,定义f(u, m):以u为根的、大小恰好为m的连通子树的数量。该值无需提前全量计算,而是在采样过程中按需计算。
2. 递归采样流程
从根节点出发,构建大小为N的子树(根已占1个位置,剩余需选k=N-1个节点):
- 对当前节点
u,遍历其所有子节点v:- 确定
v能贡献的节点数范围:最小0(不选v的任何子节点),最大min(v的子树总大小, k)(选v子树中最多k个节点)。 - 计算所有可能的节点分配组合的权重(即各子节点选
t_i个节点的f(v, t_i)乘积之和),根据权重随机选择一种分配方案。
- 确定
- 选定分配方案后,对每个子节点
v:若分配的t_i > 0,则递归处理v——需从v的子树中选t_i个节点(v必须包含在内,剩余t_i-1个从v的子节点中选取)。 - 最终收集所有选中的节点,形成结果子树。
3. 效率优化(适配巨型树与N=1000)
- 在线计算
f(u, m):采用背包思想,递归计算当前所需的f(u, m),无需预存所有节点的f值。初始状态f(u,1)=1(仅选u自身),对每个子节点v,通过滚动数组更新背包:new_f[t] += f_old[t-s] * f(v,s)(s为v贡献的节点数)。 - 剪枝处理:若子节点的子树大小小于
s,直接跳过该s值;当计算到t超过N时停止,仅保留到N的结果。 - 内存复用:背包计算使用一维滚动数组,减少内存占用。
代码实现框架(Python)
import random from functools import lru_cache class TreeNode: def __init__(self, val): self.val = val self.children = [] self.subtree_size = 0 # 预处理后存储子树总大小 def precompute_subtree_sizes(root): """遍历树,预处理每个节点的子树总大小""" stack = [(root, False)] while stack: node, processed = stack.pop() if processed: size = 1 for child in node.children: size += child.subtree_size node.subtree_size = size else: stack.append((node, True)) # 逆序入栈保证处理顺序正确 for child in reversed(node.children): stack.append((child, False)) @lru_cache(maxsize=100000) def compute_f(node_id, m): """计算f(u, m):以u为根的大小为m的连通子树数量,用节点ID做缓存键""" node = id_to_node[node_id] if m == 1: return 1 if m > node.subtree_size: return 0 # 一维背包滚动数组 dp = [0] * (m + 1) dp[1] = 1 for child in node.children: child_id = id(child) max_s = min(child.subtree_size, m - 1) # 反向遍历避免重复计算 for t in range(m, 1, -1): for s in range(1, min(max_s, t - 1) + 1): dp[t] += dp[t - s] * compute_f(child_id, s) return dp[m] def sample_subtree(node, target_size, selected): """递归采样大小为target_size的含根连通子树,selected存储选中节点""" selected.add(node) if target_size == 1: return remaining = target_size - 1 node_id = id(node) total = compute_f(node_id, target_size) children = node.children child_ids = [id(c) for c in children] # 动态规划计算每个子节点的可选分配权重 dp = [0] * (remaining + 1) dp[0] = 1 # 存储每个子节点的贡献数组,用于后续回溯分配 child_contribs = [] for c_id in child_ids: child = id_to_node[c_id] max_s = min(child.subtree_size, remaining) contrib = [0] * (remaining + 1) contrib[0] = 1 # 选0个节点的情况 for s in range(1, max_s + 1): contrib[s] = compute_f(c_id, s) child_contribs.append(contrib) # 更新背包 new_dp = [0] * (remaining + 1) for t in range(remaining + 1): if dp[t] == 0: continue for s in range(remaining - t + 1): if contrib[s] == 0: continue new_dp[t + s] += dp[t] * contrib[s] dp = new_dp # 回溯选择每个子节点的分配数量 current_remaining = remaining chosen_sizes = [] for i in reversed(range(len(children))): child = children[i] contrib = child_contribs[i] c_id = child_ids[i] max_possible = min(child.subtree_size, current_remaining) # 计算每个可能s的权重 weights = [] possible_s = [] for s in range(0, max_possible + 1): if current_remaining - s < 0: continue # 计算剩余节点在其他子节点中的组合数 temp_node = TreeNode(None) temp_node.children = children[:i] + children[i+1:] temp_node.subtree_size = node.subtree_size - child.subtree_size temp_id = id(temp_node) id_to_node[temp_id] = temp_node other_count = compute_f(temp_id, current_remaining - s + 1) del id_to_node[temp_id] weight = contrib[s] * other_count if weight > 0: possible_s.append(s) weights.append(weight) # 随机选择s if not possible_s: s_selected = 0 else: s_selected = random.choices(possible_s, weights=weights)[0] chosen_sizes.append(s_selected) current_remaining -= s_selected chosen_sizes.reverse() # 递归处理每个子节点 for i, child in enumerate(children): s = chosen_sizes[i] if s > 0: sample_subtree(child, s, selected) def get_uniform_subtree(root, N): if N < 1: return set() precompute_subtree_sizes(root) if root.subtree_size < N: raise ValueError("Tree size is smaller than target N") # 全局映射:节点ID到节点对象,用于缓存 global id_to_node id_to_node = {id(root): root} for child in root.children: id_to_node[id(child)] = child selected = set() sample_subtree(root, N, selected) # 清理缓存 compute_f.cache_clear() return selected
性能说明
- 时间复杂度:N=1000时,每个节点的背包计算量约为O(N*K)(K为子节点数量,数十个),递归深度取决于树的深度(平衡树为log级,链状树为1000级),整体运行时间可控制在秒级。
- 内存复杂度:主要消耗在
compute_f的LRU缓存(限制为100000条)和递归栈,避免了预存全量DP数组的内存开销,适配巨型树场景。
内容的提问来源于stack exchange,提问作者DrShitgenstein
相关产品推荐
相关产品推荐

