学习笔记:重链剖分

· · 算法·理论

引入

P3384 【模板】重链剖分 / 树链剖分

给定一棵有根树,需要支持以下操作:

如果只有子树加子树和,直接转到 dfs 序上就是区间加区间和,容易维护。

但是还有路径操作,如果用线段树维护,暴力做一次操作复杂度 O(n\log n),不可接受。

由于路径单独拿出来是一条链,考虑把树划分为若干条链,在每一条链上面修改。
可以将这若干条链放到一个序列上,让一条链上面的节点在连续的位置,然后使用维护序列的数据结构维护。

由于需要支持子树操作,需要让这个序列仍然是一个合法的 dfs 序,以支持子树操作。
那么每一条链都需要是后代到若干个祖先的形态。

例如下图就是一棵树的划分方式。

图中 1-3-6-9,4-8 组成两条链,2,5,7 单点分别看作一条链。

那么组成的一个合法序列对应节点就是 1,3,6,9,5,7,2,4,8。

这样链的操作就变成了若干个序列操作。

注意到这样转化后的操作的复杂度和路径被划分成的链的数量直接相关,所以需要找到一个合适的方式,使得一条路径被划分成的链的数量尽可能少。

首先如果一个非叶节点下面没有接其他的链,那么接上一条链肯定是不劣的,因为这样减少了子树内到根上的链的数量。

重链剖分

记 sz_u 为点 u 的子树大小(即子树内节点数量)。
记点 u 的重儿子是 u 子树大小最大的儿子,记为 son_u。

那么把 u-son_u-son_{son_u}-\dots 划分成一条链,直到到达叶子节点,这样的链记为一条重链,u 为这条链的链顶。
或者说,对于每个节点,它和它的重儿子连成重边,重边组成的极大链就是重链;链顶是重链中深度最小的节点。

记 (u,son_u) 是一条重边,(u,v)(v \ne son_u) 为一条轻边。

那么可以证明,任意一个节点到根的路径上最多有 O(\log_2 n) 条重链。

一个节点到根的路径上重链的数量显然为经过的轻边数量加一。

记 u 是一个链顶,那么 (fa_u,u) 是一条轻边。
所以 sz_u \leq sz_{son_{fa_u}},sz_{fa_u} \geq sz_u +sz_{son_{fa_u}}+1 > 2sz_u。
这样就有每个轻子树的大小都不超过父节点的一半。
换句话说,从 u 一直向父亲跳,每经过一条轻边,子树大小至少翻倍。

由于总结点数为 n,那么一个节点到根路径上的轻边数量最多会有 O(\log_2 n) 条。

这样,路径 u,v 上的重链数量 \leq O(\log n)+O(\log n)=O(\log n)。

所以一条路径就被重链剖分转化成了 O(\log n) 条重链,或者 dfs 序上的 O(\log n) 个区间。

以下给出代码实现。 :::info[Code]

#include<bits/stdc++.h>
using namespace std;
#define rep(a,b,c) for(int a=(b);a<=(c);++a)
const int N = 1e5 + 5;
int n, m, rt, mod;
int a[N];
vector<int> g[N];
int fa[N], sz[N], dep[N], son[N];
void dfs1(int u, int pr){
    fa[u] = pr;
    sz[u] = 1;
    dep[u] = dep[pr] + 1;
    for (int v : g[u]) if (v != pr){
        dfs1(v, u);
        sz[u] += sz[v];
        if (sz[son[u]] < sz[v]) son[u] = v;
    }
}
int top[N], dfn[N], nod[N], cnt;
void dfs2(int u, int tp){
    top[u] = tp;
    dfn[u] = ++cnt;
    nod[cnt] = u;
    if (!son[u]) return;
    dfs2(son[u], tp);
    for (int v : g[u]) if (v != fa[u] && v != son[u]) dfs2(v, v);
}
int main(){
    ios::sync_with_stdio(0);
    cin.tie(0), cout.tie(0);
    cin >> n >> m >> rt >> mod;
    rep(u, 1, n) cin >> a[u];
    rep(i, 2, n){
        int u, v;
        cin >> u >> v;
        g[u].push_back(v);
        g[v].push_back(u);
    }
    dfs1(rt, 0);
    dfs2(rt, rt);
    return 0;
}

::: 其中 fa[u] 为 u 的父节点,dep[u] 为 u 的深度(到根路径上的节点数量),son[u] 为 u 的重儿子编号,top[u] 为 u 所在重链的链顶,dfn[u] 为 u 在序列上的位置,nod[x] 为序列第 x 个位置上对应的节点编号。

根据上面的过程,可以发现重链剖分这个 O(\log n) 的上界是很严格的,并且只有特殊的数据才能卡满,例如完全二叉树。
:::info[一种常见的数据构造方式,能把树剖的复杂度卡的比较满,同时卡掉暴力] 建立一棵 \sqrt n 个结点的二叉树。
对于每个结点到其儿子的边,我们将其替换成一条长度为 \sqrt n 长度的链。

这样可以将随机询问轻重链切换次数卡到平均 \frac{\log_2 n}{2} 次,同时有 O(\sqrt n \log n) 的深度。

(本部分摘自 OI Wiki。)

于是需要区间加区间求和操作,这里使用线段树直接维护即可(代码中使用了标记永久化)。

这里讲一下路径操作的细节:

对于给定的两个节点 u,v,不断选择链顶深度更大的节点,统计 / 修改信息并向上跳到上一个重链,直到两个节点在同一条链上。
此时需要额外统计两个节点之间链的信息。

如图,top_u 深度更大,统计 [dfn_{top_u},dfn_u] 的信息,然后令 u \leftarrow fa_{top_u}。

P3384 完整实现: :::info[Code]

#include<bits/stdc++.h>
using namespace std;
#define rep(a,b,c) for(int a=(b);a<=(c);++a)
const int N = 1e5 + 5;
int n, m, rt, mod;
int a[N];
vector<int> g[N];
int fa[N], sz[N], dep[N], son[N];
void dfs1(int u, int pr){
    fa[u] = pr;
    sz[u] = 1;
    dep[u] = dep[pr] + 1;
    for (int v : g[u]) if (v != pr){
        dfs1(v, u);
        sz[u] += sz[v];
        if (sz[son[u]] < sz[v]) son[u] = v;
    }
}
int top[N], dfn[N], nod[N], cnt;
void dfs2(int u, int tp){
    top[u] = tp;
    dfn[u] = ++cnt;
    nod[cnt] = u;
    if (!son[u]) return;
    dfs2(son[u], tp);
    for (int v : g[u]) if (v != fa[u] && v != son[u]) dfs2(v, v);
}
struct node {
    int len, sum, tag;
    void add(int x){
        sum = (sum + 1ll * len * x) % mod;
        tag = (tag + x) % mod;
    }
} t[N<<2];
#define ls p<<1
#define rs p<<1|1
#define mid ((pl+pr)>>1)
void up(int p){
    t[p].sum = (t[ls].sum + t[rs].sum + 1LL * t[p].tag * t[p].len) % mod;
}
void build(int p, int pl, int pr){
    t[p].len = pr - pl + 1;
    if (pl == pr){
        t[p].sum = a[nod[pl]];
        return;
    }
    build(ls, pl, mid);
    build(rs, mid + 1, pr);
    up(p);
}
void add(int p, int pl, int pr, int l, int r, int x){
    if (l <= pl && pr <= r) return t[p].add(x);
    if (l <= mid) add(ls, pl, mid, l, r, x);
    if (r > mid) add(rs, mid + 1, pr, l, r, x);
    up(p);
}
int sum(int p, int pl, int pr, int l, int r, int tag){
    if (l <= pl && pr <= r) return(t[p].sum + 1LL * t[p].len * tag) % mod;
    tag = (tag + t[p].tag) % mod;
    int res = 0;
    if (l <= mid) res += sum(ls, pl, mid, l, r, tag);
    if (r > mid) res += sum(rs, mid + 1, pr, l, r, tag);
    return res % mod;
}
void add_subtree(int u, int x){
    add(1, 1, n, dfn[u], dfn[u] + sz[u] - 1, x);
}
void add_path(int u, int v, int x){
    while (top[u] != top[v]){
        if (dep[top[u]] < dep[top[v]]) swap(u, v);
        add(1, 1, n, dfn[top[u]], dfn[u], x);
        u = fa[top[u]];
    }
    if (dep[u] > dep[v]) swap(u, v);
    add(1, 1, n, dfn[u], dfn[v], x);
}
int sum_subtree(int u){
    return sum(1, 1, n, dfn[u], dfn[u] + sz[u] - 1, 0);
}
int sum_path(int u, int v){
    int res = 0;
    while (top[u] != top[v]){
        if (dep[top[u]] < dep[top[v]]) swap(u, v);
        res = (res + sum(1, 1, n, dfn[top[u]], dfn[u], 0)) % mod;
        u = fa[top[u]];
    }
    if (dep[u] > dep[v]) swap(u, v);
    res = (res + sum(1, 1, n, dfn[u], dfn[v], 0)) % mod;
    return res;
}
int main(){
    ios::sync_with_stdio(0);
    cin.tie(0), cout.tie(0);
    cin >> n >> m >> rt >> mod;
    rep(u, 1, n) cin >> a[u];
    rep(i, 2, n){
        int u, v;
        cin >> u >> v;
        g[u].push_back(v);
        g[v].push_back(u);
    }
    dfs1(rt, 0);
    dfs2(rt, rt);
    build(1, 1, n);
    while (m--){
        int op, u, v, x;
        cin >> op;
        if (op == 1){
            cin >> u >> v >> x;
            add_path(u, v, x);
        }
        else if (op == 2){
            cin >> u >> v;
            cout << sum_path(u, v) <<  '\n';
        }
        else if (op == 3){
            cin >> u >> x;
            add_subtree(u, x);
        }
        else if (op == 4){
            cin >> u;
            cout << sum_subtree(u) <<  '\n';
        }
    }
    return 0;
}

