【学习笔记】点分治 & 点分树

· · 算法·理论

点分治

点分治是一种树上分治算法,常用于解决各类树上路径问题。

模板题

:::info[例题 0]{open}

P3806 【模板】点分治

给定一棵有 n 个节点的树,边带权,m 次询问,每次给定 k,判断是否存在一条长度恰为 k 的路径。

0.2s,n \le 10^4m \le 100

:::

设置分治中心 rt,则树上的所有路径可被分为两类:经过 rt 的,和不经过 rt 的。对于后者,我们将其留到后续的递归中处理,在当前这一层的递归中,我们只关注经过 rt 的路径。而这些经过 rt 的路径又可被分为两类:以 rt 为端点的,和不以 rt 为端点的。后者显然可以由两条以 rt 为端点的路径合并得到。

受此启发,我们枚举 rt 的所有子节点,依次处理它们的子树。维护数组 f_i,表示已遍历过的子树中是否存在与 rt 的距离恰好为 i 的节点。设当前遍历的节点为 u,其与 rt 的距离为 dis_u,那么对于树上一条经过 rt 的路径 (u,v),其长度 k=dis_u+dis_v。要判断这样的一条路径是否存在,只需查询 f_{k-dis_u} 即可。遍历完一棵子树后,我们将该子树中所有节点的 dis 值插入 f 中,以便后续查询。递归完一层后,分别进入 rt 的每棵子树继续递归下去即可。

每层递归需要遍历当前连通块中的所有节点。当递归层数为 h 时,该算法的复杂度为 O(hn)。若每次选取重心为分治中心,复杂度将为 O(n \log n),这是因为重心的性质决定了每递归一层,树的大小至少减半。本题有 m 次询问,故复杂度为 O(nm \log n)。注意不要对每次询问都跑一次点分治,而是要把询问离线下来,一次处理完。虽然两种做法的复杂度均为 O(nm \log n),但前者会把点分治的常数开销乘到 m 上,实际运行效率要慢很多。

实现上需要注意,每层递归结束后需要清空 f 数组,但不能直接 memset 清空,应把所有改动过的位置存下来,用多少清多少。其余细节见代码。

:::success[Code]{open}

#include <bits/stdc++.h>
using namespace std;
const int MAXN = 1e4 + 10, MAXM = 110, MAXK = 1e7 + 10;
vector <pair<int, int>> adj[MAXN];
int qr[MAXM], siz[MAXN], que[MAXN], tmp[MAXN], n, m, cnt, cur;
bool vis[MAXN], flag[MAXK], ans[MAXM];
void dfs1(int u, int p){
    siz[u] = 1;
    for (auto [v, w] : adj[u]){
        if (v == p || vis[v]){
            continue;
        }
        dfs1(v, u);
        siz[u] += siz[v];
    }
    return;
}
void dfs2(int u, int p, int tot, int &rt){
    int mx = tot - siz[u];
    for (auto [v, w] : adj[u]){
        if (v == p || vis[v]){
            continue;
        }
        dfs2(v, u, tot, rt);
        mx = max(mx, siz[v]);
    }
    if (mx * 2 <= tot){
        rt = u;
    }
    return;
}
void dfs3(int u, int p, int d){
    tmp[++cnt] = d;
    for (auto [v, w] : adj[u]){
        if (v == p || vis[v]){
            continue;
        }
        dfs3(v, u, d + w);
    }
    return;
}
void solve(int u){
    cnt = cur = 0;
    dfs1(u, 0);
    int rt = 0;
    dfs2(u, 0, siz[u], rt);
    vis[rt] = flag[0] = true;
    que[++cur] = 0;
    for (auto [v, w] : adj[rt]){
        if (vis[v]){
            continue;
        }
        cnt = 0;
        dfs3(v, rt, w);
        for (int i = 1; i <= cnt; i++){
            for (int j = 1; j <= m; j++){
                if (qr[j] - tmp[i] >= 0){
                    ans[j] |= flag[qr[j] - tmp[i]];
                }
            }
        }
        for (int i = 1; i <= cnt; i++){
            if (tmp[i] < MAXK){
                que[++cur] = tmp[i];
                flag[tmp[i]] = true;
            }
        }
    }
    for (int i = 1; i <= cur; i++){
        flag[que[i]] = false;
    }
    for (auto [v, w] : adj[rt]){
        if (vis[v]){
            continue;
        }
        solve(v);
    }
    return;
}
signed main(){
    ios::sync_with_stdio(false);
    cin.tie(0);
    cout.tie(0);
    cin >> n >> m;
    for (int i = 1; i < n; i++){
        int u, v, w;
        cin >> u >> v >> w;
        adj[u].push_back({v, w});
        adj[v].push_back({u, w});
    }
    for (int i = 1; i <= m; i++){
        cin >> qr[i];
    }
    solve(1);
    for (int i = 1; i <= m; i++){
        cout << (ans[i] ? "AYE\n" : "NAY\n");
    }
    return 0;
}

:::

例题

I

:::info[例题 1]{open}

P4178 Tree

给定一棵有 n 个节点的树,边带权,求有多少条长度不超过 k 的路径。

1s,n \le 4 \times 10^4k \le 2 \times 10^4

:::

大体思路完全一致。区别只在于,本题需要计数,需将 f_i 的含义改为长度恰为 i 的路径条数。每次需对 f 数组的一段前缀求和,树状数组维护即可。

:::success[Code]

#include <bits/stdc++.h>
using namespace std;
const int MAXN = 4e4 + 10;
vector <pair<int, int>> adj[MAXN];
int siz[MAXN], que[MAXN], dis[MAXN], n, k, cur, cnt, ans;
bool vis[MAXN];
struct BIT{
    int v[MAXN];
    int lowbit(int x){
        return x & (-x);
    }
    void modify(int u, int x){
        while (u <= 4e4 + 1){
            v[u] += x;
            u += lowbit(u);
        }
        return;
    }
    int query(int u){
        int res = 0;
        while (u){
            res += v[u];
            u -= lowbit(u);
        }
        return res;
    }
}tr;
void dfs1(int u, int p){
    siz[u] = 1;
    for (auto [v, w] : adj[u]){
        if (v == p || vis[v]){
            continue;
        }
        dfs1(v, u);
        siz[u] += siz[v];
    }
    return;
}
void dfs2(int u, int p, int tot, int &rt){
    int mx = tot - siz[u];
    for (auto [v, w] : adj[u]){
        if (v == p || vis[v]){
            continue;
        }
        dfs2(v, u, tot, rt);
        mx = max(mx, siz[v]);
    }
    if (mx * 2 <= tot){
        rt = u;
    }
    return;
}
void dfs3(int u, int p, int d){
    dis[++cnt] = d;
    for (auto [v, w] : adj[u]){
        if (v == p || vis[v]){
            continue;
        }
        dfs3(v, u, d + w);
    }
    return;
}
void solve(int u){
    cur = 0;
    dfs1(u, 0);
    int rt = 0;
    dfs2(u, 0, siz[u], rt);
    vis[rt] = true;
    tr.modify(1, 1);
    que[++cur] = 1;
    for (auto [v, w] : adj[rt]){
        if (vis[v]){
            continue;
        }
        cnt = 0;
        dfs3(v, rt, w);
        for (int i = 1; i <= cnt; i++){
            if (k - dis[i] >= 0){
                ans += tr.query(k - dis[i] + 1);
            }
        }
        for (int i = 1; i <= cnt; i++){
            if (dis[i] <= 4e4){
                tr.modify(dis[i] + 1, 1);
                que[++cur] = dis[i] + 1;
            }
        }
    }
    for (int i = 1; i <= cur; i++){
        tr.modify(que[i], -1);
    }
    for (auto [v, w] : adj[rt]){
        if (vis[v]){
            continue;
        }
        solve(v);
    }
    return;
}
signed main(){
    ios::sync_with_stdio(false);
    cin.tie(0);
    cout.tie(0);
    cin >> n;
    for (int i = 1; i < n; i++){
        int u, v, w;
        cin >> u >> v >> w;
        adj[u].push_back({v, w});
        adj[v].push_back({u, w});
    }
    cin >> k;
    solve(1);
    cout << ans << "\n";
    return 0;
}

:::

II

