Loading

树链剖分 & 动态 DP 笔记

虽迟但到。

树链剖分

其实是一个很简单的技巧。

考虑如下问题:维护一棵树,支持路径加操作与路径和查询。

我们没办法直接维护,所以树剖的作用就是:将树划分成若干条链,使得任意两点之间的路径最多经过 \(\mathcal{O}(\log n)\) 条链。

这样,我们就有可行方案了:对每条链分别使用数据结构维护。


考虑一下怎么分割可以达成上述目的,接下来我讲的分割方法叫做重链剖分。

我们将某个点的 \(u\) 的子树大小记作 \(sz_u\),那么定义 \(u\) 的重儿子为 \(u\) 的儿子 \(v\)\(sz_v\) 最大的(取一个)。

顾名思义,我们对于每个点,只向她的重儿子连一条边(我们把这条边叫做重边,剩下的叫轻边),那么最后肯定是分成了若干条链。

那么考虑为啥这样是正确的:其实 \(u \to v\) 上的链数,和 \(u \to v\) 上的轻边个数同阶。考虑从 \(\text{lca}(u, v)\)\(u\) 走:遇到一条轻边,当前点所在的子树大小至少减小一半,因此最多经过 \(\mathcal{O}(\log n)\) 条轻边,\(v\) 同理。

所以,\(u \to v\) 时最多经过 \(\mathcal{O}(\log n)\) 条链。


考虑一点实现细节,因为之后也要用到。

其实解决上面的维护问题,只需要开一棵线段树即可,因为每条链其实可以作为线段树上的一个区间,不同区间互不干涉。

那么对于每个点,我们要记录:

  • 点的深度 \(dep_u\)

  • 她在线段树上的编号 \(idx_u\)(在这道题里,我们要按每个点所属的链的区间对点进行重编号)。

  • 这个点所属的链的顶端的点 \(top_u\)(为了实现 \(\mathcal{O}(1)\) 跳过某一条链)。


考虑一个难一点的问题:如果加入了子树加和子树求和该怎么办?

考虑构造一个 \(idx\):对于每一个点 \(u\),优先对重儿子 DFS,然后对轻儿子 DFS,\(dfn_u \to idx_u\)

那么这种编号方案满足:

  • 一个子树的编号是一个连续区间。

  • 一条链的编号是一个连续区间。

通过这种方式,我们不需要任何额外的处理,非常方便。

树剖求 LCA

在倍增跳里,我们预处理 \(u\)\(2^i\) 祖先是谁。在树剖之后,我们依旧采用向上跳的形式。

考虑当前有两个点 \(u, v\),我们将 \(u\) 的父亲记作 \(fa_u\),算法流程是:

  1. 找到 \(x = fa_{top_u}, y = fa_{top_v}\)

  2. 如果 \(dep_x \ge dep_y\),那么 \(x \to u\),否则 \(y \to u\)

  3. 如果 \(top_u = top_v\),深度较小的那个就是 \(\text{lca}\)

正确性显然。而且修改/查询其实都可以在求 LCA 的过程中顺带统计。而且这个代码其实奇短。

到这里,恭喜你,已经会做树剖模板了!然后你就能过很多树剖板子。

代码其实非常好写。

#include <bits/stdc++.h>
#define int long long
using namespace std;

const int N = 1e5 + 7;
int n, m, rt, MOD, sz[N], fa[N], org[N];
int lp[N], rp[N], idx, top[N], dep[N];
vector<int> g[N];

void get_sz(int u, int pre){
	sz[u] = 1, fa[u] = pre;
	dep[u] = dep[pre] + 1;
	for(int v: g[u]) if(v != pre){
		get_sz(v, u);
		sz[u] += sz[v];
	}
}
void get_dfn(int u, int pre, int tp){
	lp[u] = ++ idx, top[u] = tp;
	int id = 0;
	for(int v: g[u]) if(v != pre)
		if(sz[v] > sz[id]) id = v;
	if(id) get_dfn(id, u, tp);
	for(int v: g[u])
		if(v != pre && v != id) get_dfn(v, u, v);
	rp[u] = idx;
}

