RE悬赏一赞
查看原帖
RE悬赏一赞
1001524
UniqueYou楼主2023/7/31 15:56
#include <bits/stdc++.h>
using namespace std;
int n, m, log2n, diff[300005], pos[300005], de[300005], d[300005], le[300005], fa[300005][20], LCA[300005];
struct Edge
{
    int v, w;
};
vector<Edge> e[300005];
struct node
{
    int x, y, w, lca;
}a[300005];
bool cmp(node x, node y)
{
    return x.w > y.w;
}
void dfs(int u)
{
    for (int i = 0; i < (int)e[u].size(); i++)
    {
        int v = e[u][i].v;
        if (v == fa[u][0])
            continue;
        de[v] = de[u] + 1;
        d[v] = d[u] + e[u][i].w;
        pos[v] = e[u][i].w;
        fa[v][0] = u;
        dfs(v); 
    }
}
void dfs2(int u)
{
    for (int i = 0; i < (int)e[u].size(); i++)
    {
        int v = e[u][i].v;
        if (v == fa[u][0])
            continue;
        diff[v] += diff[u];
        dfs2(v); 
    }
}
int lca(int x, int y)
{
    if (de[x] < de[y]) swap(x, y);
    for (int i = 0, p = de[x] - de[y]; p; i++, p >>= 1)
        if (p & 1)
            x = fa[x][i];
    if (x == y)
        return x;
    for (int i = log2n; i >= 0; i--)
        if (fa[x][i] != fa[y][i])
            x = fa[x][i], y = fa[y][i];
    return fa[x][0];
}
bool check(int x)
{
    memset(diff, 0, sizeof(diff));
    int sum = 0;
    while (a[sum+1].w <= x)
    {
        sum++;
        diff[a[sum].x]++;
        diff[a[sum].y]++;
        diff[a[sum].lca] -= 2;
    }
    dfs2(1);
    int mx = 0;
    for (int i = 2; i <= n; i++)
        if (diff[i] == sum)
            mx = max(mx, pos[i]);
	
    if (mx >= a[1].w - x) 
        return 1;
	else 
        return 0;
}
int main()
{
    cin >> n >> m;
    log2n = log2(n);
    for (int i = 1; i < n; i++)
    {
        int u, v, w;
        cin >> u >> v >> w;
        e[u].push_back((Edge){v, w});
        e[v].push_back((Edge){u, w});
    }
    dfs(1);
    for (int j = 1; j <= log2(n); j++)
        for (int i = 1; i <= n; i++)
            fa[i][j] = fa[fa[i][j-1]][j-1];
    for (int i = 1; i <= m; i++)
    {
        cin >> a[i].x >> a[i].y;
        a[i].lca = lca(a[i].x, a[i].y);
        a[i].w = d[a[i].x] + d[a[i].y] - 2 * d[a[i].lca];
    }
    sort(a+1, a+m+1, cmp);
    int l = 0, r = 0x3f3f3f3f, ans = -1;
    while (l <= r)
    {
        int mid = (l + r) / 2;
        if (check(mid))
        {
            ans = mid;
            r = mid - 1;
        }
        else
            l = mid + 1;
        cout << l;
    }
    cout << ans;
    return 0;
}
2023/7/31 15:56
加载中...