:::info[例题 2]{open}

P4149 [IOI 2011] Race

给定一棵有 n 个节点的树,边带权,求所有边权和恰为 k 的路径的最小边数。

3s,n \le 2 \times 10^5k \le 10^6

:::

依旧是一样的思路。区别只在于,本题要求最小边数,需将 f_i 的含义改为长度恰为 i 的路径的最小边数,并需额外记录每个节点离 rt 的边数。

:::success[Code]

#include <bits/stdc++.h>
using namespace std;
const int MAXN = 2e5 + 10, MAXM = 1e6 + 10;
const int INF = 0x3f3f3f3f;
vector <pair<int, int>> adj[MAXN];
int siz[MAXN], minn[MAXM], que[MAXN], dis[MAXN], dep[MAXN], n, k, ans, cur, cnt;
bool vis[MAXN];
void dfs1(int u, int p){
    siz[u] = 1;
    for (auto [v, w] : adj[u]){
        if (v == p || vis[v]){
            continue;
        }
        dfs1(v, u);
        siz[u] += siz[v];
    }
    return;
}
void dfs2(int u, int p, int tot, int &rt){
    int mx = tot - siz[u];
    for (auto [v, w] : adj[u]){
        if (v == p || vis[v]){
            continue;
        }
        dfs2(v, u, tot, rt);
        mx = max(mx, siz[v]);
    }
    if (mx * 2 <= tot){
        rt = u;
    }
    return;
}
void dfs3(int u, int p, int d1, int d2){
    dis[++cnt] = d1;
    dep[cnt] = d2;
    for (auto [v, w] : adj[u]){
        if (v == p || vis[v]){
            continue;
        }
        dfs3(v, u, d1 + w, d2 + 1);
    }
    return;
}
void solve(int u){
    cur = 0;
    dfs1(u, 0);
    int rt = 0;
    dfs2(u, 0, siz[u], rt);
    vis[rt] = true;
    minn[0] = 0;
    que[++cur] = 0;
    for (auto [v, w] : adj[rt]){
        if (vis[v]){
            continue;
        }
        cnt = 0;
        dfs3(v, rt, w, 1);
        for (int i = 1; i <= cnt; i++){
            if (k - dis[i] >= 0){
                ans = min(ans, minn[k - dis[i]] + dep[i]);
            }
        }
        for (int i = 1; i <= cnt; i++){
            if (dis[i] <= 1e6){
                minn[dis[i]] = min(minn[dis[i]], dep[i]);
                que[++cur] = dis[i];
            }
        }
    }
    for (int i = 1; i <= cur; i++){
        minn[que[i]] = INF;
    }
    for (auto [v, w] : adj[rt]){
        if (!vis[v]){
            solve(v);
        }
    }
    return;
}
signed main(){
    ios::sync_with_stdio(false);
    cin.tie(0);
    cout.tie(0);
    cin >> n >> k;
    for (int i = 1; i < n; i++){
        int u, v, w;
        cin >> u >> v >> w;
        adj[u].push_back({v, w});
        adj[v].push_back({u, w});
    }
    memset(minn, 0x3f, sizeof(minn));
    ans = INF;
    solve(1);
    cout << (ans > 1e9 ? -1 : ans) << "\n";
    return 0;
}

:::

III

:::info[例题 3]{open}

P6626 [省选联考 2020 B 卷] 消息传递

给定一棵有 n 个节点的树,m 次询问,每次给定 x,k,求与节点 x 的距离恰为 k 的节点数。

2s,n,m \le 10^5

:::

仿照模板题的做法,将询问离线下来挂在 x 上。与之前的题目不同的是,对于每棵子树,我们不再只需要考虑比它先被遍历的子树,而是要考虑除它之外的所有子树。

据此,我们先遍历 rt 子树中的所有节点,将它们 dis 值全部加入 f 中,再遍历 rt 的每个子节点,先将该子节点子树中的所有 dis 值从 f 中扣掉,再处理挂在该子树上的所有询问,最后把被扣掉的部分加回去。

:::success[Code]

#include <bits/stdc++.h>
using namespace std;
const int MAXN = 1e5 + 10;
vector <int> adj[MAXN];
vector <pair<int, int>> qr[MAXN], idx[MAXN];
int siz[MAXN], qx[MAXN], qk[MAXN], cnt[MAXN], que[MAXN], ans[MAXN], n, m, cur;
bool vis[MAXN];
void dfs1(int u, int p){
    siz[u] = 1;
    for (int v : adj[u]){
        if (v == p || vis[v]){
            continue;
        }
        dfs1(v, u);
        siz[u] += siz[v];
    }
    return;
}
void dfs2(int u, int p, int tot, int &rt){
    int mx = tot - siz[u];
    for (int v : adj[u]){
        if (v == p || vis[v]){
            continue;
        }
        dfs2(v, u, tot, rt);
        mx = max(mx, siz[v]);
    }
    if (mx * 2 <= tot){
        rt = u;
    }
    return;
}
void dfs3(int rt, int u, int p, int d){
    que[++cur] = d;
    idx[rt].push_back({u, d});
    cnt[d]++;
    for (int v : adj[u]){
        if (v == p || vis[v]){
            continue;
        }
        dfs3(rt, v, u, d + 1);
    }
    return;
}
void dfs4(int u, int p, int d, int o){
    cnt[d] += o;
    for (int v : adj[u]){
        if (v == p || vis[v]){
            continue;
        }
        dfs4(v, u, d + 1, o);
    }
    return;
}
void solve(int u){
    cur = 0;
    dfs1(u, 0);
    int rt = 0;
    dfs2(u, 0, siz[u], rt);
    vis[rt] = true;
    cnt[0]++;
    que[++cur] = 0;
    for (int v : adj[rt]){
        if (vis[v]){
            continue;
        }
        dfs3(v, v, rt, 1);
    }
    for (auto [k, id] : qr[rt]){
        ans[id] += cnt[k];
    }
    for (int v : adj[rt]){
        if (vis[v]){
            continue;
        }
        dfs4(v, rt, 1, -1);
        for (auto [x, d] : idx[v]){
            for (auto [k, id] : qr[x]){
                if (k >= d){
                    ans[id] += cnt[k - d];
                }
            }
        }
        dfs4(v, rt, 1, 1);
        idx[v].clear();
    }
    for (int i = 1; i <= cur; i++){
        cnt[que[i]] = 0;
    }
    for (int v : adj[rt]){
        if (vis[v]){
            continue;
        }
        solve(v);
    }
    return;
}
signed main(){
    ios::sync_with_stdio(false);
    cin.tie(0);
    cout.tie(0);
    int t;
    cin >> t;
    while (t--){
        cin >> n >> m;
        for (int i = 1; i <= n; i++){
            adj[i].clear();
            qr[i].clear();
            vis[i] = false;
        }
        for (int i = 1; i <= m; i++){
            ans[i] = 0;
        }
        for (int i = 1; i < n; i++){
            int u, v;
            cin >> u >> v;
            adj[u].push_back(v);
            adj[v].push_back(u);
        }
        for (int i = 1; i <= m; i++){
            cin >> qx[i] >> qk[i];
            qr[qx[i]].push_back({qk[i], i});
        }
        solve(1);
        for (int i = 1; i <= m; i++){
            cout << ans[i] << "\n";
        }
    }
    return 0;
}

:::

IV

:::info[例题 4]{open}

P10421 [蓝桥杯 2023 国 A] 树上的路径

给定一棵有 n 个节点的树,求 \sum\limits_{i=1}^n \sum\limits_{j=i+1}^n dis(i,j) \cdot [L \le dis(i,j) \le R] 的值,即所有长度在 [L,R] 范围内的路径的长度和。

6s,n \le 10^6

:::

一个自然的想法是,对于一条长度为 k 的路径,我们在将其加入 f 数组时就将它的贡献乘上 k 的系数。但这样做是不行的,因为查询 f 数组时会有一个偏移量,导致我们无法正确计算贡献。

考虑将这个偏移量分离出来:

\sum\limits_{i=L}^R f_{i-k} \cdot i=\sum\limits_{i=L}^R f_{i-k} \cdot (i-k)+\sum\limits_{i=L}^R f_{i-k} \cdot k