:::

分析一下复杂度:

这里使用线段树维护区间加区间求和,算上重链剖分的辅助数组,空间复杂度是 O(n) 的。

预处理的 dfs 部分复杂度是线性的,为 O(n)。

线段树建树同样是 O(n)。

对于一次子树操作,只需要对子树对应的 dfs 序做区间修改 / 查询,只需要在线段树对应的 [dfn_u,dfn_u+sz_u-1] 区间上操作即可,复杂度是线段树的 O(\log n)。

对于一次路径操作,路径被分为了 O(\log n) 条重链,每条重链进行一次区间操作,复杂度是 O(\log^2n) 的。

总复杂度可以看作 O(n + m\log^2 n)。

重链剖分求解最近公共祖先

容易发现,在上述查询过程中,u,v 跳到同一条重链上之后,深度较小的节点就是 \operatorname{lca}(u,v)。

复杂度为跳过的重链数量,为 O(\log n)。

如果只需要求 \operatorname{lca}(u,v),代码可以精简一些,可以参考下面给出的 P3379 【模板】最近公共祖先(LCA) 实现。

时间复杂度 O(n+m \log n),空间复杂度 O(n)。 :::info[Code]

#include<bits/stdc++.h>
using namespace std;
const int N = 5e5 + 5;
int n, m, rt;
vector<int> g[N];
int fa[N], dep[N], sz[N], son[N];
void dfs1(int u, int pr){
    fa[u] = pr;
    dep[u] = dep[pr] + 1;
    sz[u] = 1;
    for (int v : g[u]) if (v != pr){
        dfs1(v, u);
        sz[u] += sz[v];
        if (sz[son[u]] < sz[v]) son[u] = v;
    }
}
int top[N];
void dfs2(int u, int tp){
    top[u] = tp;
    for (int v : g[u]) if (v != fa[u]) dfs2(v, v == son[u] ? tp : v);
}
int lca(int u, int v){
    while (top[u] != top[v]) dep[top[u]] > dep[top[v]] ? u = fa[top[u]] : v = fa[top[v]];
    return dep[u] < dep[v] ? u : v;
}
int main(){
    ios::sync_with_stdio(0);
    cin.tie(0), cout.tie(0);
    cin >> n >> m >> rt;
    for (int i = 1; i < n; i++){
        int u, v;
        cin >> u >> v;
        g[u].push_back(v);
        g[v].push_back(u);
    }
    dfs1(rt, 0);
    dfs2(rt, rt);
    while (m--){
        int u, v;
        cin >> u >> v;
        cout << lca(u, v) <<  '\n';
    }
    return 0;
}

:::

例题

P2590 [ZJOI2008] 树的统计

Description

给出一棵树,每个点有点权。
需要支持三种操作:

直接做是简单的,但是这里介绍一种减小常数的技巧:

当只有路径操作时,数据结构可以单独维护每一条链,而不是所有节点。
例如线段树,可以动态开点,每一条重链开一棵线段树,只需要预处理的时候额外记录链底即可。

Code

:::info[Code]

#include<bits/stdc++.h>
using namespace std;
#define rep(a,b,c) for(int a=(b);a<=(c);++a)
const int N = 3e4 + 5;
int n, m;
int a[N];
vector<int> g[N];
int fa[N], sz[N], dep[N], son[N];
void dfs1(int u, int pr){
    fa[u] = pr;
    sz[u] = 1;
    dep[u] = dep[pr] + 1;
    for (int v : g[u]) if (v != pr){
        dfs1(v, u);
        sz[u] += sz[v];
        if (sz[son[u]] < sz[v]) son[u] = v;
    }
}
struct node {
    int lson, rson;
    int sum, mx;
} t[N<<2];
int rt[N], ct;
int top[N], bot[N], dfn[N], nod[N], cnt;
#define ls t[p].lson
#define rs t[p].rson
#define mid ((pl+pr)>>1)
void up(int p){
    t[p].sum = t[ls].sum + t[rs].sum;
    t[p].mx = max(t[ls].mx, t[rs].mx);
}
void build(int & p, int pl, int pr){
    p = ++ct;
    if (pl == pr){
        t[p].sum = t[p].mx = a[nod[pl]];
        return;
    }
    build(ls, pl, mid);
    build(rs, mid + 1, pr);
    up(p);
}
void upd(int p, int pl, int pr, int x, int v){
    if (pl == pr){
        t[p].sum = t[p].mx = v;
        return;
    }
    x <= mid ? upd(ls, pl, mid, x, v) : upd(rs, mid + 1, pr, x, v);
    up(p);
}
int sum(int p, int pl, int pr, int l, int r){
    if (l <= pl && pr <= r) return t[p].sum;
    int res = 0;
    if (l <= mid) res += sum(ls, pl, mid, l, r);
    if (r > mid) res += sum(rs, mid + 1, pr, l, r);
    return res;
}
int mx(int p, int pl, int pr, int l, int r){
    if (l <= pl && pr <= r) return t[p].mx;
    int res =-1e9;
    if (l <= mid) res = max(res, mx(ls, pl, mid, l, r));
    if (r > mid) res = max(res, mx(rs, mid + 1, pr, l, r));
    return res;
}
void dfs2(int u, int tp){
    top[u] = tp;
    dfn[u] = ++cnt;
    bot[u] = u;
    nod[cnt] = u;
    if (!son[u]) return build(rt[tp], dfn[tp], dfn[u]);
    dfs2(son[u], tp);
    bot[u] = bot[son[u]];
    for (int v : g[u]) if (v != fa[u] && v != son[u]) dfs2(v, v);
}
void change(int u, int x){
    upd(rt[top[u]], dfn[top[u]], dfn[bot[u]], dfn[u], x);
}
int qmax(int u, int v){
    int res =-1e9;
    while (top[u] != top[v]){
        if (dep[top[u]] < dep[top[v]]) swap(u, v);
        res = max(res, mx(rt[top[u]], dfn[top[u]], dfn[bot[u]], dfn[top[u]], dfn[u]));
        u = fa[top[u]];
    }
    if (dep[u] > dep[v]) swap(u, v);
    res = max(res, mx(rt[top[u]], dfn[top[u]], dfn[bot[u]], dfn[u], dfn[v]));
    return res;
}
int qsum(int u, int v){
    int res = 0;
    while (top[u] != top[v]){
        if (dep[top[u]] < dep[top[v]]) swap(u, v);
        res += sum(rt[top[u]], dfn[top[u]], dfn[bot[u]], dfn[top[u]], dfn[u]);
        u = fa[top[u]];
    }
    if (dep[u] > dep[v]) swap(u, v);
    res += sum(rt[top[u]], dfn[top[u]], dfn[bot[u]], dfn[u], dfn[v]);
    return res;
}
int main(){
    ios::sync_with_stdio(0);
    cin.tie(0), cout.tie(0);
    cin >> n;
    rep(i, 2, n){
        int u, v;
        cin >> u >> v;
        g[u].push_back(v);
        g[v].push_back(u);
    }
    rep(u, 1, n) cin >> a[u];
    dfs1(1, 0);
    dfs2(1, 1);
    cin >> m;
    while (m--){
        string op;
        int x, y;
        cin >> op >> x >> y;
        if (op ==  "CHANGE") change(x, y);
        else if (op ==  "QMAX") cout << qmax(x, y) <<  '\n';
        else if (op ==  "QSUM") cout << qsum(x, y) <<  '\n';
    }
    return 0;
}