struct Misaka{
	int val[N << 2], lzy[N << 2];
	#define ls (x << 1)
	#define rs ((x << 1) | 1)
	#define mid ((l + r) >> 1)
	void pushup(int x){(val[x] = val[ls] + val[rs]) %= MOD;}
	void pushdown(int x, int l, int r){
		(val[ls] += (mid - l + 1) * lzy[x]) %= MOD;
		(val[rs] += (r - mid) * lzy[x]) %= MOD;
		(lzy[ls] += lzy[x]) %= MOD, (lzy[rs] += lzy[x]) %= MOD;
		lzy[x] = 0;
	}
	void add(int x, int l, int r, int ql, int qr, int k){
		if(ql <= l && r <= qr){
			(val[x] += (r - l + 1) * k) %= MOD;
			(lzy[x] += k) %= MOD; return;
		}
		pushdown(x, l, r);
		if(ql <= mid) add(ls, l, mid, ql, qr, k);
		if(mid < qr) add(rs, mid + 1, r, ql, qr, k);
		pushup(x);
 	}
 	int sum(int x, int l, int r, int ql, int qr){
 		if(ql <= l && r <= qr) return val[x];
 		pushdown(x, l, r);
 		int res = 0;
 		if(ql <= mid) (res += sum(ls, l, mid, ql, qr)) %= MOD;
 		if(mid < qr)  (res += sum(rs, mid + 1, r, ql, qr)) %= MOD;
 		return res;
 	}
};
Misaka Mikoto;

int opr_ch(int opt, int x, int y, int k){
	int res = 0;
	while(top[x] != top[y]){
		if(dep[top[x]] >= dep[top[y]]){
			if(opt == 0) Mikoto.add(1, 1, n, lp[top[x]], lp[x], k);
			else (res += Mikoto.sum(1, 1, n, lp[top[x]], lp[x])) %= MOD;
			x = fa[top[x]];
		}
		else{
			if(opt == 0) Mikoto.add(1, 1, n, lp[top[y]], lp[y], k);
			else (res += Mikoto.sum(1, 1, n, lp[top[y]], lp[y])) %= MOD;
			y = fa[top[y]];
		}
	}
	if(dep[x] > dep[y]) swap(x, y);
	if(opt == 0) Mikoto.add(1, 1, n, lp[x], lp[y], k);
	else (res += Mikoto.sum(1, 1, n, lp[x], lp[y])) %= MOD;
	return res;
}

int opr_sub(int opt, int u, int k){
	int res = 0;
	if(opt == 0) Mikoto.add(1, 1, n, lp[u], rp[u], k);
	else (res += Mikoto.sum(1, 1, n, lp[u], rp[u])) %= MOD;
	return res;
}

signed main(){
  ios::sync_with_stdio(0), cin.tie(0);
  
  cin >> n >> m >> rt >> MOD;
  for(int i = 1; i <= n; i ++) cin >> org[i];
  for(int i = 1; i < n; i ++){
  	int x, y; cin >> x >> y;
  	g[x].push_back(y);
  	g[y].push_back(x);
  }
  
  get_sz(rt, 0);
  get_dfn(rt, 0, rt);
  
  for(int i = 1; i <= n; i ++) Mikoto.add(1, 1, n, lp[i], lp[i], org[i]);
  
  while(m --){
  	int opt, x; cin >> opt >> x;
  	if(opt == 1){
  		int y, z; cin >> y >> z;
  		opr_ch(0, x, y, z);
  	}
  	else if(opt == 2){
  		int y; cin >> y;
  		cout << opr_ch(1, x, y, 0) << "\n";
  	}
  	else if(opt == 3){
  		int z; cin >> z;
  		opr_sub(0, x, z);
  	}
  	else cout << opr_sub(1, x, 0) << "\n";
  }
  
  return 0;
}

动态 DP

似乎叫 DDP。

我们都知道,DP 只适用于静态的信息。我们先考虑一个问题:

我们现在有一个 \([1, n]\) 的 DP,每次问你 \([l, r]\) 的 DP 结果是啥(DP 并不满足逆运算,也就是说你无法返回 \(dp_r \cdot \text{inv}(dp_{l - 1})\))。

