https://loj.ac/s/1844505 https://www.luogu.com.cn/record/117743528
#include <bits/stdc++.h>
using namespace std;
const int MAXN = 2e5 + 10;
int n, k, ans = INT_MAX;
vector<int> G[MAXN], c[MAXN];
int col[MAXN], siz[MAXN], fa[MAXN];
bool used[MAXN], vis[MAXN], vis2[MAXN], vis3[MAXN];
void init(int u, int fa) {
siz[u] = 1;
vis2[u] = true;
::fa[u] = fa;
for (int v : G[u]) {
if (used[v] || v == fa) continue;
init(v, u);
siz[u] += siz[v];
}
}
int mx, _;
void get_rt(int u, int siz) {
int ___ = 0;
vis[u] = true;
for (int v : G[u]) {
if (vis[v] || used[v]) continue;
get_rt(v, siz);
___ = max(___, ::siz[v]);
}
___ = max(___, siz - ::siz[u]);
if (___ < mx) {
_ = u;
mx = ___;
}
}
int cnt;
queue<int> q;
bool chk(int w) {
vis3[w] = true;
for (int x : c[w]) {
if (!vis2[x]) return false;
q.push(x);
}
cnt++;
return true;
}
int get(int p) {
cnt = 0;
while (q.size()) q.pop();
if (!chk(col[p])) return INT_MAX;
while (q.size()) {
int u = q.front();
q.pop();
if (!vis3[col[fa[u]]] && !used[fa[u]] && fa[u]) {
if (!chk(col[fa[u]])) return INT_MAX;
}
}
return cnt - 1;
}
void Solve(int rt) {
init(rt, 0);
used[rt] = true;
ans = min(ans, get(rt));
memset(vis, false, sizeof vis);
memset(vis2, false, sizeof vis2);
memset(vis3, false, sizeof vis3);
for (int v : G[rt]) {
if (used[v]) continue;
mx = INT_MAX;
get_rt(v, siz[v]);
Solve(_);
}
}
int main() {
ios::sync_with_stdio(false); cin.tie(0); cout.tie(0);
cin >> n >> k;
for (int i = 1; i < n; i++) {
static int u, v;
cin >> u >> v;
G[u].push_back(v);
G[v].push_back(u);
}
for (int i = 1; i <= n; i++) {
cin >> col[i];
c[col[i]].push_back(i);
}
init(1, 0);
mx = INT_MAX;
get_rt(1, siz[1]);
Solve(_);
cout << ans << endl;
return 0;
}