:::

P4114 Qtree1

Description

给出一棵树,每条边有边权。
需要支持两种操作:

Solution

介绍一个技巧:边转点。

由于每个节点的父节点唯一,所以可以把一条边的边权转成深度较大的一个端点的点权,这样就可以直接做了。
特别地,根节点的权值没有定义。

需要注意的是,边转点之后路径操作不能涉及到 \operatorname{lca}(u,v) 的点权。

Code

:::info[Code]

#include<bits/stdc++.h>
using namespace std;
#define rep(a,b,c) for(int a=(b);a<=(c);++a)
const int N = 1e5 + 5;
int n, m;
int eu[N], ev[N];
vector<pair<int, int>> g[N];
int fa[N], sz[N], dep[N], son[N], wt[N];
void dfs1(int u, int pr){
    fa[u] = pr;
    sz[u] = 1;
    dep[u] = dep[pr] + 1;
    for (auto[v, w] : g[u]) if (v != pr){
        wt[v] = w;
        dfs1(v, u);
        sz[u] += sz[v];
        if (sz[son[u]] < sz[v]) son[u] = v;
    }
}
struct node {
    int lson, rson;
    int mx;
} t[N<<2];
int rt[N], ct;
int top[N], bot[N], dfn[N], nod[N], cnt;
#define ls t[p].lson
#define rs t[p].rson
#define mid ((pl+pr)>>1)
void up(int p){
    t[p].mx = max(t[ls].mx, t[rs].mx);
}
void build(int & p, int pl, int pr){
    p = ++ct;
    if (pl == pr){
        t[p].mx = wt[nod[pl]];
        return;
    }
    build(ls, pl, mid);
    build(rs, mid + 1, pr);
    up(p);
}
void upd(int p, int pl, int pr, int x, int v){
    if (pl == pr){
        t[p].mx = v;
        return;
    }
    x <= mid ? upd(ls, pl, mid, x, v) : upd(rs, mid + 1, pr, x, v);
    up(p);
}
int mx(int p, int pl, int pr, int l, int r){
    if (l <= pl && pr <= r) return t[p].mx;
    int res =-1e9;
    if (l <= mid) res = max(res, mx(ls, pl, mid, l, r));
    if (r > mid) res = max(res, mx(rs, mid + 1, pr, l, r));
    return res;
}
void dfs2(int u, int tp){
    top[u] = tp;
    dfn[u] = ++cnt;
    bot[u] = u;
    nod[cnt] = u;
    if (!son[u]) return build(rt[tp], dfn[tp], dfn[u]);
    dfs2(son[u], tp);
    bot[u] = bot[son[u]];
    for (auto [v, w] : g[u]) if (v != fa[u] && v != son[u]) dfs2(v, v);
}
void change(int i, int x){
    int u = dep[eu[i]] > dep[ev[i]] ? eu[i] : ev[i];
    upd(rt[top[u]], dfn[top[u]], dfn[bot[u]], dfn[u], x);
}
int qmax(int u, int v){
    int res = 0;
    while (top[u] != top[v]){
        if (dep[top[u]] < dep[top[v]]) swap(u, v);
        res = max(res, mx(rt[top[u]], dfn[top[u]], dfn[bot[u]], dfn[top[u]], dfn[u]));
        u = fa[top[u]];
    }
    if (dep[u] > dep[v]) swap(u, v);
    if (u != v) res = max(res, mx(rt[top[u]], dfn[top[u]], dfn[bot[u]], dfn[u] + 1, dfn[v]));
    return res;
}
int main(){
    ios::sync_with_stdio(0);
    cin.tie(0), cout.tie(0);
    cin >> n;
    rep(i, 1, n - 1){
        int u, v, w;
        cin >> u >> v >> w;
        g[u].push_back({v, w});
        g[v].push_back({u, w});
        eu[i] = u, ev[i] = v;
    }
    dfs1(1, 0);
    dfs2(1, 1);
    string op;
    while (cin >> op){
        if (op ==  "DONE") break;
        int x, y;
        cin >> x >> y;
        if (op ==  "CHANGE") change(x, y);
        else if (op ==  "QUERY") cout << qmax(x, y) <<  '\n';
    }
    return 0;
}

:::

P4116 Qtree3

Description

给出一棵 n 个点的树,节点有黑白两种颜色,初始全为白。

有两种操作:

Solution

看作维护一个点集,每次查询 1 到 u 路径上最浅的节点的编号。

考虑分成若干条重链,查询每条重链上最浅的节点的编号。

可以直接维护一个以 dfn_u 为键值的 set,每次直接在链上使用 s.lower_bound(dfn[top[u]]) 查询即可。

但是发现如果每条重链开一个 set,每次只需要取出 s[top[u]].begin(),判断其是否在 [dfn_{top_u},dfn_u] 内即可,单次询问复杂度 O(\log n)。

时间复杂度 O(n+m\log n),空间复杂度 O(n)。

Code

:::info[Code]

#include<bits/stdc++.h>
using namespace std;
#define rep(a,b,c) for(int a=(b);a<=(c);++a)
const int N = 1e5 + 5;
int n, q;
int a[N];
vector<int> g[N];
int fa[N], sz[N], dep[N], son[N];
void dfs1(int u, int pr){
    fa[u] = pr;
    sz[u] = 1;
    dep[u] = dep[pr] + 1;
    for (int v : g[u]) if (v != pr){
        dfs1(v, u);
        sz[u] += sz[v];
        if (sz[son[u]] < sz[v]) son[u] = v;
    }
}
int top[N];
void dfs2(int u, int tp){
    top[u] = tp;
    for (int v : g[u]) if (v != fa[u]) dfs2(v, v == son[u] ? tp : v);
}
set<pair<int, int>> s[N];
void upd(int u){
    if (a[u]) s[top[u]].erase({dep[u], u});
    else s[top[u]].insert({dep[u], u});
    a[u] ^= 1;
}
int query(int u){
    int res =-1;
    while (u){
        if (!s[top[u]].empty() && (*s[top[u]].begin()).first <= dep[u]) res = (*s[top[u]].begin()).second;
        u = fa[top[u]];
    }
    return res;
}
int main(){
    ios::sync_with_stdio(0);
    cin.tie(0), cout.tie(0);
    cin >> n >> q;
    rep(i, 2, n){
        int u, v;
        cin >> u >> v;
        g[u].push_back(v);
        g[v].push_back(u);
    }
    dfs1(1, 0);
    dfs2(1, 1);
    while (q--){
        int op, u;
        cin >> op >> u;
        if (!op) upd(u);
        else cout << query(u) <<  '\n';
    }
    return 0;
}

:::

