P12357 [eJOI 2024] 贸易搭配 Many Pairs

代码难度极高。

对树上每个总部 \(h\),删除 \(h\) 后至多选择两个相邻方向。协议两端都在被选范围内便获得价值,求每个 \(h\) 的最大收益。

删除 \(h\) 后,每个邻点对应一个方向,每份协议唯一属于:

  • 两个端点都是 \(h\) ,恒定收益 \(S_h\)
  • 两个端点都在一个方向 \(U_d\)
  • 两个端点分局两个方向 \(C_{d,e}\)

固定总部 \(h\),删除后每个邻点对应一个方向。每条协议 \((a,b,w)\) 相对于 \(h\) 只有三类:两端都是 \(h\)(恒定收益 \(S_h\))、两端在同一方向(一元收益 \(U_d\))、两端分属两个方向(交叉收益 \(C_{d,e}\))。若选择方向集合大小不超过 2,总收益为

\[S_h+\sum_{d\in T}U_d+\sum_{\{d,e\}\subseteq T}C_{d,e}, \]

故答案为

\[S_h+\max\Bigl(\max_d U_d,\ \max_{d<e}(U_d+U_e+C_{d,e})\Bigr). \]

任选根,对非根节点 \(u\) 定义 \(\mathrm{inside}[u]\)(两端在子树内)、\(\mathrm{cut}[u]\)(跨父边)、\(\mathrm{down}[u]\)(一端是 parent,另一端在子树内)、\(\mathrm{up}[u]\)(一端是 \(u\),另一端在父方向)。这些可通过树上差分和子树汇总求得:\(\mathrm{inside}\) 在 LCA 累加后子树求和;\(\mathrm{cut}\) 在两端加、LCA 减 \(2w\) 后子树求和;\(\mathrm{down}\) 在直连父子时加到儿子;\(\mathrm{up}\) 在端点非 LCA 时加到该端点。

对总部 \(h\),儿子方向 \(v\) 的一元收益为 \(U_v=\mathrm{inside}[v]+\mathrm{down}[v]\),父方向为 \(U_{\mathrm{parent}}=\text{total}-\mathrm{inside}[h]-\mathrm{cut}[h]+\mathrm{up}[h]\)。交叉收益:儿子 \(v\) 与父方向的交叉为 \(\mathrm{cut}[v]-\mathrm{down}[v]-\mathrm{incident}[v]\),其中 \(\mathrm{incident}[v]\) 是一端在 \(v\) 内、另一端在兄弟方向的协议总价值;两儿子间的交叉直接从挂在 \(h\) 上的非直连协议收集,同时累加 \(\mathrm{incident}\)。

DFS 过程中,对每个节点 \(u\),先递归处理儿子并收集儿子列表,计算父方向收益 \(cfa\)。处理挂在 \(u\) 上的协议:自环时扣除并记录附加值;直连儿子时修正其 inside 和 del;非直连时将交叉边存入儿子对映射。用线段树维护每个儿子的 inside 值,枚举每个儿子作为选中方向,动态调整线段树和 \(cfa\),用线段树最大值和 \(cfa\) 分别更新选两个方向或只选父方向的答案。最后清空线段树。总复杂度 \(O((n+Q)\log n)\),空间 \(O(n\log n+n+Q)\)。

#include <bits/stdc++.h>
#define int long long
#define rep(i, l, r) for (int i = (l); i <= (r); ++ i)
#define per(i, r, l) for (int i = (r); i >= (l); -- i)
#define pb push_back
using namespace std;

const int N = 2e5 + 10;

int n, Q;
int d[N], up[N][22], lg[N], dfn[N], sz[N], idx;
vector<int> p[N];
struct node { int u, v, w; };
vector<node> q[N];

int tot = 0;
int ans[N], f[N], val[N], del[N];
vector<pair<int, int>> mp[N];

struct tree {
    int tr[N << 2];
    void clear(int p, int l, int r) {
        if (l == r) { tr[p] = 0; return; }
        int mid = (l + r) >> 1;
        clear(p << 1, l, mid);
        clear(p << 1 | 1, mid + 1, r);
        tr[p] = 0;
    }
    void upd(int x, int v, int p = 1, int l = 1, int r = n) {
        if (l == r) { tr[p] += v; return; }
        int mid = (l + r) >> 1;
        if (x <= mid) upd(x, v, p << 1, l, mid);
        else upd(x, v, p << 1 | 1, mid + 1, r);
        tr[p] = max(tr[p << 1], tr[p << 1 | 1]);
    }
} ds;

