#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 条评论

目前还没有评论...