P3979 遥远的国度

Description

给定一棵树,有点权,支持以下操作:

Solution

路径赋值是简单的,考虑如何维护换根操作。

重点在于如何求出以 rt 为根时,点 u 子树对应的节点是哪些。

如图,考虑 rt 的位置不同时,u 的子树对应哪些节点。
(下文中祖先关系都是在以 1 为根的情况下讨论。)

如果 rt \not \in subtree_u,那么 rt 为根时 u 的子树还是 subtree_u,对应 dfs 序为 [dfn_u,dfn_u+sz_u)。
如果 rt=u,那么 u 的子树就是整棵树,对应 dfs 序为 [1,n]。
否则,记 rt 祖先中 u 的儿子节点为 v,那么 rt 为根时 u 的子树为 T \setminus subtree_v,即扣掉 v 子树之后的所有点,对应 dfs 序为 [1,dfn_v) \cup [dfn_v+sz_v,n]。

发现这样仍然可以用一棵全局的线段树维护,于是只需要考虑 rt 在 u 子树内时如何找到 v。

这里提供三种方法。
第一种是长剖找 rt 的第 dep_{rt}-dep_u+1 级祖先,但是我不会,所以不做介绍。
第二种是利用树剖重链 dfs 序连续的特性直接向上跳,到了同一条链上可以 O(1) 访问 k 级祖先,复杂度 O(\log n)。
第三种是把存边的 vector 只保留每个点连向儿子的边,然后把重儿子放到第一个再继续 dfs,访问时每个儿子的 dfn_v 自然是递增的,可以直接二分查找最后一个 dfn_v \leq dfn_{rt} 的儿子,复杂度同样 O(\log n)。

预处理复杂度 O(n),修改复杂度 O(\log ^2n),查询复杂度 O(\log n)。
空间线性。

Code

第二种扩展性强一点,这里给出二三两种方式实现的代码。
:::info[Code]

#include<bits/stdc++.h>
using namespace std;
#define ll long long
#define rep(a,b,c) for(int a=(b);a<=(c);++a)
const int N = 1e5 + 5, INF = 0x7fffffff;
int n, m, rt;
int a[N];
vector<int> g[N];
int fa[N], dep[N], sz[N], son[N];
void dfs1(int u, int pr){
    fa[u] = pr;
    sz[u] = 1;
    dep[u] = dep[pr] + 1;
    if (pr) g[u].erase(find(g[u].begin(), g[u].end(), pr));  //sol 2
    for (int v : g[u]){
        dfs1(v, u);
        sz[u] += sz[v];
        if (sz[son[u]] < sz[v]) son[u] = v;
    }
}
int top[N], dfn[N], nod[N], cnt;
void dfs2(int u, int tp){
    top[u] = tp;
    dfn[u] = ++cnt;
    nod[cnt] = u;
    rep(i, 0, (int)g[u].size() - 1) if (g[u][i] == son[u]) swap(g[u][0], g[u][i]);  //sol 2
    if (!son[u]) return;
    dfs2(son[u], tp);
    for (int v : g[u]) if (v != son[u]) dfs2(v, v);
}
int find1(int x, int u){//sol 1
    int dept = dep[u] + 1;
    while (dep[top[x]] > dept) x = fa[top[x]];
    return nod[dfn[x] - dep[x] + dept];
}
int find2(int x, int u){//sol 2
    return g[u][upper_bound(g[u].begin(), g[u].end(), x, [](int x, int y){return dfn[x] < dfn[y];}) - g[u].begin() - 1];
}
struct node {
    int mn, tag;
    void cov(int x){
        mn = tag = x;
    }
} t[N<<2];
#define ls p<<1
#define rs p<<1|1
#define mid ((pl+pr)>>1)
void up(int p){
    t[p].mn = min(t[ls].mn, t[rs].mn);
}
void down(int p){
    if (t[p].tag !=-1){
        t[ls].cov(t[p].tag);
        t[rs].cov(t[p].tag);
        t[p].tag =-1;
    }
}
void build(int p, int pl, int pr){
    t[p].tag =-1;
    if (pl == pr){
        t[p].mn = a[nod[pl]];
        return;
    }
    build(ls, pl, mid);
    build(rs, mid + 1, pr);
    up(p);
}
void upd(int p, int pl, int pr, int l, int r, int x){
    if (l <= pl && pr <= r) return t[p].cov(x);
    down(p);
    if (l <= mid) upd(ls, pl, mid, l, r, x);
    if (r > mid) upd(rs, mid + 1, pr, l, r, x);
    up(p);
}
int query(int p, int pl, int pr, int l, int r){
    if (l > r) return INF;
    if (l <= pl && pr <= r) return t[p].mn;
    down(p);
    int res = INF;
    if (l <= mid) res = min(res, query(ls, pl, mid, l, r));
    if (r > mid) res = min(res, query(rs, mid + 1, pr, l, r));
    return res;
}
void updt(int u, int v, int x){
    while (top[u] != top[v]){
        if (dep[top[u]] < dep[top[v]]) swap(u, v);
        upd(1, 1, n, dfn[top[u]], dfn[u], x);
        u = fa[top[u]];
    }
    if (dep[u] > dep[v]) swap(u, v);
    upd(1, 1, n, dfn[u], dfn[v], x);
}
int qry(int u, int x){
    if (u == x) return t[1].mn;
    if (dfn[x] < dfn[u] || dfn[u] + sz[u] <= dfn[x]) return query(1, 1, n, dfn[u], dfn[u] + sz[u] - 1);
    int v = find2(x, u); // find1(x,u)
    return min(query(1, 1, n, 1, dfn[v] - 1), query(1, 1, n, dfn[v] + sz[v], n));
}
int main(){
    ios::sync_with_stdio(0);
    cin.tie(0), cout.tie(0);
    cin >> n >> m;
    rep(i, 2, n){
        int u, v;
        cin >> u >> v;
        g[u].push_back(v);
        g[v].push_back(u);
    }
    rep(u, 1, n) cin >> a[u];
    dfs1(1, 0);
    dfs2(1, 1);
    build(1, 1, n);
    cin >> rt;
    while (m--){
        int op, u, v, x;
        cin >> op;
        if (op == 1) cin >> rt;
        else if (op == 2) cin >> u >> v >> x, updt(u, v, x);
        else if (op == 3) cin >> u, cout << qry(u, rt) <<  '\n';
    }
    return 0;
}

:::

P4211 [LNOI2014] LCA

Description

给出一棵以 1 为根的有根树(原题题面为 0-indexed),定义一个点的深度为这个点到根的路径上结点的数量。

