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

基于最大边权的树中点对多查询问题求解(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)数量。

输入输出格式

输入

  1. 第一行两个整数n和m,分别是顶点数和查询数(1 ≤ n, m ≤ 100,000)。
  2. 接下来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)
  3. 最后一行包含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)结合边排序和查询排序的方法,大致思路是:

  1. 按边权从小到大排序所有边。
  2. 按查询的q_i从小到大排序所有查询(同时记录原始索引,方便最后还原结果顺序)。
  3. 用并查集逐步合并边权不超过当前q_i的边,维护连通分量的大小。
  4. 通过连通分量的大小计算满足条件的点对数量:每次合并两个大小为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,提问作者讗讜讛讚 讙讜诇讚讘专讙

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.14 17:03:12