超详细讲解,
2026-07-18 19:22:27
发布于:云南
超详细讲解,非AI,教师讲解。
一、 题意解析与思路排除
这道题给出了一棵 n 个节点的无向带权树,有 m 次询问,每次询问任意两点之间的最短距离。
有些人可能会想到用 Floyd 算法或者对每次询问做 BFS、Dijkstra。但需要数据范围:n 最大为 1 万,m 最大为 2 万。如果每次询问都遍历全图,复杂度会达到 O(n*m),在极限数据下会严重超时。
解题的关键突破口在于:题目保证输入是一棵树,有 n 个点和 n-1 条边。在树结构中,两点之间的路径是唯一的。因此,所谓的最短距离,就是这条唯一路径上的边权总和。我们需要找到一种高效的预处理方法,来快速求出任意两点之间的路径长度。
二、 核心算法原理:LCA 与前缀距离
解决该问题的标准算法是最近公共祖先,简称 LCA,配合根节点前缀距离。
首先,选定任意一个节点作为根节点,比如节点 1。我们定义两个数组:
depth[i] 表示节点 i 的深度,根节点的深度为 0。
dist[i] 表示从根节点到节点 i 的路径边权总和。
对于一次查询 (x, y),设它们的最近公共祖先为 l,即 l = LCA(x, y)。
因为 x 到 l 的路径和 y 到 l 的路径在 l 处汇合,所以 x 到 y 的最短距离等于 dist[x] 加上 dist[y] 减去两倍的 dist[l]。
这是因为 dist[x] 和 dist[y] 都包含了根节点到 l 的那一段距离,重复计算了两次,所以需要减掉两次。
关于 LCA 的求法,我们采用二进制倍增法。我们需要预处理每个节点向上跳 2 的 j 次方步到达的祖先节点。查询时,先将较深的节点向上跳跃,与较浅的节点对齐到同一深度。然后,如果两个节点不相同,就一起向上跳跃,直到它们的父节点相同,这个父节点就是最近公共祖先。
三、 关键变量说明
这里说明一下代码中关键变量的作用:
adj[u]:邻接表,存储与 u 相连的节点以及对应的边权。
depth[i]:节点 i 的深度,用于判断跳跃的步数。
dist[i]:根节点到节点 i 的累计距离。
up[i][j]:节点 i 向上跳 2^j 步到达的祖先节点。
LOG:最大倍增层数。由于 n 最大为 1 万,2 的 14 次方等于 16384,所以取 LOG 为 17 就足够安全了。
四、 算法流程与复杂度分析
预处理阶段:通过一次深度优先搜索从根节点 1 出发遍历整棵树。在遍历过程中,我们计算出每个节点的深度和前缀距离,同时利用递推公式 up[v][j] = up[ up[v][j-1] ][j-1] 填充倍增表。这一步的时间复杂度是 O(n log n)。
查询阶段:每次询问执行 LCA 函数,内部的循环次数与 log n 成正比。总时间复杂度为 O((n + m) log n),空间复杂度为 O(n log n),完全可以满足题目 1 秒的时间限制。
五、 代码(%100AC)
#include <bits/stdc++,h>
using namespace std;
const int MAXN = 10005;
const int LOG = 17;
vector<pair<int, int>> adj[MAXN];
int depth[MAXN];
int dist[MAXN];
int up[MAXN][LOG];
void dfs(int u, int p) {
for (auto &edge : adj[u]) {
int v = edge.first;
int w = edge.second;
if (v == p) continue;
depth[v] = depth[u] + 1;
dist[v] = dist[u] + w;
up[v][0] = u;
for (int j = 1; j < LOG; j++) {
up[v][j] = up[ up[v][j-1] ][j-1];
}
dfs(v, u);
}
}
int lca(int u, int v) {
if (depth[u] < depth[v]) swap(u, v);
int diff = depth[u] - depth[v];
for (int j = 0; j < LOG; j++) {
if (diff & (1 << j)) {
u = up[u][j];
}
}
if (u == v) return u;
for (int j = LOG - 1; j >= 0; j--) {
if (up[u][j] != up[v][j]) {
u = up[u][j];
v = up[v][j];
}
}
return up[u][0];
}
int main() {
ios::sync_with_stdio(false);
cin.tie(nullptr);
int n, m;
cin >> n >> m;
for (int i = 1; i <= n - 1; i++) {
int x, y, k
cin >> x >> y >> k;
adj[x].push_back({y, k});
adj[y].push_back({x, k});
}
depth[1] = 0;
dist[1] = 0;
up[1][0] = 1;
for (int j = 1; j < LOG:j++) {
up[1][j] = 1;
dfs(1, 1);
while (m--) {
int x, y;
cin >> x >> y;
int z = lca(x, y)
cout << dist[x] + dist[y] - 2 * dist[z] << '\n';
}
return 0
}//已开启放抄袭
这里空空如也








有帮助,赞一个