#### Solution 首先 $\sum_{i=l}^r dep_{\operatorname{lca}(u,i)}$ 可以拆成前缀和相减,变成 $\sum_{i=1}^r dep_{\operatorname{lca}(u,i)}-\sum_{i=1}^{l-1} dep_{\operatorname{lca}(u,i)}$。 于是按照编号顺序把询问离线到节点上面,只需要解决若干个形如 $\sum_{i=1}^x dep_{\operatorname{lca}(u,i)}$ 的询问。 由于一个点的深度为这个点到根的路径上结点的数量,那么可以考虑 $1$ 到 $u$ 路径上的节点 $x$ 对 $dep_{\operatorname{lca}(u,i)}$ 的贡献。 根据定义,$x$ 的贡献就是 $x$ 子树内 $i$ 的数量。 于是每当 $i\leftarrow i'+1$,就把 $1$ 到 $i$ 的路径上的点权加 $1$,每个 $u$ 的查询就是 $u$ 到 $1$ 上面的点权和。 可以直接树剖维护。 由于只有重链上前缀加前缀和,可以对每一个重链开一个 BIT,不过有点难写,可以自己实现。 这里给出每个重链开一棵线段树的解法。 时间复杂度 $O((n+m)\log^2 n)$,空间复杂度 $O(n+m)$。 #### Code :::info[Code] ```cpp #include<bits/stdc++.h> using namespace std; #define rep(a,b,c) for(int a=(b);a<=(c);++a) const int N = 5e4 + 5, P = 201314; int n, m; vector<int> g[N]; int fa[N], sz[N], dep[N], son[N]; void dfs1(int u){ sz[u] = 1; dep[u] = dep[fa[u]] + 1; for (int v : g[u]){ dfs1(v); sz[u] += sz[v]; if (sz[son[u]] < sz[v]) son[u] = v; } } int top[N], dfn[N], bot[N], nod[N], cnt; struct node { int lson, rson; int len, sum, tag; void add(){ sum = (sum + len) % P; ++tag; } } t[N<<2]; int rt[N], ct; #define ls t[p].lson #define rs t[p].rson #define mid ((pl+pr)>>1) void up(int p){ t[p].sum = (t[ls].sum + t[rs].sum + 1LL * t[p].tag * t[p].len) % P; } void build(int & p, int pl, int pr){ p = ++ct; t[p].len = pr - pl + 1; if (pl == pr) return; build(ls, pl, mid); build(rs, mid + 1, pr); up(p); } void add(int p, int pl, int pr, int l, int r){ if (l <= pl && pr <= r) return t[p].add(); if (l <= mid) add(ls, pl, mid, l, r); if (r > mid) add(rs, mid + 1, pr, l, r); up(p); } int sum(int p, int pl, int pr, int l, int r, int tag){ if (l <= pl && pr <= r) return(t[p].sum + 1LL * t[p].len * tag) % P; tag = (tag + t[p].tag) % P; int res = 0; if (l <= mid) res += sum(ls, pl, mid, l, r, tag); if (r > mid) res += sum(rs, mid + 1, pr, l, r, tag); return res % P; } void dfs2(int u, int tp){ top[u] = tp; dfn[u] = ++cnt; nod[cnt] = u; bot[u] = u; if (!son[u]) return build(rt[tp], dfn[tp], dfn[u]); dfs2(son[u], tp); bot[u] = bot[son[u]]; for (int v : g[u]) if (v != fa[u] && v != son[u]) dfs2(v, v); } void add_path(int u){ while (u){ add(rt[top[u]], dfn[top[u]], dfn[bot[u]], dfn[top[u]], dfn[u]); u = fa[top[u]]; } } int sum_path(int u){ int res = 0; while (top[u]){ res = (res + sum(rt[top[u]], dfn[top[u]], dfn[bot[u]], dfn[top[u]], dfn[u], 0)) % P; u = fa[top[u]]; } return res; } struct query { int u, id, typ; }; vector<query> qry[N]; int ans[N]; int main(){ ios::sync_with_stdio(0); cin.tie(0), cout.tie(0); cin >> n >> m; rep(u, 2, n){ cin >> fa[u]; ++fa[u]; g[fa[u]].push_back(u); } dfs1(1); dfs2(1, 1); rep(i, 1, m){ int l, r, u; cin >> l >> r >> u; ++l, ++r, ++u; qry[r].push_back({u, i, 1}); qry[l - 1].push_back({u, i,-1}); } rep(u, 1, n){ add_path(u); for (auto [v, id, x] : qry[u]) ans[id] += x * sum_path(v); } rep(i, 1, m) cout << (ans[i] % P + P) % P << '\n'; return 0; } ``` ::: ### [P2486 [SDOI2011] 染色](https://www.luogu.com.cn/problem/P2486) #### Description 给定一棵 $n$ 个节点的无根树,每个点有点权,共有 $m$ 个操作,分为两种类型: 1. 路径赋值。 2. 查询路径上点权的极长相等段数量。 #### Solution 考虑路径上每一对相邻位置。 对于两个相邻的位置,如果这两个位置的颜色相同,那么答案不变。 否则段数加一。 考虑维护路径上相邻异色对数量。 可以使用线段树维护。 考虑如何合并两个节点。 发现合并时只会添加区间中点位置的一对颜色,于是对于每一个区间额外维护左右端点颜色,合并时相邻异色对数量加上中点贡献即可。 由于有路径覆盖,拆到重链上就是区间覆盖。 对于目前节点维护的信息来说,覆盖是简单的,更新左右端点颜色,并将相邻异色对数量设置为 $0$ 即可。 由于只有路径操作,可以每条重链开一个线段树。 查询有一些细节。 首先由于维护的信息比较复杂,需要分别维护 $u$ 向上部分的信息和 $v$ 向上部分的信息。 向上跳的时候,把链顶深度大的拿出来并更新信息。 由于线段树维护的是 dfs 序区间,而一条重链的 dfs 序随深度递增,所以向上跳的过程中,查询区间 $[l,r]$ 是单调左移并且两两不交的。 所以应该让旧的信息作为右区间,查询到的信息作为左区间执行合并。 最后是两点在同一条重链上的部分。 让深度大的节点向上跳,并更新这个节点的信息即可。 这样就得到了 $path[\operatorname{lca}(u,v),u]$ 和 $path(\operatorname{lca}(u,v),v]$ 的信息。 显然此时如果直接合并方向是错完了的。 所以应该翻转 $u$ 的信息,然后再将 $u$ 的信息作为左区间和 $v$ 的信息合并,得到 $path(u,v)$ 上面相邻异色对数。 信息的翻转也是容易的,直接交换节点区间左右端点颜色即可。 理论上也可以写标记永久化,但是有点麻烦,这里没有写。 时间复杂度 $O(n+m \log^2n)$,空间复杂度 $O(n)$。 #### Code :::info[Code] ```cpp #include<bits/stdc++.h> using namespace std; #define rep(a,b,c) for(int a=(b);a<=(c);++a) const int N = 1e5 + 5; int n, m; int a[N]; vector<int> g[N]; int fa[N], sz[N], dep[N], son[N]; void dfs1(int u, int pr){ fa[u] = pr; sz[u] = 1; dep[u] = dep[pr] + 1; for (int v : g[u]) if (v != pr){ dfs1(v, u); sz[u] += sz[v]; if (sz[son[u]] < sz[v]) son[u] = v; } } struct node { int lson, rson; int lc, rc, sum, tag; void cov(int x){ lc = rc = tag = x; sum = 0; } void flip(){ swap(lc, rc); } node & operator = (const node & x){ lc = x.lc, rc = x.rc, sum = x.sum, tag = x.tag; return * this; } } t[N<<2]; int rt[N], ct; int top[N], bot[N], dfn[N], nod[N], cnt; #define ls t[p].lson #define rs t[p].rson #define mid ((pl+pr)>>1) void merge(node & x, const node l, const node r){ if (!l.lc){x = r; return;} if (!r.rc){x = l; return;} x.lc = l.lc, x.rc = r.rc; x.sum = l.sum + r.sum + (l.rc != r.lc); x.tag = 0; } void up(int p){ merge(t[p], t[ls], t[rs]); } void down(int p){ if (t[p].tag){ t[ls].cov(t[p].tag); t[rs].cov(t[p].tag); t[p].tag = 0; } } void build(int & p, int pl, int pr){ p = ++ct; t[p].tag = 0; if (pl == pr){ t[p].lc = t[p].rc = a[nod[pl]]; return; } build(ls, pl, mid); build(rs, mid + 1, pr); up(p); } void upd(int p, int pl, int pr, int l, int r, int v){ if (l <= pl && pr <= r) return t[p].cov(v); down(p); if (l <= mid) upd(ls, pl, mid, l, r, v); if (r > mid) upd(rs, mid + 1, pr, l, r, v); up(p); } node query(int p, int pl, int pr, int l, int r){ if (l <= pl && pr <= r) return t[p]; node res = {0, 0, 0, 0, 0, 0}; down(p); if (l <= mid) merge(res, res, query(ls, pl, mid, l, r)); if (r > mid) merge(res, res, query(rs, mid + 1, pr, l, r)); return res; } void dfs2(int u, int tp){ top[u] = tp; dfn[u] = ++cnt; bot[u] = u; nod[cnt] = u; if (!son[u]) return build(rt[tp], dfn[tp], dfn[u]); dfs2(son[u], tp); bot[u] = bot[son[u]]; for (int v : g[u]) if (v != fa[u] && v != son[u]) dfs2(v, v); } void updt(int u, int v, int x){ while (top[u] != top[v]){ if (dep[top[u]] < dep[top[v]]) swap(u, v); upd(rt[top[u]], dfn[top[u]], dfn[bot[u]], dfn[top[u]], dfn[u], x); u = fa[top[u]]; } if (dep[u] > dep[v]) swap(u, v); upd(rt[top[u]], dfn[top[u]], dfn[bot[u]], dfn[u], dfn[v], x); } int qry(int u, int v){ node ures = {0, 0, 0, 0, 0, 0}, vres = {0, 0, 0, 0, 0, 0}; while (top[u] != top[v]){ if (dep[top[u]] > dep[top[v]]){ merge(ures, query(rt[top[u]], dfn[top[u]], dfn[bot[u]], dfn[top[u]], dfn[u]), ures); u = fa[top[u]]; } else { merge(vres, query(rt[top[v]], dfn[top[v]], dfn[bot[v]], dfn[top[v]], dfn[v]), vres); v = fa[top[v]]; } } if (dep[u] < dep[v]) merge(vres, query(rt[top[v]], dfn[top[v]], dfn[bot[v]], dfn[u], dfn[v]), vres); else merge(ures, query(rt[top[u]], dfn[top[u]], dfn[bot[u]], dfn[v], dfn[u]), ures); ures.flip(); merge(ures, ures, vres); return ures.sum + 1; } int main(){ ios::sync_with_stdio(0); cin.tie(0), cout.tie(0); cin >> n >> m; rep(u, 1, n) cin >> a[u]; rep(i, 2, n){ int u, v; cin >> u >> v; g[u].push_back(v); g[v].push_back(u); } dfs1(1, 0); dfs2(1, 1); while (m--){ char op; int u, v, x; cin >> op; if (op == 'C'){ cin >> u >> v >> x; updt(u, v, x); } else { cin >> u >> v; cout << qry(u, v) << '\n'; } } return 0; } ``` ::: ## 习题 可以参考我的[题单](https://www.luogu.com.cn/training/851940#problems),里面有几个长剖题,不过大部分以重剖为主。 # 拓展 ## 树链剖分维护树上动态 DP ### 总介绍 树链剖分维护树上动态 DP 是通过将 DP 状态转移方程写成广义矩阵乘法的形式,利用数据结构维护矩阵乘积转移,做到支持修改和查询操作。 需要注意,每个节点维护的是转移矩阵。 ### 路径类转移 这种题比较简单,相当于把一个序列上的 dp 放到了有向路径上面,直接按照题意写出序列上面的转移方程,然后用矩阵表示出来。 这样一次路径查询就是查询初始矩阵和路径上的转移矩阵依次相乘得到的结果,可以直接线段树维护。 如果不带修,可以用倍增查询路径矩阵乘积,空间差一些但是理论时间更加优秀。 也可以使用猫树等科技优化树剖的查询复杂度。 需要注意的是,由于大部分的路径类转移都和节点顺序有关,需要注意计算路径乘积时的顺序。 可以参考代码实现理解。 下面给出一道例题。 #### [P4679 [ZJOI2011] 道馆之战](https://www.luogu.com.cn/problem/P4679) ##### Description 给定一棵 $n$ 个点的树,每个节点分为 A、B 两个区域。 每个区域为冰块(`.`)或障碍物(`#`),障碍物不可以走。 沿着边行走的时候不能更换区域。 $m$ 次操作,分为两种类型。 - 类型 1:修改点 $u$ 的两个区域。 - 类型 2:从 $u$ 走到 $v$,询问可以经过的最大冰块区域数量(只能经过路径上的点)。 ##### Solution 先考虑序列上怎么做。 不妨考虑 DP。 设 $f_{u,A/B}$ 为处于节点 $u$ 的 $A/B$ 区域时,经过的冰块数量的最大值。 令 $g_u$ 为走到 $u$ 时的答案。 从 $u$ 走到 $v$ 时,显然转移与节点的区域情况有关。 开始分讨。 1. 当 $v$ 为 `..` 时: 有 $f_{v,A}=\max\{-\infty,f_{u,A}+1,f_{u,B}+2\}$,$f_{v,B}=\max\{-\infty,f_{u,A}+2,f_{u,B}+1\}$,$g_v=\max\{-\infty,f_{v,A},f_{v,B},g_{u}\}$。 2. 当 $v$ 为 `.#` 时: 有 $f_{v,A}=\max\{-\infty,f_{u,A}+1\}$,$f_{v,B}=-\infty$,$g_v=\max\{-\infty,f_{v,A},g_u\}$。 3. 当 $v$ 为 `#.` 时: 有 $f_{v,A}=-\infty$,$f_{v,B}=\max\{-\infty,f_{u,B}+1\}$,$g_v=\max\{-\infty,f_{v,B},g_u\}$。 4. 当 $v$ 为 `##` 时: 有 $f_{v,A}=f_{v,B}=-\infty$,$g_v=g_u$。 最终的 $g_v$ 就是答案。 发现转移方程都是 max-plus 的形式,而且只与上一个状态有关。 考虑使用矩阵加速。 定义矩阵乘法为 max-plus 形式。 将状态用矩阵表示为 $\begin{bmatrix}f_{u,A}&f_{u,B}&g_u\end{bmatrix}$,那么上面四种情况的转移对应的矩阵分别为: 1:$\begin{bmatrix}1&2&2\\2&1&2\\-\infty&-\infty&0\end{bmatrix}