有一个思想是:如果你的 DP 可以使用矩阵乘法转移,并且这种矩阵乘法具有结合律,那么我们其实可以对每次转移构造一个矩阵。

然后将这些矩阵挂在线段树上,那么每次只需要查询树上的 \(\mathcal{O}(\log n)\) 个节点,并且用矩阵乘法把这些区间拼接起来,就可以 \(\mathcal{O}(\log n)\) 查询结果了。

显然这个肯定支持单点修改操作,因为只需要修改线段树上的一个矩阵即可。


顺带说说怎么证明矩阵乘法的结合律。假如我们现在想证明 \((\min, +)\) 矩阵的结合律。

那么证明过程如下:

\[(AB)_{i, j} = \min_k(A_{i, k} + B_{k, j}) \]

于是我们有

\[[(AB)C]_{i, j} = \min_s(\min_k(A_{i, k} + B_{k, s}) + C_{s, j}) \]

显然我们可以把 \(\min\) 挪到外面,所以

\[[(AB)C]_{i, j} = \min_{s, k} (A_{i, k} + B_{k, s} + C_{s, j}) \]

那么显然 \((AB)C\)\(A(BC)\) 得到的结果都是这个,因此 \((\min, +)\) 的矩阵乘法满足结合律。


那么如果你把这个放到一棵树上,显然就是把每个点的转移矩阵处理出来,用线段树维护一下矩阵乘。

我们用如下例题讲述具体做法:单点修改权值,求一条链的最大带权独立集。

显然我们可以定义 \(f(i, 0/1)\) 代表在 \(i\) 所在的链里由下至上考虑到 \(i\),不选/选 \(i\) 点的最大带权独立集。

转移是容易的,可以很轻易地化成矩阵。

然后单点修改,即对一个矩阵进行修改,着重说一下怎么查询,以及矩阵结合的顺序:

当我们想查询 \(u \to v\) 时,我们可以在上述的 \(\text{lca}\) 算法中统计。

具体地,我们在求 \(\text{lca}\) 的过程中分别维护矩阵 \(U, V\),她们最终代表 \(u \to \text{lca}(u, v)\)\(v \to \text{lca}(u, v)\) 的转移结果(注意这两条链不包含 \(\text{lca}(u, v)\))。

然后枚举一下 \(\text{lca}(u, v)\) 是否选取即可。


考虑上述问题的加强版:每次求整棵树的最大带权独立集。

显然,只维护 \(f\) 数组肯定是不行了。为了迎合树剖,我们设置 \(g(u, 0/1)\) 代表只考虑 \(u\) 的轻儿子时,\(u\) 不选/选,\(u\) 子树的最大带权独立集,我们记 \(u\) 的重儿子为 \(v\),有转移

\[f(u, 0) = \max(f(v, 0), f(v, 1)) + g(u, 0) \]

\[f(u, 1) = f(v, 0) + g(u, 1) \]

因为转移矩阵所在的节点编号应该比 \(f(j, 0/1)\) 所在的节点编号小,也就是说,转移矩阵应该放在左边,所以这个转移的矩阵形式,是

\[\left[ \begin{array}{cc} g(u, 0) & g(u, 0) \\ g(u, 1) & -\infty \end{array} \right] \left[ \begin{array}{cc} f(v, 0) \\ f(v, 1) \end{array} \right] = \left[ \begin{array}{cc} f(u, 0) \\ f(u, 1) \end{array} \right]\]

考虑对一个点进行修改时,假设 \(a_u\) 的变化量为 \(\Delta\),那么显然

\[g(u, 1) = g(u, 1) + \Delta \]

并且因为这条链上的 \(g\) 值发生了变化,那么 \(f(top_u, *)\) 便会伴随着矩阵乘法被改变。那么 \(g(fa_{top_u}, *)\) 也会变。就这样,每通过一条轻边,就会有一次 \(f, g\) 的重计算。由于 \(u\) 到根一共只会有 \(\mathcal{O}(\log n)\) 条轻边,所以加上数据结构的复杂度,一次修改的总复杂度其实是 \(\mathcal{O}(\log ^ 2 n)\) 的。

树剖常数小哇。

posted @ 2026-08-27 13:00  Trent900  阅读(18)  评论(0)    收藏  举报