题目背景
对应的选择、判断题:试题 - GESP 202503 C++ 八级 - 洛谷有题
题目描述
小杨有一棵包含 n 个节点的树,其中节点的编号从 1 到 n。
小杨设置了 a 个好点对 {⟨u1,v1⟩,⟨u2,v2⟩,...,⟨ua,va⟩} 和一个坏点对 ⟨bu,bv⟩。一个节点能被删除,当且仅当:
- 删除该节点后对于所有的 1≤i≤a,好点对 ui 和 vi 仍然连通;
- 删除该节点后坏点对 bu 和 bv 不连通。
如果点对中的任意一个节点被删除,其视为不连通。
小杨想知道,还有多少个节点能被删除。
输入格式
第一行包含两个非负整数 n, a,含义如下题面所示。
接下来 n−1 行,每行包含两个正整数 xi,yi,代表存在一条连接节点 xi 和 yi 的边;
之后 a 行,每行包含两个正整数 ui,vi,代表一个好点对 ⟨ui,vi⟩;
最后一行包含两个正整数 bu,bv,代表坏点对 ⟨bu,bv⟩。
输出格式
输出一个非负整数,代表删除的节点个数。
输入输出样例
输入 #1复制
6 2
1 3
1 5
3 6
3 2
5 4
5 4
5 3
2 6
输出 #1复制
2
说明/提示
| 子任务编号 | 分值 | n | a |
|---|---|---|---|
| 1 | 20 | =10 | =0 |
| 2 | 20 | ≤100 | ≤100 |
| 3 | 60 | ≤106 | ≤105 |
对于全部数据,保证有 1≤n≤106, 0≤a≤105, ui=vi, bu=bv。
代码实现:
cpp
#include <iostream>
#include <vector>
#include <algorithm>
using namespace std;
const int MAXN = 1e6 + 10;
const int LOG = 20;
vector<int> g[MAXN];
int dep[MAXN], fa[LOG][MAXN];
long long diff[MAXN];
void dfs(int u, int father)
{
fa[0][u] = father;
dep[u] = dep[father] + 1;
for (int v : g[u])
{
if (v != father)
dfs(v, u);
}
}
int lca(int u, int v)
{
if (dep[u] < dep[v]) swap(u, v);
for (int k = LOG - 1; k >= 0; k--)
{
if (dep[u] - (1 << k) >= dep[v])
u = fa[k][u];
}
if (u == v) return u;
for (int k = LOG - 1; k >= 0; k--)
{
if (fa[k][u] != fa[k][v])
{
u = fa[k][u];
v = fa[k][v];
}
}
return fa[0][u];
}
void add_path(int u, int v)
{
int l = lca(u, v);
diff[u]++;
diff[v]++;
diff[l]--;
if (fa[0][l]) diff[fa[0][l]]--;
}
void get_cnt(int u, int father)
{
for (int v : g[u])
{
if (v != father)
{
get_cnt(v, u);
diff[u] += diff[v];
}
}
}
void get_road(int u, int v, vector<int>& path)
{
int l = lca(u, v);
while (u != l)
{
path.push_back(u);
u = fa[0][u];
}
path.push_back(l);
vector<int> tmp;
while (v != l)
{
tmp.push_back(v);
v = fa[0][v];
}
reverse(tmp.begin(), tmp.end());
for (int x : tmp) path.push_back(x);
}
int main()
{
ios::sync_with_stdio(false);
cin.tie(nullptr);
int n, a;
cin >> n >> a;
for (int i = 1; i <= n - 1; i++)
{
int x, y;
cin >> x >> y;
g[x].push_back(y);
g[y].push_back(x);
}
dep[0] = 0;
dfs(1, 0);
for (int k = 1; k < LOG; k++)
{
for (int i = 1; i <= n; i++)
{
fa[k][i] = fa[k-1][fa[k-1][i]];
}
}
for (int i = 1; i <= a; i++)
{
int u, v;
cin >> u >> v;
add_path(u, v);
}
get_cnt(1,0);
int bu, bv;
cin >> bu >> bv;
vector<int> bad_path;
get_road(bu, bv, bad_path);
int ans = 0;
for (int x : bad_path)
{
if (diff[x] == 0) ans++;
}
cout << ans << endl;
return 0;
}