2:\begin{bmatrix}1&-\infty&1\\-\infty&-\infty&-\infty\\-\infty&-\infty&0\end{bmatrix}

3:\begin{bmatrix}-\infty&-\infty&-\infty\\-\infty&1&1\\-\infty&-\infty&0\end{bmatrix}

4:\begin{bmatrix}-\infty&-\infty&-\infty\\-\infty&-\infty&-\infty\\-\infty&-\infty&0\end{bmatrix}

初始的状态为 \begin{bmatrix}0&0&0\end{bmatrix}。

进行转移只要把路径上所有的转移矩阵全部乘一遍就行了,最后取所得矩阵 M 的 g_v 即可。

只需要对每一个节点维护一个矩阵,使用线段树维护区间从左到右的乘积 mul1 和从右到左的 mul2 即可。

修改时,直接更新叶子节点的矩阵,然后更新祖先节点维护的乘积即可。 查询时,根据查询的方向选择矩阵相乘得到总转移矩阵,再让初始矩阵与转移矩阵相乘即可得到最终的矩阵 $M$,答案就是 $M_{0,2}$。 记矩阵边长 $a=3$。 时间复杂度: 树剖预处理 $O(n)$,线段树建树 $O(na^3)$,修改 $O(a^3\log n)$,查询 $O(a^3 \log^2 n)$。 空间复杂度: 树剖相关 $O(n)$,线段树 $O(a^2n)$,修改 $O(1)$,查询 $O(a^2)$。 由于矩阵运算很慢,尽可能减少运算次数以可以大幅优化常数。 ##### Code :::info[Code] ```cpp #include<bits/stdc++.h> using namespace std; #define rep(a,b,c) for(int a=(b);a<=(c);++a) #define per(a,b,c) for(int a=(b);a>=(c);--a) const int N = 5e4 + 5, INF = 0x3f3f3f3f; int n, m; vector<int> g[N]; char c[N][2]; int fa[N], sz[N], dep[N], son[N]; void dfs1(int u, int pr){ fa[u] = pr; sz[u] = 1; dep[u] = dep[pr] + 1; for (int v : g[u]) if (v != pr){ dfs1(v, u); sz[u] += sz[v]; if (sz[son[u]] < sz[v]) son[u] = v; } } int top[N], bot[N], dfn[N], nod[N], cnt; struct mat { int a[3][3]; }; const mat I = {{{0,-INF,-INF}, {-INF, 0,-INF}, {-INF,-INF, 0}}}, E = {{{-INF,-INF,-INF}, {-INF,-INF,-INF}, {-INF,-INF,-INF}}}; mat operator * (const mat & x, const mat & y){ mat res = E; rep(k, 0, 2) rep(i, 0, 2){ int v = x.a[i][k]; if (v !=-INF) rep(j, 0, 2) res.a[i][j] = max(res.a[i][j], v + y.a[k][j]); } return res; } const mat F[2][2] = { { {{{1, 2, 2}, {2, 1, 2}, {-INF,-INF, 0}}}, {{{1,-INF, 1}, {-INF,-INF,-INF}, {-INF,-INF, 0}}} } , { {{{-INF,-INF,-INF}, {-INF, 1, 1}, {-INF,-INF, 0}}}, {{{-INF,-INF,-INF}, {-INF,-INF,-INF}, {-INF,-INF, 0}}} } }; struct node { mat mul1, mul2; int lson, rson; } t[N<<2]; const node EN = {I, I, 0, 0}; int rt[N], ct; #define ls t[p].lson #define rs t[p].rson #define mid ((pl+pr)>>1) int typ(char ch){return ch == '#' ? 1 : 0;} node Node(int l, int r){ return {F[l][r], F[l][r], 0, 0}; } void merge(node & x, const node l, const node r){ x.mul1 = l.mul1 * r.mul1, x.mul2 = r.mul2 * l.mul2; } void build(int & p, int pl, int pr){ p = ++ct; if (pl == pr){ int u = nod[pl]; t[p] = Node(typ(c[u][0]), typ(c[u][1])); return; } build(ls, pl, mid); build(rs, mid + 1, pr); merge(t[p], t[ls], t[rs]); } void dfs2(int u, int tp){ top[u] = tp; bot[u] = u; dfn[u] = ++cnt; nod[cnt] = u; if (!son[u]) return build(rt[tp], dfn[tp], dfn[u]); dfs2(son[u], tp); bot[u] = bot[son[u]]; for (int v : g[u])if (v != fa[u] && v != son[u]) dfs2(v, v); } void upd(int p, int pl, int pr, int x, char l, char r){ if (pl == pr){ t[p] = Node(typ(l), typ(r)); return; } if (x <= mid) upd(ls, pl, mid, x, l, r); else upd(rs, mid + 1, pr, x, l, r); merge(t[p], t[ls], t[rs]); } node query(int p, int pl, int pr, int l, int r){ if (l <= pl && pr <= r) return t[p]; if (r <= mid) return query(ls, pl, mid, l, r); if (l > mid) return query(rs, mid + 1, pr, l, r); node res; merge(res, query(ls, pl, mid, l, r), query(rs, mid + 1, pr, l, r)); return res; } int qy(int u, int v){ mat f0 = {{{0, 0, 0}, {-INF,-INF,-INF}, {-INF,-INF,-INF}}}; node ures = EN, vres = EN; while (top[u] != top[v]){ if (dep[top[u]] > dep[top[v]]) merge(ures, query(rt[top[u]], dfn[top[u]], dfn[bot[u]], dfn[top[u]], dfn[u]), ures), u = fa[top[u]]; else merge(vres, query(rt[top[v]], dfn[top[v]], dfn[bot[v]], dfn[top[v]], dfn[v]), vres), v = fa[top[v]]; } if (dep[u] > dep[v]) merge(ures, query(rt[top[u]], dfn[top[u]], dfn[bot[u]], dfn[v], dfn[u]), ures); else merge(vres, query(rt[top[v]], dfn[top[v]], dfn[bot[v]], dfn[u], dfn[v]), vres); return(f0 * ures.mul2 * vres.mul1).a[0][2]; } int main(){ ios::sync_with_stdio(0); cin.tie(0), cout.tie(0); cin >> n >> m; rep(i, 2, n){ int u, v; cin >> u >> v; g[u].push_back(v); g[v].push_back(u); } rep(i, 1, n) cin >> c[i][0] >> c[i][1]; dfs1(1, 0); dfs2(1, 1); while (m--){ char op, l, r; int u, v; cin >> op; if (op == 'C'){ cin >> u >> l >> r; c[u][0] = l, c[u][1] = r; upd(rt[top[u]], dfn[top[u]], dfn[bot[u]], dfn[u], l, r); } else { cin >> u >> v; cout << qy(u, v) << '\n'; } } return 0; } ``` ::: ### 子树类转移 考虑一类常见的树形 DP 转移。 例如 $f_u=\sum_{fa_v=u}g_v$ 或者 $f_u =\max_{fa_v=u}g_v+a_v$ 类的转移。 这些转移都可以写成广义矩阵乘法的形式。 考虑如何用树剖维护。 重儿子的转移是容易的,直接以重儿子的值为初始贡献即可。 轻儿子的贡献需要额外维护。 一般可以记 $g_u$ 为 $u$ 和所有轻儿子对 $f_u$ 转移的贡献,那么需要在树剖之后对于每一个 $u$ 维护 $g_u$ 作为转移矩阵,查询时把 $u$ 的重链乘积拿出来和初始状态相乘即可。 需要注意数据结构维护时矩阵相乘的顺序。 修改要复杂一些,$u$ 和 $u$ 到根路径上轻边父亲的 $g$ 都会发生变化,需要额外维护。 由于轻边只有 $O(\log n)$ 条,一般可以直接暴力修改。 需要注意更新顺序的影响,一般是一边跳一边更新每一个 $u$。 这里以典题最大独立集为例。 #### [P4719 【模板】动态 DP](https://www.luogu.com.cn/problem/P4719) ##### Description 给定一棵 $n$ 个点的带点权树,有 $m$ 次单点修改,在每次修改后求出树的最大权独立集的权值。 ##### Solution 先考虑静态问题。 记 $f_{u,1}$ 为子树内为选 $u$ 的最大权独立集的权值,$f_{u,0}$ 为子树内不选 $u$ 的最大权独立集的权值。 有转移 $f_{u,1}=a_u+\sum_{fa_v=u}f_{v,0}$,$f_{u,0}=\sum_{fa_v=u}\max(f_{v,0},f_{v,1})$。 考虑描述 $u$ 和轻儿子 $v$ 的贡献。 根据转移方程,可以得到 $g_{u,1}=a_u+\sum_{fa_v=u \land v\ne son_u}f_{v,0}$,$g_{u,0}=\sum_{fa_v=u \land v \ne son_u}\max(f_{v,0},f_{v,1})$。 那么转移就变成了 $f_{u,1}=g_{u,1}+f_{son_u,0}$,$f_{u,0}=\max(g_{u,0}+f_{son_u,0},g_{u,0}+f_{son_u,1})$。 写成矩阵的形式,由于有 $\max$,可以写成 $(\max,+)$ 矩阵。 维护一个矩阵 $\begin{bmatrix}f_{u,1}&f_{u,0}\end{bmatrix}$ 作为当前状态的 dp 值。 根据上面的新方程得到 $u$ 的转移矩阵为 $\begin{bmatrix}-\infty&g_{u,0}\\g_{u,1}&g_{u,0}\end{bmatrix}$。 于是可以直接线段树维护每一条重链的转移。 初始矩阵显然为 $\begin{bmatrix}0&0\end{bmatrix}$,依次乘上重链的转移矩阵就可以得到 $f_u$。 对于修改,从 $u$ 不断向上跳,更新 $u$ 和每一个链顶父节点的转移矩阵即可,具体实现可以看代码。 代码中使用线段树维护矩阵乘积。 由于一条重链按照 dfs 序从小到大深度递增,线段树需要维护区间从右到左矩阵乘积。 时空复杂度多带了矩阵边长的系数,记矩阵边长 $a=2$,那么时间复杂度 $O((n+m \log^2n)a^3)$,空间复杂度 $O(na^2)$。 ##### Code :::info[Code] ```cpp #include<bits/stdc++.h> using namespace std; #define rep(a,b,c) for(int a=(b);a<=(c);++a) const int N = 1e5 + 5, INF = 0x3f3f3f3f; int n, m; int a[N]; vector<int> e[N]; int fa[N], sz[N], dep[N], son[N]; int f[N][2], g[N][2]; void dfs1(int u, int pr){ fa[u] = pr; sz[u] = 1; dep[u] = dep[pr] + 1; f[u][1] = a[u]; for (int v : e[u]) if (v != pr){ dfs1(v, u); sz[u] += sz[v]; if (sz[son[u]] < sz[v]) son[u] = v; f[u][1] += f[v][0]; f[u][0] += max(f[v][0], f[v][1]); } } int top[N], bot[N], dfn[N], nod[N], cnt; struct mat { int a[2][2]; }; const mat I = {{{0,-INF}, {-INF, 0}}}; const mat E = {{{-INF,-INF}, {-INF,-INF}}}; const mat F0 = {{{0, 0}, {-INF,-INF}}}; mat operator * (const mat & x, const mat & y){ return {{ {max(x.a[0][0] + y.a[0][0], x.a[0][1] + y.a[1][0]), max(x.a[0][0] + y.a[0][1], x.a[0][1] + y.a[1][1])}, {max(x.a[1][0] + y.a[0][0], x.a[1][1] + y.a[1][0]), max(x.a[1][0] + y.a[0][1], x.a[1][1] + y.a[1][1])} } }; } struct node { mat mul; int lson, rson; } t[N<<2]; int rt[N], ct; #define ls t[p].lson #define rs t[p].rson #define mid ((pl+pr)>>1) mat getmat(int u){return {{{-INF, g[u][0]}, {g[u][1], g[u][0]}}};} void up(int p){t[p].mul = t[rs].mul * t[ls].mul;} void build(int & p, int pl, int pr){ p = ++ct; if (pl == pr){ t[p].mul = getmat(nod[pl]); return; } build(ls, pl, mid); build(rs, mid + 1, pr); up(p); } void dfs2(int u, int tp){ top[u] = tp; bot[u] = u; dfn[u] = ++cnt; nod[cnt] = u; g[u][1] = a[u]; if (!son[u]) return; dfs2(son[u], tp); bot[u] = bot[son[u]]; for (int v : e[u]) if (v != fa[u] && v != son[u]){ dfs2(v, v); g[u][0] += max(f[v][0], f[v][1]); g[u][1] += f[v][0]; } } int stk[20], st; void upd(int u){ int p = rt[top[u]], pl = dfn[top[u]], pr = dfn[bot[u]], x = dfn[u]; while (pl != pr){ stk[++st] = p; x <= mid ? (pr = mid, p = ls) : (pl = mid + 1, p = rs); } t[p].mul = getmat(u); while (st) up(stk[st--]); } void updt(int u, int x){ g[u][1] += x - a[u]; a[u] = x; while (u){ int tp = top[u], p = fa[tp]; mat F = F0 * t[rt[tp]].mul; int f1 = F.a[0][0], f0 = F.a[0][1]; upd(u); mat nF = F0 * t[rt[tp]].mul; int nf1 = nF.a[0][0], nf0 = nF.a[0][1]; if (p) g[p][0] += max(nf1, nf0) - max(f1, f0), g[p][1] += nf0 - f0; u = p; } } int main(){ ios::sync_with_stdio(false); cin.tie(0), cout.tie(0); cin >> n >> m; rep(i, 1, n) cin >> a[i]; rep(i, 2, n){ int u, v; cin >> u >> v; e[u].push_back(v); e[v].push_back(u); } dfs1(1, 0); dfs2(1, 1); rep(u, 1, n) if (top[u] == u) build(rt[u], dfn[u], dfn[bot[u]]); while (m--){ int u, x; cin >> u >> x; updt(u, x); mat res = F0 * t[rt[1]].mul; cout << max(res.a[0][0], res.a[0][1]) << '\n'; } return 0; } ``` ::: 由于代码的常数优化,这个代码也可以直接通过[加强版](https://www.luogu.com.cn/problem/P4751)。 由于矩阵运算很慢,实现上可以尽量减少矩阵运算的次数,可以得到更快的速度。 卡常方式也很多,例如矩乘展开,递归改迭代,加快读等等。 ### 习题 同样有一份[题单](https://www.luogu.com.cn/training/1007000),其中部分题目为序列上的 DDP。 ## 树上启发式合并 / dsu on tree ### 介绍 启发式合并是利用重链剖分轻重儿子的形式进行子树统计的一种方式。 先考虑一个简单问题:[静态子树数颜色](https://www.luogu.com.cn/problem/U41492)。 显然可以转到 dfs 序上变成区间数颜色。 这里介绍另一个角度的方法。 先考虑暴力。 维护一个颜色出现次数的桶,对于每一个节点,把子树内所有的节点的颜色都加入桶里面,就可以统计出一个点子树的颜色数量。 统计完之后删除桶里面的颜色,然后解决下一个子树。 这样的时间复杂度是所有子树大小之和,最坏可以达到 $O(n^2)$,在树是一条链的情况取到。 但是由于节点 $u$ 的任意一个儿子 $v$ 的子树都包含在 $u$ 的子树内,所以可以考虑处理完某个儿子之后不删除,直接统计 $u$ 的子树。 贪心地想,每个 $u$ 继承子树大小最大的儿子肯定最好,也就是说选择重儿子。 这样复杂度就变成了所有轻子树的大小之和(将整棵树也看作一个轻子树),可以证明所有轻子树的大小之和是 $O(n \log n)$ 的。 考虑拆贡献,每个节点 $u$ 对轻子树大小和的贡献是 $u$ 到根路径上面轻边的数量加一,根据上文的证明,这个数量是 $O(\log n)$ 的。 所以这样可以得到正确的复杂度。 实现上,先进行 dfs 求出来每一个节点的重儿子,并保存 dfs 序以便遍历子树。 然后从根开始 dfs 每一个节点。 先递归所有轻儿子计算,每个轻儿子计算完之后立刻删除轻子树。 然后如果当前节点不是叶子,递归计算重儿子,计算完之后不删除。 然后加入当前节点和节点的轻子树,计算挂在节点上面的询问。 容易分析,上面的过程对于每个节点统计了其所有轻儿子的子树和节点本身,复杂度就是轻子树大小和,为 $O(n \log n)$。 以及启发式合并的“删除”操作其实是清空,有些题目删除的复杂度不是很优秀但是清空很方便。 这里给出例题的代码实现。 :::info[Code] ```cpp #include<bits/stdc++.h> using namespace std; #define rep(a,b,c) for(int a=(b);a<=(c);++a) #define subtree(i,u) for(int i=dfn[u];i<dfn[u]+sz[u];++i) const int N = 1e5 + 5; int n, m; int c[N]; vector<int> g[N]; int sz[N], son[N]; int dfn[N], nod[N], cnt; void dfs(int u, int pr){ sz[u] = 1; dfn[u] = ++cnt; nod[cnt] = u; for (int v : g[u]) if (v != pr){ dfs(v, u); sz[u] += sz[v]; if (sz[son[u]] < sz[v]) son[u] = v; } } int ans[N]; int buc[N], res; void add(int x){ res += !buc[x]++; } void del(int x){ --buc[x]; } void solve(int u, int pr){ for (int v : g[u]) if (v != pr && v != son[u]){ solve(v, u); subtree(i, v) del(c[nod[i]]); res = 0; } if (son[u]) solve(son[u], u); for (int v : g[u]) if (v != pr && v != son[u]) subtree(i, v) add(c[nod[i]]); add(c[u]); ans[u] = res; } int main(){ ios::sync_with_stdio(0); cin.tie(0), cout.tie(0); cin >> n; rep(i, 2, n){ int u, v; cin >> u >> v; g[u].push_back(v); g[v].push_back(u); } rep(u, 1, n) cin >> c[u]; dfs(1, 0); solve(1, 0); cin >> m; while (m--){ int u; cin >> u; cout << ans[u] << '\n'; } return 0; } ``` ::: ### 习题 由于个人题单数量限制,我没有新的题单位置了,建议直接找[树上启发式合并标签的题目](https://www.luogu.com.cn/problem/list?tag=55)做。