那么对于左边这部分,就可以用上面的方法计算了,右边的部分系数固定,也是容易计算的。具体地,开两棵树状数组,一棵在加入时乘上系数(维护左半部分),另一棵不乘系数(维护右半部分),查询时求区间和即可。

:::success[Code]

#include <bits/stdc++.h>
#define int long long
using namespace std;
const int MAXN = 1e6 + 10;
vector <int> adj[MAXN];
int siz[MAXN], que[MAXN], dis[MAXN], n, L, R, cur, cnt, ans;
bool vis[MAXN];
struct BIT{
    int v[MAXN];
    int lowbit(int x){
        return x & (-x);
    }
    void modify(int u, int x){
        while (u < MAXN){
            v[u] += x;
            u += lowbit(u);
        }
        return;
    }
    int query(int u){
        int res = 0;
        while (u){
            res += v[u];
            u -= lowbit(u);
        }
        return res;
    }
}tr1, tr2;
void dfs1(int u, int p){
    siz[u] = 1;
    for (int v : adj[u]){
        if (v == p || vis[v]){
            continue;
        }
        dfs1(v, u);
        siz[u] += siz[v];
    }
    return;
}
void dfs2(int u, int p, int tot, int &rt){
    int mx = tot - siz[u];
    for (int v : adj[u]){
        if (v == p || vis[v]){
            continue;
        }
        dfs2(v, u, tot, rt);
        mx = max(mx, siz[v]);
    }
    if (mx * 2 <= tot){
        rt = u;
    }
    return;
}
void dfs3(int u, int p, int d){
    dis[++cnt] = d;
    for (int v : adj[u]){
        if (v == p || vis[v]){
            continue;
        }
        dfs3(v, u, d + 1);
    }
    return;
}
void solve(int u){
    cur = 0;
    dfs1(u, 0);
    int rt = 0;
    dfs2(u, 0, siz[u], rt);
    tr2.modify(1, 1);
    vis[rt] = true;
    que[++cur] = 1;
    for (int v : adj[rt]){
        if (vis[v]){
            continue;
        }
        cnt = 0;
        dfs3(v, rt, 1);
        for (int i = 1; i <= cnt; i++){
            if (dis[i] > R){
                continue;
            }
            else if (dis[i] > L){
                ans += tr1.query(R - dis[i] + 1) + dis[i] * tr2.query(R - dis[i] + 1);
            }
            else{
                ans += tr1.query(R - dis[i] + 1) - tr1.query(L - dis[i]) + dis[i] * (tr2.query(R - dis[i] + 1) - tr2.query(L - dis[i]));
            }
        }
        for (int i = 1; i <= cnt; i++){
            que[++cur] = dis[i] + 1;
            tr1.modify(dis[i] + 1, dis[i]);
            tr2.modify(dis[i] + 1, 1);
        }
    }
    for (int i = 1; i <= cur; i++){
        tr1.modify(que[i], 1 - que[i]);
        tr2.modify(que[i], -1);
    }
    for (int v : adj[rt]){
        if (vis[v]){
            continue;
        }
        solve(v);
    }
    return;
}
signed main(){
    ios::sync_with_stdio(false);
    cin.tie(0);
    cout.tie(0);
    cin >> n >> L >> R;
    for (int i = 2; i <= n; i++){
        int p;
        cin >> p;
        adj[i].push_back(p);
        adj[p].push_back(i);
    }
    solve(1);
    cout << ans << "\n";
    return 0;
}

:::

V

:::info[例题 5]{open}

P3714 [BJOI2017] 树的难题

给定一棵有 n 个节点的树,边有 1 \sim m 之间的颜色,颜色 i 的权值为 c_i。对于一条路径,将其经过的所有边的颜色按顺序组成一个序列,该路径的权值为该序列中各极长颜色段的颜色权值之和。求所有长度在 [l,r] 范围内的路径的最大权值。

2s,n,m \le 2 \times 10^5

:::

对于分治中心 rt,将 (rt,u),(rt,v) 两段路径拼接时,若它们起始边的颜色不同,则可将它们的权值直接相加,否则还要减掉起始边的颜色权值。

考虑分别维护这两种情况,开两棵线段树,均以路径长度为下标,一棵维护与当前子树起始边颜色不同的路径的最大权值,另一棵维护与当前子树起始边颜色相同的路径的最大权值。为了让颜色相同的边一起被处理,在跑点分治之前,我们将每个节点的出边按颜色排序。这样我们每处理完一棵子树,就先把该子树的信息插入第二棵线段树,然后在处理下一棵子树前,如果起始边的颜色发生了改变,再把第二棵线段树中的信息暴力合并到第一棵线段树中。

:::success[Code]

#include <bits/stdc++.h>
#define lc (u << 1)
#define rc ((u << 1) | 1)
#define mid ((l + r) >> 1)
using namespace std;
const int MAXN = 2e5 + 10;
const int INF = 0x3f3f3f3f;
vector <pair<int, int>> adj[MAXN];
int c[MAXN], siz[MAXN], que1[MAXN], que2[MAXN], dis[MAXN], sum[MAXN], n, m, L, R, cur1, cur2, cur3, ans = -INF;
bool vis[MAXN];
struct Segment_tree{
    int mx[MAXN * 4];
    void pushup(int u){
        mx[u] = max(mx[lc], mx[rc]);
        return;
    }
    void build(int u, int l, int r){
        mx[u] = -INF;
        if (l == r){
            return;
        }
        build(lc, l, mid);
        build(rc, mid + 1, r);
        return;
    }
    void chkmax(int u, int l, int r, int pos, int val){
        if (l == r){
            mx[u] = max(mx[u], val);
            return;
        }
        if (pos <= mid){
            chkmax(lc, l, mid, pos, val);
        }
        else{
            chkmax(rc, mid + 1, r, pos, val);
        }
        pushup(u);
        return;
    }
    void assign(int u, int l, int r, int pos, int val){
        if (l == r){
            mx[u] = val;
            return;
        }
        if (pos <= mid){
            assign(lc, l, mid, pos, val);
        }
        else{
            assign(rc, mid + 1, r, pos, val);
        }
        pushup(u);
        return;
    }
    int query(int u, int l, int r, int ql, int qr){
        if (ql <= l && r <= qr){
            return mx[u];
        }
        int res = -INF;
        if (ql <= mid){
            res = max(res, query(lc, l, mid, ql, qr));
        }
        if (qr > mid){
            res = max(res, query(rc, mid + 1, r, ql, qr));
        }
        return res;
    }
}tr1, tr2;
void dfs1(int u, int p){
    siz[u] = 1;
    for (auto [col, v] : adj[u]){
        if (v == p || vis[v]){
            continue;
        }
        dfs1(v, u);
        siz[u] += siz[v];
    }
    return;
}
void dfs2(int u, int p, int tot, int &rt){
    int mx = tot - siz[u];
    for (auto [col, v] : adj[u]){
        if (v == p || vis[v]){
            continue;
        }
        dfs2(v, u, tot, rt);
        mx = max(mx, siz[v]);
    }
    if (mx * 2 <= tot){
        rt = u;
    }
    return;
}
void dfs3(int u, int p, int d, int x, int last){
    dis[++cur3] = d;
    sum[cur3] = x;
    for (auto [col, v] : adj[u]){
        if (v == p || vis[v]){
            continue;
        }
        dfs3(v, u, d + 1, x + (col != last) * c[col], col);
    }
    return;
}
void solve(int u){
    cur1 = cur2 = 0;
    dfs1(u, 0);
    int rt = 0;
    dfs2(u, 0, siz[u], rt);
    vis[rt] = true;
    tr2.chkmax(1, 0, n - 1, 0, 0);
    que1[++cur1] = 0;
    que2[++cur2] = 0;
    int last = 0;
    for (auto [col, v] : adj[rt]){
        if (vis[v]){
            continue;
        }
        if (col != last){
            for (int i = 1; i <= cur2; i++){
                int tmp = tr2.query(1, 0, n - 1, que2[i], que2[i]);
                tr1.chkmax(1, 0, n - 1, que2[i], tmp);
                tr2.assign(1, 0, n - 1, que2[i], -INF);
            }
            cur2 = 0;
        }
        cur3 = 0;
        dfs3(v, rt, 1, c[col], col);
        for (int i = 1; i <= cur3; i++){
            if (L >= dis[i]){
                int tmp1 = tr1.query(1, 0, n - 1, L - dis[i], R - dis[i]) + sum[i];
                int tmp2 = tr2.query(1, 0, n - 1, L - dis[i], R - dis[i]) - c[col] + sum[i];
                ans = max(ans, max(tmp1, tmp2));
            }
            else if (R >= dis[i]){
                int tmp1 = tr1.query(1, 0, n - 1, 0, R - dis[i]) + sum[i];
                int tmp2 = tr2.query(1, 0, n - 1, 0, R - dis[i]) - c[col] + sum[i];
                ans = max(ans, max(tmp1, tmp2));
            }
        }
        for (int i = 1; i <= cur3; i++){
            tr2.chkmax(1, 0, n - 1, dis[i], sum[i]);
            que1[++cur1] = dis[i];
            que2[++cur2] = dis[i];
        }
        last = col;
    }
    for (int i = 1; i <= cur1; i++){
        tr1.assign(1, 0, n - 1, que1[i], -INF);
        tr2.assign(1, 0, n - 1, que1[i], -INF);
    }
    for (auto [col, v] : adj[rt]){
        if (!vis[v]){
            solve(v);
        }
    }
    return;
}
signed main(){
    ios::sync_with_stdio(false);
    cin.tie(0);
    cout.tie(0);
    cin >> n >> m >> L >> R;
    for (int i = 1; i <= m; i++){
        cin >> c[i];
    }
    for (int i = 1; i < n; i++){
        int u, v, col;
        cin >> u >> v >> col;
        adj[u].push_back({col, v});
        adj[v].push_back({col, u});
    }
    for (int i = 1; i <= n; i++){
        sort(adj[i].begin(), adj[i].end());
    }
    tr1.build(1, 0, n - 1);
    tr2.build(1, 0, n - 1);
    solve(1);
    cout << ans << "\n";
    return 0;
}

