基于最大边权的树中点对多查询问题求解(Union-Find实现报错求助)
基于最大边权的树中点对多查询问题求解(Union-Find实现报错求助)
我现在在处理带权树的多查询问题,需要高效回答每个查询:找出所有点对(u, v)(u < v),使得u到v的简单路径上的最大边权不超过给定的q_i。需要按顺序输出所有查询的结果。
定义
Tree
- 连通无环图,包含n个顶点,n-1条边。
- 每条边带有权重w_i,连接两个顶点u_i和v_i。
Simple path
- 两个顶点之间不重复经过任何顶点的路径。
Query
- 给定整数q_i,统计满足u < v且u到v的简单路径上的最大边权 ≤ q_i的点对(u, v)数量。
输入输出格式
输入
- 第一行两个整数n和m,分别是顶点数和查询数(1 ≤ n, m ≤ 100,000)。
- 接下来n-1行,每行三个整数w_i, u_i, v_i,表示一条边:
- w_i:边的权重(1 ≤ w_i ≤ 200,000)
- u_i, v_i:边连接的两个顶点(1 ≤ u_i, v_i ≤ n)
- 最后一行包含m个整数q_1, q_2, ..., q_m,每个q_i对应一个查询的阈值(1 ≤ q_i ≤ 200,000)。
输出
一行包含m个整数,依次对应每个查询的结果:即满足条件的点对(u, v)(u < v)的数量。
示例
示例1
输入:
7 5 1 2 1 3 2 3 2 4 1 4 5 2 5 7 4 3 6 2 5 2 3 4 1
输出:
21 7 15 21 3
示例2
输入:
1 2 1 2
输出:
0 0
示例3
输入:
3 3 1 2 1 2 3 2 1 3 2
输出:
1 3 3
我的解决思路与代码
我尝试用并查集(Disjoint Set Union, DSU)结合边排序和查询排序的方法,大致思路是:
- 按边权从小到大排序所有边。
- 按查询的q_i从小到大排序所有查询(同时记录原始索引,方便最后还原结果顺序)。
- 用并查集逐步合并边权不超过当前q_i的边,维护连通分量的大小。
- 通过连通分量的大小计算满足条件的点对数量:每次合并两个大小为a和b的分量,新增的点对数量是a*b,累加得到当前总点对数。
但是我写的Python代码在学校的测试平台上无法通过部分测试用例,以下是我的代码:
from collections import defaultdict class UnionFind: def __init__(self, size): self.parent = list(range(size)) self.size = [1] * size def find(self, x): if self.parent[x] != x: self.parent[x] = self.find(self.parent[x]) return self.parent[x] def union(self, x, y): root_x = self.find(x) root_y = self.find(y) if root_x != root_y: if self.size[root_x] < self.size[root_y]: root_x, root_y = root_y, root_x self.parent[root_y] = root_x self.size[root_x] += self.size[root_y] return self.size[root_x] * self.size[root_y] return 0 def count_pairs(n, m, edges, queries): edges.sort() queries = [(q, i) for i, q in enumerate(queries)] queries.sort() uf = UnionFind(n + 1) results = [0] * m current_pairs = 0 edge_index = 0 for q, idx in queries: while edge_index < len(edges) and edges[edge_index][0] <= q: w, u, v = edges[edge_index] current_pairs += uf.union(u, v) edge_index += 1 results[idx] = current_pairs return results import sys input = sys.stdin.read data = input().splitlines() n, m = map(int, data[0].split()) edges = [tuple(map(int, line.split())) for line in data[1:n]] queries = list(map(int, data[n].split())) print(*count_pairs(n, m, edges, queries))
希望大家能帮我看看代码哪里有问题,为什么过不了测试用例?谢谢大家!
备注:内容来源于stack exchange,提问作者讗讜讛讚 讙讜诇讚讘专讙
相关产品推荐
相关产品推荐

