E. A and B and Lecture Rooms 解题思路
核心问题分析
题面包装得很文艺:一所大学是棵 个点、 条边的树。每天 A 在房间 出题、B 在房间 出题,他俩想找一间房子开会,要求这间房到 和到 的距离相等。树上两点距离就是它们之间唯一路径的边数。每次询问 ,问有多少个点到 、 等距。
数据范围 ,每次询问都暴力扫一遍 个点显然要爆。我们得把"到 等距的点集"刻画成一个能 算出大小的几何对象。突破口在于:树上路径是唯一的,所谓"等距"本质上是绕着 路径的中点做文章。
1. 先把三种平凡情形拍死
先把代码开头那几个 continue 解决掉,省得后面纠结边界:
if (x == y) {
cout << n << endl;
continue;
}
int l = lca(x, y);
int dist = depth[x] + depth[y] - 2 * depth[l];
if (dist % 2 != 0) {
cout << 0 << endl;
continue;
}- :到自己距离 、到对方也距离 ,那全树每个点到这两个"同一个点"的距离当然都相等,答案就是 。
- 距离为奇数:树上 到 的距离 。如果 是奇数,那根本不存在"正中间"那个点——任何点到 的距离奇偶性必然一个奇一个偶,永远凑不成相等。直接输出 。
只有 为偶数时才有戏。设 ,路径上从 走 步落到的那个点,就是这条路径的中点 。所有等距点必然到 别急,我们慢慢看。
2. 倍增三件套:LCA、第 k 步、找中点
要快速做这些事,先把倍增祖先表 fa[u][i] 建好( 表示 往上跳 步的祖先),同时记下每个点的深度 depth。LOG = 20 足够覆盖 。
有了倍增表,get_k(u, v, k) 就能在 内求出 到 路径上、从 数第 个点:先求 LCA,判断第 步落在 这半边还是 那半边,然后按 的二进制位往上跳:
auto get_k = [&](int u, int v, int k) -> int {
int l = lca(u, v);
int dist_u_to_lca = depth[u] - depth[l];
int total_dist = depth[u] + depth[v] - 2 * depth[l];
if (k > total_dist) return 0;
if (k <= dist_u_to_lca) {
int curr = u;
for (int i = LOG - 1; i >= 0; --i)
if ((k >> i) & 1) curr = fa[curr][i];
return curr;
} else {
int steps_from_v = total_dist - k;
...
}
};如果 还没跨过 LCA( 到 LCA 的距离),就直接从 往上跳 步;否则换个角度——从 那头往回数 步,等价地跳过去。这个小技巧避免了"从 LCA 往下走"的麻烦(树上往下走没有唯一方向,但往上跳是确定的)。
中点就是 get_k(x, y, half),记作 。
3. 核心几何观察:等距点都挂在中点 M 上
来到本题的灵魂。设中点为 ,它两侧路径上各有一个贴着 的邻居:靠 那侧的记 ,靠 那侧的记 。代码里用 get_ans 一把抓出来:
auto get_ans = [&](int u, int v, int dist) -> pair<int, int> {
if (dist <= 0) return {0, 0};
int prev_u = get_k(u, v, dist - 1); // 离中点最近、偏 u 一侧的点
int prev_v = get_k(v, u, dist - 1); // 离中点最近、偏 v 一侧的点
return {prev_u, prev_v};
};get_k(x, y, half - 1) 是从 走 步的点——也就是中点 在 这一侧的紧邻;同理另一边。
现在断言:一个点 到 等距,当且仅当 到 的路径都恰好在 处汇合,再说直白点, 必须"挂"在 上而不能跑进 或 所代表的那两个分支里。
为啥?任取一点 ,它到 、到 的两条路径,从 出发一定会在某处并入 – 主路径,记交汇点为 。于是 、。两者相等 。也就是说,所有等距点,到主路径的接入口必须正好是中点 。
4. 子树相减:把答案数出来
知道了"等距点 = 接入口为 的点",剩下就是计数。这分两种情况,恰好对应代码最后那个 if:
情形一:( 就是 LCA)。
此时中点 正是 的最近公共祖先。以 为根,整棵树有 个点;要排除的是那些"接入口不是 "的点——也就是会先拐进 子树或 子树的点。这两个子树里的点,到 、 的距离一定不等(它们偏向了某一边)。所以答案是
if (depth[x] == depth[y]) {
cout << n - sons[no_x] - sons[no_y] << endl;
}这里 sons[u] 是以 为根时 的子树大小,提前一遍 DFS 求好。 是 的两个孩子方向,它们的子树两两不交、也不含 本身,所以减一减就把"会偏向 一侧"和"会偏向 一侧"的点干净地剔除了,剩下的全是合法等距点(包括 自己,以及挂在 其它分支上的点)。
情形二:( 不是 LCA,深的那侧在 子树里)。
代码先 swap 保证 是更深的那个:
if (depth[x] < depth[y]) {
swap(x, y);
swap(no_x, no_y);
}
cout << sons[now] - sons[no_x] << endl;注意 now 就是中点 (代码里 int now = get_k(x, y, half);)。既然 更深,那中点 一定落在 到 LCA 这段上,是 LCA 的某个真后代。这时合法的等距点只可能藏在 的子树内部——因为 子树外的点接入主路径时根本到不了 (会从 上方接入,偏向 )。而 子树里,又要扣掉偏向 的那一支,也就是 的子树。于是
这里 是 在 一侧的孩子,它的整棵子树都更靠近 ,必须排除;而 一侧根本不在 子树内,天然就不用管,所以这次只减一项。
5. 复盘一下整套流程
把上面拼起来,每次询问的逻辑就是:
- 求 LCA、算距离 ; 输出 , 为奇输出 。
- ,倍增找中点 和它两侧紧邻 。
- 同深度 ;否则让 取深的一侧,输出 。
预处理倍增表与子树大小是 ,每次询问几次倍增跳跃都是 ,总复杂度
对 的规模轻轻松松,2 秒时限里随便跑。一道看着吓人、其实把"中点 + 子树相减"想通就秒掉的好题。
CPP 代码实现
// E. A and B and Lecture Rooms
#include <bits/stdc++.h>
#define lg(x) (63 - __builtin_clzll(x))
#define all(x) (x).begin(), (x).end()
#define low_bit(x) ((x) & (-x))
#define pb push_back
#define db long double
#define int long long
#define sz(x) (int)x.size()
#define endl "\n"
using namespace std;
void solve() {
int n;
cin >> n;
vector<vector<int>> g(n + 1);
for (int i = 1; i < n; i++) {
int u, v;
cin >> u >> v;
g[u].pb(v);
g[v].pb(u);
}
const int LOG = 20;
vector<int> depth(n + 1, 0);
vector<vector<int>> fa(n + 1, vector<int>(LOG, 0));
auto dfs = [&] (auto&& self, int u, int p) -> void {
fa[u][0] = p;
for (int i = 1; i < LOG; ++i) {
fa[u][i] = fa[fa[u][i - 1]][i - 1];
}
for (auto it : g[u]) {
if (it == p) continue;
depth[it] = depth[u] + 1;
self(self, it, u);
}
};
dfs(dfs, 1, 0);
auto lca = [&](int u, int v) -> int {
if (depth[u] < depth[v]) swap(u, v);
for (int i = LOG - 1; i >= 0; --i) {
if (depth[fa[u][i]] >= depth[v]) {
u = fa[u][i];
}
}
if (u == v) return u;
for (int i = LOG - 1; i >= 0; --i) {
if (fa[u][i] != fa[v][i]) {
u = fa[u][i];
v = fa[v][i];
}
}
return fa[u][0];
};
auto get_k = [&](int u, int v, int k) -> int {
int l = lca(u, v);
int dist_u_to_lca = depth[u] - depth[l];
int total_dist = depth[u] + depth[v] - 2 * depth[l];
if (k > total_dist) return 0;
if (k <= dist_u_to_lca) {
int curr = u;
for (int i = LOG - 1; i >= 0; --i) {
if ((k >> i) & 1) curr = fa[curr][i];
}
return curr;
} else {
int steps_from_v = total_dist - k;
int curr = v;
for (int i = LOG - 1; i >= 0; --i) {
if ((steps_from_v >> i) & 1) curr = fa[curr][i];
}
return curr;
}
};
auto get_ans = [&](int u, int v, int dist) -> pair<int, int> {
if (dist <= 0) return {0, 0};
int prev_u = get_k(u, v, dist - 1);
int prev_v = get_k(v, u, dist - 1);
return {prev_u, prev_v};
};
vector<int> sons(n + 1, 0);
auto get_son = [&] (auto&& self, int u, int p) -> void {
int sum = 1;
for (auto it : g[u]) {
if (it == p) continue;
self(self, it, u);
sum += sons[it];
}
sons[u] = sum;
};
get_son(get_son, 1, 0);
int q;
cin >> q;
for (int i = 1; i <= q; i++) {
int x, y;
cin >> x >> y;
if (x == y) {
cout << n << endl;
continue;
}
int l = lca(x, y);
int dist = depth[x] + depth[y] - 2 * depth[l];
if (dist % 2 != 0) {
cout << 0 << endl;
continue;
}
int half = dist / 2;
auto temp = get_ans(x, y, half);
int no_x = temp.first;
int no_y = temp.second;
int now = get_k(x, y, half); // 中点
if (depth[x] == depth[y]) {
cout << n - sons[no_x] - sons[no_y] << endl;
} else {
if (depth[x] < depth[y]) {
swap(x, y);
swap(no_x, no_y);
}
cout << sons[now] - sons[no_x] << endl;
}
}
}
signed main() {
ios_base::sync_with_stdio(false);
cin.tie(nullptr);
int t = 1;
// cin >> t;
while (t--) {
solve();
}
}