int get(int u, int v) {
    int dep = d[u] + 1;
    per(i, 20, 0) if (d[up[v][i]] >= dep) v = up[v][i];
    return v;
}

void dfs0(int u, int fa) {
    d[u] = d[fa] + 1;
    up[u][0] = fa;
    dfn[u] = ++ idx;
    sz[u] = 1;
    for (auto v : p[u]) if (v != fa) { dfs0(v, u); sz[u] += sz[v]; }
}

int lca(int x, int y) {
    if (d[x] < d[y]) swap(x, y);
    int t = d[x] - d[y];
    for (int i = 0; i < 20; i ++) if (t & (1 << i)) x = up[x][i];
    if (x == y) return x;
    per(i, 20, 0) if (up[x][i] != up[y][i]) x = up[x][i], y = up[y][i];
    return up[x][0];
}

void dfs(int u, int fa) {
    vector<pair<int, int>> son;
    for (auto &e : q[u]) f[u] += e.w;

    int w_init = val[u];

    for (auto v : p[u]) if (v != fa) {
        dfs(v, u);
        son.pb({dfn[v], v});
        val[u] += val[v];
        f[u] += f[v];
    }

    int cfa = tot + w_init - val[u] + f[u];
    int ad = 0;

    for (auto &e : q[u]) {
        if (e.u == u && e.v == u) {
            cfa -= e.w;
            ad += e.w;
            continue;
        }
        if (e.u == u) {
            int v2 = (-- upper_bound(son.begin(), son.end(), make_pair(dfn[e.v], n))) -> second;
            f[v2] += e.w;
            cfa -= e.w;
            del[v2] += e.w;
            continue;
        }
        int v1 = (-- upper_bound(son.begin(), son.end(), make_pair(dfn[e.u], n))) -> second;
        int v2 = (-- upper_bound(son.begin(), son.end(), make_pair(dfn[e.v], n))) -> second;
        mp[v1].pb({v2, e.w});
        mp[v2].pb({v1, e.w});
    }

    for (auto v : p[u]) if (v != fa) ds.upd(v, f[v]);

    for (auto v : p[u]) if (v != fa) {
        ds.upd(v, -f[v]);
        cfa += val[v] - f[v] * 2 + del[v];

        for (auto tmp : mp[v]) {
            int v2 = tmp.first, w = tmp.second;
            ds.upd(v2, w);
            cfa -= w;
        }

        ans[u] = max(ans[u], max(fa ? cfa : 0, ds.tr[1] + ad) + f[v]);

        for (auto tmp : mp[v]) {
            int v2 = tmp.first, w = tmp.second;
            ds.upd(v2, -w);
            cfa += w;
        }

        ds.upd(v, f[v]);
        cfa -= val[v] - f[v] * 2 + del[v];
    }

    for (auto v : p[u]) if (v != fa) ds.upd(v, -f[v]);

    ans[u] = max(ans[u], cfa);
}

void Jail() {
    cin >> n >> Q;
    rep(i, 2, n) lg[i] = lg[i >> 1] + 1;
    for (int i = 1; i < n; i ++) {
        int u, v;
        cin >> u >> v;
        p[u].pb(v);
        p[v].pb(u);
    }

    dfs0(1, 0);

    for (int j = 1; j <= 20; j ++)
        for (int i = 1; i <= n; i ++)
            up[i][j] = up[up[i][j - 1]][j - 1];

    for (int i = 1; i <= Q; i ++) {
        int u, v, w;
        cin >> u >> v >> w;
        if (dfn[u] > dfn[v]) swap(u, v);
        val[u] += w;
        val[v] += w;
        tot += w;
        int t = lca(u, v);
        q[t].pb({u, v, w});
    }
    dfs(1, 0);
    for (int i = 1; i <= n; i ++) cout << ans[i] << (i == n ? '\n' : ' ');
}
signed main() {
    // freopen("a.in", "r", stdin);
    // freopen("a.out", "w", stdout);
    ios::sync_with_stdio(0);
    cin.tie(0);
    cout.tie(0);
    Jail();
    return 0;
}
posted @ 2026-07-22 14:58  Koswel  阅读(7)  评论(0)    收藏  举报