#include <bits/stdc++.h>
using namespace std;
const int N = 100005;
int n, m, u, v, ecnt, dfn, res, a[N], cnt[N], pos[N], l[N], r[N], ans[N];
vector <int> G[N];
void dfs(int u, int fa) {
l[u] = ++dfn;
for (int i = 0; i < G[u].size(); ++i) {
if (G[u][i] != fa) dfs(G[u][i], u);
}
r[u] = dfn;
}
struct node {
int l, r, id;
} q[N];
bool cmp(node a, node b) { return (pos[a.l] < pos[b.l]) || (pos[a.l] == pos[b.l] && a.r < b.r); }
void add(int x) { ++cnt[a[x]]; if (cnt[a[x]] == 1) ++res; }
void del(int x) { --cnt[a[x]]; if (!cnt[a[x]]) --res; }
int main() {
cin >> n;
int len = sqrt(n);
for (int i = 1; i < n; ++i) {
cin >> u >> v;
G[u].push_back(v);
G[v].push_back(u);
}
for (int i = 1; i <= n; ++i) cin >> a[i];
dfs(1, 0);
cin >> m;
for (int i = 1; i <= n; ++i) pos[i] = i / len;
while (m--) {
int x;
cin >> x;
++ecnt;
q[ecnt].l = l[x], q[ecnt].r = r[x], q[ecnt].id = ecnt;
}
sort(q + 1, q + 1 + ecnt, cmp);
int l = 1, r = 0;
for (int i = 1; i <= ecnt; ++i) {
while (l > q[i].l) add(--l);
while (r < q[i].r) add(++r);
while (l < q[i].l) del(l++);
while (r > q[i].r) del(r--);
ans[q[i].id] = res;
}
for (int i = 1; i <= ecnt; ++i) printf("%d\n", ans[i]);
return 0;
}
题目