:::

VI

:::info[例题 6]{open}

P5306 [COCI 2018/2019 #5] Transport

给定一棵有 n 个节点的树,点带权,边带权。从某个节点出发,每经过一条边需要支付与边权相等的代价,每经过一个点可以获得与点权相等的点数。求有多少个有序对 (u,v),满足 u \ne v,且可以从 u 出发抵达 v

1s,n \le 10^5

:::

对于分治中心 rt,记 s_i 为路径 rt \rightarrow i 的点权和,t_i 为路径 rt \rightarrow i 的边权和。对于一条经过 rt 的路径 u \rightarrow v,将其拆为 u \rightarrow rtrt \rightarrow v 两段:

rt 子树内的所有节点遍历一遍,存下所有的合法的 rt_y-s_{fa_y},排序后双指针匹配即可。注意需要扣除 u,v 在同一棵子树内的贡献,也就是还要对每棵子树跑一遍双指针。

:::success[Code]

#include <bits/stdc++.h>
using namespace std;
typedef long long ll;
const int MAXN = 1e5 + 10;
const ll INF = 0x3f3f3f3f3f3f3f3f;
vector <pair<int, ll>> adj[MAXN];
int a[MAXN], n;
int siz[MAXN], cur1, cur2, cur3, cur4;
bool vis[MAXN];
ll b[MAXN], c[MAXN], d[MAXN], e[MAXN], ans;
void dfs1(int u, int p){
    siz[u] = 1;
    for (auto [v, w] : adj[u]){
        if (v == p || vis[v]){
            continue;
        }
        dfs1(v, u);
        siz[u] += siz[v];
    }
    return;
}
void dfs2(int u, int p, int tot, int &rt){
    int mx = tot - siz[u];
    for (auto [v, w] : adj[u]){
        if (v == p || vis[v]){
            continue;
        }
        dfs2(v, u, tot, rt);
        mx = max(mx, siz[v]);
    }
    if (mx * 2 <= tot){
        rt = u;
    }
    return;
}
void dfs3(int st, int u, int p, ll mx, ll suma, ll sumd){
    mx = max(mx, suma - sumd);
    if (mx <= suma - sumd){
        b[++cur1] = d[++cur3] = suma - sumd;
    }
    for (auto [v, w] : adj[u]){
        if (v == p || vis[v]){
            continue;
        }
        dfs3(st, v, u, mx, suma + a[v], sumd + w);
    }
    return;
}
void dfs4(int st, int u, int p, ll mx, ll suma, ll sumd){
    mx = max(mx, sumd - suma);
    c[++cur2] = e[++cur4] = mx;
    for (auto [v, w] : adj[u]){
        if (v == p || vis[v]){
            continue;
        }
        dfs4(st, v, u, mx, suma + a[u], sumd + w);
    }
    return;
}
void build(int u){
    cur1 = cur2 = 0;
    dfs1(u, 0);
    int rt = 0;
    dfs2(u, 0, siz[u], rt);
    vis[rt] = true;
    b[++cur1] = a[rt];
    c[++cur2] = -INF;
    for (auto [v, w] : adj[rt]){
        if (vis[v]){
            continue;
        }
        cur3 = cur4 = 0;
        dfs3(v, v, rt, a[rt], a[rt] + a[v], w);
        dfs4(v, v, rt, -INF, 0, w);
        sort(d + 1, d + cur3 + 1, greater<ll>());
        sort(e + 1, e + cur4 + 1, greater<ll>());
        for (int i = 1, j = 0; i <= cur4; i++){
            while (j < cur3 && e[i] <= d[j + 1]){
                j++;
            }
            ans -= j;
        }
    }
    sort(b + 1, b + cur1 + 1, greater<ll>());
    sort(c + 1, c + cur2 + 1, greater<ll>());
    for (int i = 1, j = 0; i <= cur2; i++){
        while (j < cur1 && c[i] <= b[j + 1]){
            j++;
        }
        ans += j;
    }
    for (auto [v, w] : adj[rt]){
        if (!vis[v]){
            build(v);
        }
    }
    return;
}
signed main(){
    ios::sync_with_stdio(false);
    cin.tie(0);
    cout.tie(0);
    cin >> n;
    for (int i = 1; i <= n; i++){
        cin >> a[i];
    }
    for (int i = 1; i < n; i++){
        int u, v, w;
        cin >> u >> v >> w;
        adj[u].push_back({v, w});
        adj[v].push_back({u, w});
    }
    build(1);
    cout << ans - n << "\n";
    return 0;
}

:::

点分树

你可能已经注意到了,上面的所有题目都没有涉及修改操作。若涉及修改,或题目要求强制在线,则需要用到点分树。

模板题

:::info[例题 0]{open}

P6329 【模板】点分树 / 震波

给定一棵有 n 个节点的树,点带权。m 次操作,分为两种类型:

2s,强制在线,n,m \le 10^5

:::

点分树,顾名思义,是一棵在点分治过程中构造出来的树。具体的构造方法是:将每一层递归的分治中心与上一层的分治中心连边。

点分树有以下两条关键性质:

其中,第一条性质用于确保复杂度。有了这条性质,许多看似暴力的操作(如在点分树上跳祖先、对每个节点分别开一个数据结构维护子树内信息)都具有了正确的复杂度。

第二条性质用于处理路径信息。回到本题。对于节点 u,我们在点分树上跳它的祖先,设当前跳到的节点为 p,则原树上的任意一条经过 p 的路径 (u,v) 均可被拆解为 (u,p),(p,v) 两段。对每个节点开一个数据结构,维护该节点子树中所有节点到该节点的距离(注意是在原树上的距离)。对于修改,影响的只是被修改节点到根节点路径上的 O(\log n) 个节点,暴力修改即可。对于查询,还是暴力跳祖先,跳到祖先 p 时,只需查询与 p 距离不超过 k-dis_{u,p} 的节点个数即可。这就是求一段前缀和,树状数组可以维护。

但这样做有个问题:在上面的做法中,我们只考虑了路径经过 p 的情况,而对于没有经过 p 的路径,它的贡献在前面的节点已经被统计过了。若我们仍按 dis_{u,p}+dis_{p,v} 计算,就多往返了 (p,v) 这段距离,这就错了。也就是说,我们需要把 p 到节点 u 方向上的这棵子树的贡献扣掉。为此,我们对每个节点再开一棵树状数组,维护该节点子树中所有节点到该节点的父亲的距离。每次被扣掉的贡献也是 \le k-dis_{u,son_p} 的这段前缀和。

对于修改和查询,跳祖先都是 O(\log n) 的,在每个祖先处查一次前缀和,总复杂度 O(n \log^2 n)。可以看出,由于需在每个节点处维护数据结构,点分树的复杂度一般都是 2log 的。实现时需要注意,肯定不能把每棵树状数组都开到 10^5,应该用多少开多少。以及,最好写基于欧拉序和 RMQ 的 O(n \log n)-O(1) LCA。其余细节见代码。

:::success[Code]{open}

#include <bits/stdc++.h>
using namespace std;
const int MAXN = 1e5 + 10;
vector <int> adj[MAXN];
int a[MAXN], fa[MAXN], siz[MAXN], dep[MAXN], ord[MAXN * 2], in[MAXN], mxd1[MAXN], mxd2[MAXN], st[MAXN * 2][20], lg[MAXN * 2], n, m, cur;
bool vis[MAXN];
struct BIT{
    vector <int> sum;
    void init(int len){
        sum.resize(len + 2, 0);
        return;
    }
    int lowbit(int x){
        return x & (-x);
    }
    void modify(int u, int x){
        u++;
        while (u < (int)sum.size()){
            sum[u] += x;
            u += lowbit(u);
        }
        return;
    }
    int query(int u){
        u++;
        u = min(u, (int)sum.size() - 1);
        int res = 0;
        while (u){
            res += sum[u];
            u -= lowbit(u);
        }
        return res;
    }
}tr1[MAXN], tr2[MAXN];
void dfs_lca(int u, int p){
    dep[u] = dep[p] + 1;
    ord[++cur] = u;
    in[u] = cur;
    for (int v : adj[u]){
        if (v != p){
            dfs_lca(v, u);
            ord[++cur] = u;
        }
    }
}
int get_min(int u, int v){
    return (dep[u] < dep[v] ? u : v);
}
void init_st(){
    lg[1] = 0;
    for (int i = 2; i <= cur; i++){
        lg[i] = lg[i / 2] + 1;
    }
    for (int i = 1; i <= cur; i++){
        st[i][0] = ord[i];
    }
    for (int i = 1; i <= lg[cur]; i++){
        for (int j = 1; j <= cur - (1 << i) + 1; j++){
            st[j][i] = get_min(st[j][i - 1], st[j + (1 << (i - 1))][i - 1]);
        }
    }
    return;
}
int get_lca(int u, int v){
    int l = in[u], r = in[v];
    if (l > r){
        swap(l, r);
    }
    int k = lg[r - l + 1];
    return get_min(st[l][k], st[r - (1 << k) + 1][k]);
}
int get_dis(int u, int v){
    int lca = get_lca(u, v);
    return dep[u] + dep[v] - 2 * dep[lca];
}
void dfs1(int u, int p){
    siz[u] = 1;
    for (int v : adj[u]){
        if (v == p || vis[v]){
            continue;
        }
        dfs1(v, u);
        siz[u] += siz[v];
    }
    return;
}
void dfs2(int u, int p, int tot, int &rt){
    int mx = tot - siz[u];
    for (int v : adj[u]){
        if (v == p || vis[v]){
            continue;
        }
        dfs2(v, u, tot, rt);
        mx = max(mx, siz[v]);
    }
    if (mx * 2 <= tot){
        rt = u;
    }
    return;
}
void build(int u, int p){
    dfs1(u, 0);
    int rt = 0;
    dfs2(u, 0, siz[u], rt);
    vis[rt] = true;
    fa[rt] = p;
    for (int v : adj[rt]){
        if (!vis[v]){
            build(v, rt);
        }
    }
    return;
}
void modify_tree(int x, int val){
    for (int u = x; u; u = fa[u]){
        tr1[u].modify(get_dis(u, x), val);
        if (fa[u]){
            tr2[u].modify(get_dis(fa[u], x), val);
        }
    }
    return;
}
int query_tree(int x, int k){
    int sum = tr1[x].query(k);
    for (int u = x; fa[u]; u = fa[u]){
        int d = get_dis(fa[u], x);
        if (k >= d){
            sum += tr1[fa[u]].query(k - d);
            sum -= tr2[u].query(k - d);
        }
    }
    return sum;
}
signed main(){
    ios::sync_with_stdio(false);
    cin.tie(0);
    cout.tie(0);
    cin >> n >> m;
    for (int i = 1; i <= n; i++){
        cin >> a[i];
    }
    for (int i = 1; i < n; i++){
        int u, v;
        cin >> u >> v;
        adj[u].push_back(v);
        adj[v].push_back(u);
    }
    dfs_lca(1, 0);
    init_st();
    build(1, 0);
    for (int i = 1; i <= n; i++){
        for (int u = i; u; u = fa[u]){
            mxd1[u] = max(mxd1[u], get_dis(u, i));
            if (fa[u]){
                mxd2[u] = max(mxd2[u], get_dis(fa[u], i));
            }
        }
    }
    for (int i = 1; i <= n; i++){
        tr1[i].init(mxd1[i]);
        tr2[i].init(mxd2[i]);
    }
    for (int i = 1; i <= n; i++){
        modify_tree(i, a[i]);
    }
    int ans = 0;
    while (m--){
        int op;
        cin >> op;
        if (op == 0){
            int x, k;
            cin >> x >> k;
            x ^= ans;
            k ^= ans;
            ans = query_tree(x, k);
            cout << ans << "\n";
        }
        else{
            int x, y;
            cin >> x >> y;
            x ^= ans;
            y ^= ans;
            modify_tree(x, y - a[x]);
            a[x] = y;
        }
    }
    return 0;
}

:::

例题

I

:::info[例题 1]{open}

P10603 BZOJ4372 烁烁的游戏

给定一棵有 n 个节点的树,点带权。m 次操作,分为两种类型:

3s,n,m \le 10^5

:::

也就是把模板题的修改和查询反了过来。容易发现这并没有本质区别。对于修改,跳 u 的祖先,假设跳到了 p,那么与 p 的距离不超过 k-dis_{u,p} 的所有节点的权值均会受到影响。这相当于一段区间加,可用树状数组维护差分解决。对于查询,跳祖先累加贡献即可。

:::success[Code]

#include <bits/stdc++.h>
using namespace std;
const int MAXN = 1e5 + 10;
vector <int> adj[MAXN];
int a[MAXN], fa[MAXN], siz[MAXN], dep[MAXN], ord[MAXN * 2], in[MAXN], mxd1[MAXN], mxd2[MAXN], st[MAXN * 2][20], lg[MAXN * 2], n, m, cur;
bool vis[MAXN];
struct BIT{
    vector <int> sum;
    void init(int len){
        sum.resize(len + 2, 0);
        return;
    }
    int lowbit(int x){
        return x & (-x);
    }
    void modify(int u, int x){
        u++;
        while (u < (int)sum.size()){
            sum[u] += x;
            u += lowbit(u);
        }
        return;
    }
    int query(int u){
        u++;
        u = min(u, (int)sum.size() - 1);
        int res = 0;
        while (u){
            res += sum[u];
            u -= lowbit(u);
        }
        return res;
    }
}tr1[MAXN], tr2[MAXN];
void dfs_lca(int u, int p){
    dep[u] = dep[p] + 1;
    ord[++cur] = u;
    in[u] = cur;
    for (int v : adj[u]){
        if (v != p){
            dfs_lca(v, u);
            ord[++cur] = u;
        }
    }
}
int get_min(int u, int v){
    return (dep[u] < dep[v] ? u : v);
}
void init_st(){
    lg[1] = 0;
    for (int i = 2; i <= cur; i++){
        lg[i] = lg[i / 2] + 1;
    }
    for (int i = 1; i <= cur; i++){
        st[i][0] = ord[i];
    }
    for (int i = 1; i <= lg[cur]; i++){
        for (int j = 1; j <= cur - (1 << i) + 1; j++){
            st[j][i] = get_min(st[j][i - 1], st[j + (1 << (i - 1))][i - 1]);
        }
    }
    return;
}
int get_lca(int u, int v){
    int l = in[u], r = in[v];
    if (l > r){
        swap(l, r);
    }
    int k = lg[r - l + 1];
    return get_min(st[l][k], st[r - (1 << k) + 1][k]);
}
int get_dis(int u, int v){
    int lca = get_lca(u, v);
    return dep[u] + dep[v] - 2 * dep[lca];
}
void dfs1(int u, int p){
    siz[u] = 1;
    for (int v : adj[u]){
        if (v == p || vis[v]){
            continue;
        }
        dfs1(v, u);
        siz[u] += siz[v];
    }
    return;
}
void dfs2(int u, int p, int tot, int &rt){
    int mx = tot - siz[u];
    for (int v : adj[u]){
        if (v == p || vis[v]){
            continue;
        }
        dfs2(v, u, tot, rt);
        mx = max(mx, siz[v]);
    }
    if (mx * 2 <= tot){
        rt = u;
    }
    return;
}
void build(int u, int p){
    dfs1(u, 0);
    int rt = 0;
    dfs2(u, 0, siz[u], rt);
    vis[rt] = true;
    fa[rt] = p;
    for (int v : adj[rt]){
        if (!vis[v]){
            build(v, rt);
        }
    }
    return;
}
void modify_tree(int x, int d, int val){
    for (int u = x; u; u = fa[u]){
        if (d >= get_dis(u, x)){
            tr1[u].modify(0, val);
            tr1[u].modify(d - get_dis(u, x) + 1, -val);
        }
        if (fa[u] && d >= get_dis(fa[u], x)){
            tr2[u].modify(0, val);
            tr2[u].modify(d - get_dis(fa[u], x) + 1, -val);
        }
    }
    return;
}
int query_tree(int x){
    int sum = tr1[x].query(0);
    for (int u = x; fa[u]; u = fa[u]){
        int d = get_dis(fa[u], x);
        sum += tr1[fa[u]].query(d);
        sum -= tr2[u].query(d);
    }
    return sum;
}
signed main(){
    ios::sync_with_stdio(false);
    cin.tie(0);
    cout.tie(0);
    cin >> n >> m;
    for (int i = 1; i < n; i++){
        int u, v;
        cin >> u >> v;
        adj[u].push_back(v);
        adj[v].push_back(u);
    }
    dfs_lca(1, 0);
    init_st();
    build(1, 0);
    for (int i = 1; i <= n; i++){
        for (int u = i; u; u = fa[u]){
            mxd1[u] = max(mxd1[u], get_dis(u, i));
            if (fa[u]){
                mxd2[u] = max(mxd2[u], get_dis(fa[u], i));
            }
        }
    }
    for (int i = 1; i <= n; i++){
        tr1[i].init(mxd1[i]);
        tr2[i].init(mxd2[i]);
    }
    while (m--){
        char op;
        cin >> op;
        if (op == 'Q'){
            int x;
            cin >> x;
            cout << query_tree(x) << "\n";
        }
        else{
            int x, d, w;
            cin >> x >> d >> w;
            modify_tree(x, d, w);
        }
    }
    return 0;
}

:::

II

:::info[例题 2]{open}

P2056 [ZJOI2007] 捉迷藏

给定一棵有 n 个节点的树,点有黑色和白色两种颜色。初始时,所有节点均为黑色。q 次操作,分为两种类型:

5s,n \le 10^5q \le 5 \times 10^5

:::

称两端点均为黑点的路径为黑点路径。对于任意一条经过了节点 p 的黑点路径 (u,v),将它拆成 (p,u),(p,v) 两部分,则我们希望让这两部分分别取最大值和次大值。

要动态维护最大值和次大值,不难想到 multiset。我们对每个节点 u 开两个 multiset S_u, T_uS_u 维护 u 的子树内所有黑点到 fa_u 的距离,T_u 维护 u 的各个儿子 v\max(S_v),则 T_u 中最大值与次大值的和,就是经过 p 的最长黑点路径的长度。同时维护一个全局 multiset ans,存所有节点的最长路径长度。修改时跳祖先进行相应的更新,查询时取 ans 的最大值即可。

实现上需要注意,不要用 STL multiset,会被卡常,应用手写可删堆代替。

:::success[Code]

#include <bits/stdc++.h>
using namespace std;
const int MAXN = 1e5 + 10;
vector <int> adj[MAXN];
int fa[MAXN], siz[MAXN], dep[MAXN], ord[MAXN * 2], in[MAXN], st[MAXN * 2][20], lg[MAXN * 2], n, q, cur, cnt;
bool vis[MAXN], flag[MAXN];
struct Heap{
    priority_queue <int> q1, q2;
    void push(int x){
        q1.push(x);
        return;
    }
    void erase(int x){
        q2.push(x);
        return;
    }
    void clean(){
        while (!q2.empty() && q1.top() == q2.top()){
            q1.pop();
            q2.pop();
        }
        return;
    }
    int size(){
        return q1.size() - q2.size();
    }
    int top(){
        clean();
        return q1.empty() ? -1 : q1.top();
    }
    void pop(){
        clean();
        if (!q1.empty()){
            q1.pop();
        }
        return;
    }
    int get_two_max(){
        if (size() < 2){
            return -1;
        }
        int t1 = top();
        pop();
        int t2 = top();
        push(t1);
        return t1 + t2;
    }
}s1[MAXN], s2[MAXN], ans;
void dfs_lca(int u, int p){
    dep[u] = dep[p] + 1;
    ord[++cur] = u;
    in[u] = cur;
    for (int v : adj[u]){
        if (v != p){
            dfs_lca(v, u);
            ord[++cur] = u;
        }
    }
}
int get_min(int u, int v){
    return (dep[u] < dep[v] ? u : v);
}
void init_st(){
    lg[1] = 0;
    for (int i = 2; i <= cur; i++){
        lg[i] = lg[i / 2] + 1;
    }
    for (int i = 1; i <= cur; i++){
        st[i][0] = ord[i];
    }
    for (int i = 1; i <= lg[cur]; i++){
        for (int j = 1; j <= cur - (1 << i) + 1; j++){
            st[j][i] = get_min(st[j][i - 1], st[j + (1 << (i - 1))][i - 1]);
        }
    }
    return;
}
int get_lca(int u, int v){
    int l = in[u], r = in[v];
    if (l > r){
        swap(l, r);
    }
    int k = lg[r - l + 1];
    return get_min(st[l][k], st[r - (1 << k) + 1][k]);
}
int get_dis(int u, int v){
    int lca = get_lca(u, v);
    return dep[u] + dep[v] - 2 * dep[lca];
}
void dfs1(int u, int p){
    siz[u] = 1;
    for (int v : adj[u]){
        if (v == p || vis[v]){
            continue;
        }
        dfs1(v, u);
        siz[u] += siz[v];
    }
    return;
}
void dfs2(int u, int p, int tot, int &rt){
    int mx = tot - siz[u];
    for (int v : adj[u]){
        if (v == p || vis[v]){
            continue;
        }
        dfs2(v, u, tot, rt);
        mx = max(mx, siz[v]);
    }
    if (mx * 2 <= tot){
        rt = u;
    }
    return;
}
void build(int u, int p){
    dfs1(u, 0);
    int rt = 0;
    dfs2(u, 0, siz[u], rt);
    vis[rt] = true;
    fa[rt] = p;
    for (int v : adj[rt]){
        if (!vis[v]){
            build(v, rt);
        }
    }
    return;
}
int get_ans(int i){
    return s2[i].get_two_max();
}
void modify_tree(int i){
    if (!flag[i]){
        int tmp = get_ans(i);
        if (tmp != -1){
            ans.erase(tmp);
        }
        s2[i].push(0);
        tmp = get_ans(i);
        if (tmp != -1){
            ans.push(tmp);
        }
        for (int u = i; fa[u]; u = fa[u]){
            tmp = get_ans(fa[u]);
            if (tmp != -1){
                ans.erase(tmp);
            }
            if (s1[u].size() > 0){
                s2[fa[u]].erase(s1[u].top());
            }
            s1[u].push(get_dis(i, fa[u]));
            if (s1[u].size() > 0){
                s2[fa[u]].push(s1[u].top());
            }
            tmp = get_ans(fa[u]);
            if (tmp != -1){
                ans.push(tmp);
            }
        }
        flag[i] = true;
        cnt++;
    }
    else{
        int tmp = get_ans(i);
        if (tmp != -1){
            ans.erase(tmp);
        }
        s2[i].erase(0);
        tmp = get_ans(i);
        if (tmp != -1){
            ans.push(tmp);
        }
        for (int u = i; fa[u]; u = fa[u]){
            tmp = get_ans(fa[u]);
            if (tmp != -1){
                ans.erase(tmp);
            }
            if (s1[u].size() > 0){
                s2[fa[u]].erase(s1[u].top());
            }
            s1[u].erase(get_dis(i, fa[u]));
            if (s1[u].size() > 0){
                s2[fa[u]].push(s1[u].top());
            }
            tmp = get_ans(fa[u]);
            if (tmp != -1){
                ans.push(tmp);
            }
        }
        flag[i] = false;
        cnt--;
    }
    return;
}
signed main(){
    ios::sync_with_stdio(false);
    cin.tie(0);
    cout.tie(0);
    cin >> n;
    for (int i = 1; i < n; i++){
        int u, v;
        cin >> u >> v;
        adj[u].push_back(v);
        adj[v].push_back(u);
    }
    dfs_lca(1, 0);
    init_st();
    build(1, 0);
    for (int i = 1; i <= n; i++){
        modify_tree(i);
    }
    cin >> q;
    while (q--){
        char op;
        cin >> op;
        if (op == 'C'){
            int i;
            cin >> i;
            modify_tree(i);
        }
        else{
            if (cnt == 0){
                cout << "-1\n";
            }
            else if (cnt == 1){
                cout << "0\n";
            }
            else{
                cout << ans.top() << "\n";
            }
        }
    }
    return 0;
}

:::

III

:::info[例题 3]{open}

P3345 [ZJOI2015] 幻想乡战略游戏

给定一棵有 n 个节点的树,点带权,点 i 的权值为 d_i,初始时 d_i=0q 次操作,每次给定 u,e,表示令 d_u \leftarrow d_u+e。定义节点 u 的代价为 \sum_{1 \le v \le n} d_v \cdot dis(u,v)。每次操作后,求所有节点的最小代价。

6s,n,q \le 10^5,任意节点的度数不超过 20

:::

对于任意一个节点 u,钦定它为根,设 v 是它的一个儿子,siz_vv 的子树内所有节点的权值和,S 为整棵树上所有节点的权值和,w 为边 (u,v) 的边权,则 u,v 的代价之差 f(v)-f(u)=(S-siz_v) \cdot w-siz_v \cdot w。要使 v 的代价更小,就要使 siz_v>S/2。这意味着,代价比 u 更小的 v 至多只有 1 个。据此,可以初步得出一个做法:从任意一个节点出发,每步都走向相邻节点中唯一一个更优的节点,直到走不动为止,此时停下来的位置就是代价最小的位置。

直接在原树上做,最坏情况下要走 O(n) 步。把这个过程搬到点分树上,每次移动时直接跳到该节点所在连通块的分治中心,即可将步数控制在 O(\log n) 级别。再用点分树维护每个节点的代价,支持 O(\log n) 查询,即可做到 O(n \log^2 n \cdot \text{deg}),其中 \text{deg}=20。实际上根本跑不满,最大点才 600ms。

:::success[Code]

#include <bits/stdc++.h>
using namespace std;
typedef long long ll;
const int MAXN = 1e5 + 10;
const ll INF = 0x3f3f3f3f3f3f3f3f;
vector <pair<int, ll>> adj[MAXN], g[MAXN];
int n, q;
int dep[MAXN], ord[MAXN * 2], in[MAXN], lg[MAXN * 2], st[MAXN * 2][20], tim;
int siz[MAXN], fa[MAXN], cnt[MAXN], fi;
ll sum1[MAXN], sum2[MAXN], dis[MAXN];
bool vis[MAXN];
void init_lca(int u, int p){
    dep[u] = dep[p] + 1;
    ord[++tim] = u;
    in[u] = tim;
    for (auto [v, w] : adj[u]){
        if (v == p){
            continue;
        }
        dis[v] = dis[u] + w;
        init_lca(v, u);
        ord[++tim] = u;
    }
    return;
}
int get_min(int u, int v){
    return dep[u] < dep[v] ? u : v;
}
void init_st(){
    lg[1] = 0;
    for (int i = 2; i <= tim; i++){
        lg[i] = lg[i / 2] + 1;
    }
    for (int i = 1; i <= tim; i++){
        st[i][0] = ord[i];
    }
    for (int i = 1; i <= lg[tim]; i++){
        for (int j = 1; j <= tim - (1 << i) + 1; j++){
            st[j][i] = get_min(st[j][i - 1], st[j + (1 << (i - 1))][i - 1]);
        }
    }
    return;
}
int get_lca(int u, int v){
    int l = in[u], r = in[v];
    if (l > r){
        swap(l, r);
    }
    int k = lg[r - l + 1];
    return get_min(st[l][k], st[r - (1 << k) + 1][k]);
}
int get_dis(int u, int v){
    int lca = get_lca(u, v);
    return dis[u] + dis[v] - 2 * dis[lca];
}
void dfs1(int u, int p){
    siz[u] = 1;
    for (auto [v, w] : adj[u]){
        if (v == p || vis[v]){
            continue;
        }
        dfs1(v, u);
        siz[u] += siz[v];
    }
    return;
}
void dfs2(int u, int p, int tot, int &rt){
    int mx = tot - siz[u];
    for (auto [v, w] : adj[u]){
        if (v == p || vis[v]){
            continue;
        }
        dfs2(v, u, tot, rt);
        mx = max(mx, siz[v]);
    }
    if (mx * 2 <= tot){
        rt = u;
    }
    return;
}
int build(int u, int p){
    dfs1(u, 0);
    int rt = 0;
    dfs2(u, 0, siz[u], rt);
    fa[rt] = p;
    vis[rt] = true;
    for (auto [v, w] : adj[rt]){
        if (!vis[v]){
            int nxt = build(v, rt);
            g[rt].push_back({v, nxt});
        }
    }
    return rt;
}
void modify(int x, int val){
    cnt[x] += val;
    for (int u = x; fa[u]; u = fa[u]){
        int p = fa[u], d = get_dis(x, p);
        sum1[p] += 1ll * d * val;
        sum2[u] += 1ll * d * val;
        cnt[p] += val;
    }
    return;
}
ll query(int x){
    if (!x){
        return INF;
    }
    ll res = sum1[x];
    for (int u = x; fa[u]; u = fa[u]){
        int p = fa[u], d = get_dis(x, p);
        res += sum1[p] - sum2[u];
        res += 1ll * (cnt[p] - cnt[u]) * d;
    }
    return res;
}
ll find(int u){
    ll val = query(u);
    for (auto [v, nxt] : g[u]){
        if (query(v) < val){
            return find(nxt);
        }
    }
    return val;
}
signed main(){
    ios::sync_with_stdio(false);
    cin.tie(0);
    cout.tie(0);
    cin >> n >> q;
    for (int i = 1; i < n; i++){
        int u, v, w;
        cin >> u >> v >> w;
        adj[u].push_back({v, w});
        adj[v].push_back({u, w});
    }
    init_lca(1, 0);
    init_st();
    fi = build(1, 0);
    while (q--){
        int u, x;
        cin >> u >> x;
        modify(u, x);
        cout << find(fi) << "\n";
    }
    return 0;
}

:::

IV

:::info[例题 4]{open}

P17141 [NOI 2026] 传送

给定一棵有 n 个节点的树,节点编号为 0 \sim n-1。要求对每个节点 u 确定一个移动方式 a_u \in [0,n],若 a_u<n,则下一步移动到节点 a_u,若 a_u=n,则下一步等概率传送到 n 个节点之一。m 次询问,每次给定 x,y,求所有方案中从节点 x 移动到节点 y 的期望移动次数的最小值。

5s,n \le 5 \times 10^5m \le 10^6

:::

可以发现以下几条性质:

  1. 传送会覆盖掉之前的所有移动,故我们一定不会移动几次后再传送。这意味着,我们的策略一定是让终点所在的某个连通块内的所有节点移动,其余节点传送;
  2. 钦定终点为根,则同深度节点选择的策略一定相同。结合性质 1,这意味着,以某个阈值 dep 为界,深度 \le dep 的所有节点将选择移动,深度 >dep 的所有节点将选择传送;
  3. 传送一次后,期望步数就和起点没有任何关系了。这意味着,我们只需对每个终点 y 预处理传送一次后的期望步数 P,查询时将其与 dis_{x} 取 min 即可。

S 为选择移动的节点的集合,k=|S|,则有 P=1+\frac{(n-k)P+\sum_{u \in S} dis_{u}}{n},整理得 P=\frac{n+\sum_{u \in S} dis_{u}}{k}。对于一个固定的阈值 dep,它可能是最优解,当且仅当 P > dep。这很好理解,因为如果 P \le dep 的话,向外扩展一步一定更优。这个条件显然可以二分判定。点分树维护 \sum dis_u,k,可做到 O(n \log^2 n+m),可获得 [75,100] 分。实现上需要注意,本题没有修改,可用前缀和代替树状数组,且必须用 O(1) LCA,否则会变成 3log。

继续观察性质。可以发现,若确定了节点 uP_u,则任意一个与其相邻的节点 v 都满足 |P_u-P_v| \le 1,它们的阈值 dep 之差也不超过 1。故对一个节点进行一次 O(\log^2n) 的二分,然后 dfs 扩展即可。复杂度 O(n \log n+m)

:::success[Code]

#include <bits/stdc++.h>
using namespace std;
typedef long long ll;
const int MAXN = 5e5 + 10;
vector <int> adj[MAXN], pre3[MAXN], pre4[MAXN];
vector <ll> pre1[MAXN], pre2[MAXN];
int a[MAXN], fa[MAXN], siz[MAXN], dep[MAXN], mxd1[MAXN], mxd2[MAXN], cur;
int ord[MAXN * 2], in[MAXN], st[MAXN * 2][20], lg[MAXN * 2];
ll pa[MAXN];
int pb[MAXN], pd[MAXN];
bool vis[MAXN];
void dfs_lca(int u, int p){
    dep[u] = dep[p] + 1;
    ord[++cur] = u;
    in[u] = cur;
    for (int v : adj[u]){
        if (v != p){
            dfs_lca(v, u);
            ord[++cur] = u;
        }
    }
}
int get_min(int u, int v){
    return (dep[u] < dep[v] ? u : v);
}
void init_st(){
    lg[1] = 0;
    for (int i = 2; i <= cur; i++){
        lg[i] = lg[i / 2] + 1;
    }
    for (int i = 1; i <= cur; i++){
        st[i][0] = ord[i];
    }
    for (int i = 1; i <= lg[cur]; i++){
        for (int j = 1; j <= cur - (1 << i) + 1; j++){
            st[j][i] = get_min(st[j][i - 1], st[j + (1 << (i - 1))][i - 1]);
        }
    }
    return;
}
int get_lca(int u, int v){
    int l = in[u], r = in[v];
    if (l > r){
        swap(l, r);
    }
    int k = lg[r - l + 1];
    return get_min(st[l][k], st[r - (1 << k) + 1][k]);
}
int get_dis(int u, int v){
    int lca = get_lca(u, v);
    return dep[u] + dep[v] - 2 * dep[lca];
}
void dfs1(int u, int p){
    siz[u] = 1;
    for (int v : adj[u]){
        if (v == p || vis[v]){
            continue;
        }
        dfs1(v, u);
        siz[u] += siz[v];
    }
    return;
}
void dfs2(int u, int p, int tot, int &rt){
    int mx = tot - siz[u];
    for (int v : adj[u]){
        if (v == p || vis[v]){
            continue;
        }
        dfs2(v, u, tot, rt);
        mx = max(mx, siz[v]);
    }
    if (mx * 2 <= tot){
        rt = u;
    }
    return;
}
void build(int u, int p){
    dfs1(u, 0);
    int rt = 0;
    dfs2(u, 0, siz[u], rt);
    vis[rt] = true;
    fa[rt] = p;
    for (int v : adj[rt]){
        if (!vis[v]){
            build(v, rt);
        }
    }
    return;
}
void modify_tree(int x){
    for (int u = x; u; u = fa[u]){
        int d1 = get_dis(u, x), d2 = get_dis(fa[u], x);
        pre1[u][d1] += d1;
        pre3[u][d1]++;
        if (fa[u]){
            pre2[u][d2] += d2;
            pre4[u][d2]++;
        }
    }
    return;
}
pair <ll, int> query_tree(int x, int k){
    int tmp = min(k, (int)pre1[x].size() - 1);
    ll fi = pre1[x][tmp];
    int se = pre3[x][tmp];
    for (int u = x; fa[u]; u = fa[u]){
        int d = get_dis(fa[u], x);
        if (k >= d){
            int t1 = min(k - d, (int)pre1[fa[u]].size() - 1);
            int t2 = min(k - d, (int)pre2[u].size() - 1);
            ll sum = pre1[fa[u]][t1] - pre2[u][t2];
            int cnt = pre3[fa[u]][t1] - pre4[u][t2];
            fi += sum + 1ll * d * cnt;
            se += cnt;
        }
    }
    return {fi, se};
}
void dfs3(int u, int p, int d, int n){
    for (int v : adj[u]){
        if (v == p){
            continue;
        }
        for (int j = min(d + 1, n); j >= max(d - 1, 0); j--){
            pair <ll, int> res = query_tree(v, j);
            if (n + res.first > 1ll * j * res.second){
                pa[v] = n + res.first;
                pb[v] = res.second;
                pd[v] = j;
                dfs3(v, u, j, n);
                break;
            }
        }
    }
    return;
}
vector <pair<ll, int>> teleport(int c, int n, int m, vector <int> u, vector <int> v, vector <int> x, vector <int> y){
    for (int i = 0; i < n - 1; i++){
        int ui = u[i] + 1;
        int vi = v[i] + 1;
        adj[ui].push_back(vi);
        adj[vi].push_back(ui);
    }
    dfs_lca(1, 0);
    init_st();
    build(1, 0);
    for (int i = 1; i <= n; i++){
        for (int u = i; u; u = fa[u]){
            mxd1[u] = max(mxd1[u], get_dis(u, i));
            if (fa[u]){
                mxd2[u] = max(mxd2[u], get_dis(fa[u], i));
            }
        }
    }
    for (int i = 1; i <= n; i++){
        pre1[i].resize(mxd1[i] + 5, 0);
        pre2[i].resize(mxd2[i] + 5, 0);
        pre3[i].resize(mxd1[i] + 5, 0);
        pre4[i].resize(mxd2[i] + 5, 0);
    }
    for (int i = 1; i <= n; i++){
        modify_tree(i);
    }
    for (int i = 1; i <= n; i++){
        for (int j = 1; j < (int)pre1[i].size(); j++){
            pre1[i][j] += pre1[i][j - 1];
            pre3[i][j] += pre3[i][j - 1];
        }
        for (int j = 1; j < (int)pre2[i].size(); j++){
            pre2[i][j] += pre2[i][j - 1];
            pre4[i][j] += pre4[i][j - 1];
        }
    }
    int l = 0, r = n;
    while (r - l >= 5){
        int mid = (l + r) >> 1;
        pair <ll, int> res = query_tree(1, mid);
        if (n + res.first > 1ll * mid * res.second){
            l = mid;
        }
        else{
            r = mid;
        }
    }
    for (int j = r; j >= l; j--){
        pair <ll, int> res = query_tree(1, j);
        if (n + res.first > 1ll * j * res.second){
            pa[1] = n + res.first;
            pb[1] = res.second;
            pd[1] = j;
            break;
        }
    }
    dfs3(1, 0, pd[1], n);
    vector <pair<ll, int>> ans(m);
    for (int i = 0; i < m; i++){
        int xi = x[i] + 1;
        int yi = y[i] + 1;
        if (pa[yi] < 1ll * get_dis(xi, yi) * pb[yi]){
            ll t = gcd(pa[yi], pb[yi]);
            ans[i] = {pa[yi] / t, (int)(pb[yi] / t)};
        }
        else{
            ans[i] = {get_dis(xi, yi), 1};
        }
    }
    return ans;
}

:::