代码如下:
#include<bits/stdc++.h>
using namespace std;
vector<vector<int>> adj;
vector<int> depth;
vector<vector<int>> parent;
int LOG;
void dfs(int u, int p) {
parent[0][u] = p;
for (int v : adj[u]) {
if (v != p) {
depth[v] = depth[u] + 1;
dfs(v, u);
}
}
}
void preprocess(int n) {
LOG = log2(n) + 1;
parent.resize(LOG, vector<int>(n + 1));
depth.resize(n + 1);
dfs(1, -1);
for (int k = 1; k < LOG; ++k) {
for (int v = 1; v <= n; ++v) {
if (parent[k - 1][v] != -1) {
parent[k][v] = parent[k - 1][parent[k - 1][v]];
} else {
parent[k][v] = -1;
}
}
}
}
int lca(int u, int v) {
if (depth[u] < depth[v]) swap(u, v);
for (int k = LOG - 1; k >= 0; --k) {
if (depth[u] - (1 << k) >= depth[v]) {
u = parent[k][u];
}
}
if (u == v) return u;
for (int k = LOG - 1; k >= 0; --k) {
if (parent[k][u] != -1 && parent[k][u] != parent[k][v]) {
u = parent[k][u];
v = parent[k][v];
}
}
return parent[0][u];
}
int distance(int u, int v) {
return depth[u] + depth[v] - 2 * depth[lca(u, v)];
}
int main() {
int n, m;
cin >> n >> m;
adj.resize(n + 1);
for (int i = 0; i < n - 1; ++i) {
int x, y;
cin >> x >> y;
adj[x].push_back(y);
adj[y].push_back(x);
}
preprocess(n);
vector<int> c(m);
for (int i = 0; i < m; ++i) {
cin >> c[i];
}
set<int> S;
vector<int> entry(n + 1), exit_(n + 1);
int time = 0;
vector<int> stack;
stack.push_back(1);
vector<bool> visited(n + 1);
visited[1] = true;
while (!stack.empty()) {
int u = stack.back();
stack.pop_back();
entry[u] = ++time;
for (int v : adj[u]) {
if (!visited[v] && v != parent[0][u]) {
visited[v] = true;
stack.push_back(v);
}
}
}
auto cmp = [&](int a, int b) {
return entry[a] < entry[b];
};
set<int, decltype(cmp)> nodes_in_order(cmp);
int total_edges = 0;
for (int i = 0; i < m; ++i) {
int ci = c[i];
if (S.find(ci) == S.end()) {
S.insert(ci);
if (nodes_in_order.empty()) {
nodes_in_order.insert(ci);
total_edges = depth[ci];
} else {
auto it = nodes_in_order.lower_bound(ci);
int pred = -1, succ = -1;
if (it != nodes_in_order.end()) {
succ = *it;
}
if (it != nodes_in_order.begin()) {
--it;
pred = *it;
}
int lca_pred = (pred == -1) ? -1 : lca(pred, ci);
int lca_succ = (succ == -1) ? -1 : lca(ci, succ);
int add = depth[ci];
if (pred != -1) {
add -= depth[lca_pred];
}
if (succ != -1) {
add -= depth[lca_succ];
}
if (pred != -1 && succ != -1) {
add += depth[lca(pred, succ)];
}
total_edges += add;
nodes_in_order.insert(ci);
}
}
int dist = distance(1, ci);
int ans = 2 * total_edges - dist;
cout << ans << '\n';
}
return 0;
}