- Network
1
- @ 2026-7-25 11:52:22
#include <bits/stdc++.h>
using namespace std;
#define ll long long
#define pii pair<int, int>
struct ST {
vector<vector<int>> f;
int n, m;
vector<int> a;
ST(vector<int> a_) {
a = a_;
n = a.size();
m = __lg(n);
f.assign(m + 1, vector(n, 0));
for (int i = 1; i < n; i++) f[0][i] = i;
for (int j = 1; j <= m; j++) {
for (int i = 1; i < n; i++) {
if (i + (1 << j - 1) >= n) break;
// f[j][i] = max(f[j - 1][i], f[j - 1][i + (1 << j - 1)]);
f[j][i] = a[f[j-1][i]] < a[f[j-1][i+(1<<j-1)]] ? f[j-1][i] : f[j-1][i+(1<<j-1)];
}
}
}
int get(int l, int r) {
int len = r - l + 1;
int j = __lg(len);
return a[f[j][l]] < a[f[j][r-(1<<j)+1]] ? f[j][l] : f[j][r-(1<<j)+1];
}
};
void solve() {
int n, m;
cin >> n >> m;
vector e(n + 1, vector<int>());
for (int i = 1, u, v, w; i < n; i++) {
cin >> u >> v;
e[u].push_back(v);
e[v].push_back(u);
}
vector<pii> g(m);
for (int i = 0; i < m; i++) {
cin >> g[i].first >> g[i].second;
}
vector<int> dep(n + 1), len(n + 1), f(n + 1), dfn(n + 1), node(n + 1);
int cnt = 0;
auto dfs = [&] (auto &&self, int u, int fa)->void {
dfn[u] = ++cnt;
node[cnt] = u;
for (auto v : e[u]) {
if (v == fa) continue;
dep[v] = dep[u] + 1, len[v] = len[u] + 1;
f[v] = u;
self(self, v, u);
}
};
dfs(dfs, 1, 0);
vector<int> dep1(n + 1);
for (int i = 1; i <= n; i++) dep1[dfn[i]] = dep[i];
ST st(dep1);
auto lca = [&] (int x, int y)->int {
if (x == y) return x;
x = dfn[x], y = dfn[y];
if (x > y) swap(x, y);
auto son = node[st.get(x + 1, y)];
return f[son];
};
vector<int> d(n + 1);
for (auto [x, y] : g) {
auto l = lca(x, y);
d[x]++, d[y]++, d[l] -= 2;
}
auto dfs1 = [&] (auto &&self, int u, int fa)->void {
for (auto v : e[u]) {
if (v == fa) continue;
self(self, v, u);
d[u] += d[v];
}
};
dfs1(dfs1, 1, 0);
int cnt0 = 0;
for (int i = 2; i <= n; i++) cnt0 += d[i] == 0;
vector<int> s(n + 1);
auto dfs2 = [&] (auto &&self, int u, int fa)->void {
for (auto v : e[u]) {
if (v == fa) continue;
s[v] = s[u] + (d[v] == 1);
self(self, v, u);
}
};
dfs2(dfs2, 1, 0);
ll ans = 0;
for (auto [x, y] : g) {
auto l = lca(x, y);
ans += cnt0 + s[x] + s[y] - s[l] - s[l];
}
cout << ans << '\n';
}
signed main() {
ios::sync_with_stdio(0); cin.tie(0);
int t = 1;
// cin >> t;
while (t--) solve();
return 0;
}
/*
g++ -std=c++20 1.cpp -o 1 && 1 < in.txt > out.txt
*/
0 条评论
目前还没有评论...
信息
- ID
- 517
- 时间
- ms
- 内存
- MiB
- 难度
- 9
- 标签
- 递交数
- 10
- 已通过
- 6
- 上传者