动态 DP

创建时间:2026-04-07


广义矩阵乘法

我们知道,一般矩阵乘法的定义如下:设有 \(m \times n\) 的矩阵 \(A\)\(n \times s\) 的矩阵,它们的乘积 \(AB\) 被定义一个 \(m \times s\) 的矩阵 \(C\),满足

\[C_{i, j} = \sum_{k=1}^n A_{i, k} B_{k, j} \]

一般矩阵乘法满足结合律,但不满足交换律。

定义广义矩阵乘法 \((\max, +)\),表示将一般矩阵乘法中的 \(+\) 替换为 \(\max\)\(\times\) 替换为 \(+\)。形式化的,广义矩阵乘法 \((\max, +)\)\(C = AB\) 满足

\[C_{i, j} = \max_{k = 1}^n\{A_{i, k} + B_{k, j}\} \]

显然,广义矩阵乘法同样 \((\max, +)\) 满足结合律,不满足交换律。类似的,我们也可以定义广义矩阵乘法 \((\min, +)\) 等。

广义矩阵乘法常用于矩阵快速幂优化 DP(包括定长最短路,动态 DP 等)。

动态 DP

动态 DP(Dynamic Dynamic Programming, DDP),即动态维护树上 DP。

P4719为例:给出一棵 \(n\) 个节点的树和每个点的点权 \(a_i\),有 \(m\) 次修改,每次修改形如 \(a_x := a_y\),在每次修改后求这棵树的最大权独立集。

下文默认点 \(1\) 是根,\(v\)\(u\) 的儿子。

考虑静态最大权独立集怎么求。

这是一个很经典的树上 DP 问题。令 \(1\) 为根,设 \(f_{i, 0}\)\(i\) 子树中不选 \(i\) 的答案,\(f_{i, 1}\)\(i\) 子树中选 \(i\) 的答案,状态转移如下:

\[f_{u, 0} = \sum_{v} \max(f_{v, 0}, f_{v, 1}) \]

\[f_{u, 1} = a_u + \sum_{v} f_{v, 0} \]

答案即为 \(\max(f_{1, 0}, f_{1, 1})\)。到现在,我们可以以 \(O(nm)\) 的时间复杂度解决问题。

上面算法的瓶颈在于每次修改都重算了所有节点的 \(f\),这看起来太傻了。注意到修改 \(x\) 的点权只对 \(x\)\(x\) 的祖先的 \(f\) 有影响,故考虑重链剖分。

\(son_u\) 表示 \(u\) 的重儿子,设

\[g_{u, 0} = \sum_{v \neq son_u} \max(f_{v, 0}, f_{v, 1}) \]

\[g_{u, 1} = a_u + \sum_{v \neq son_u} f_{v, 0} \]

\(f_u\) 可表示为

\[\begin{matrix} f_{u, 0} = & \max & (g_{u, 0} + f_{son_u, 0}, & g_{u, 0} + f_{son_u, 1}) \\ f_{u, 1} = & & g_{u, 1} + f_{son_u, 0} \end{matrix}\]

这像极了广义矩阵乘法 \((\max, +)\),具体的,我们有:

\[\begin{bmatrix} f_{u, 0} \\ f_{u, 1} \end{bmatrix} = \begin{bmatrix} g_{u, 0} & g_{u, 0} \\ g_{u, 1} & -\infty \end{bmatrix} \begin{bmatrix} f_{son_u, 0} \\ f_{son_u, 1} \end{bmatrix}\]

不难推出,一条重链 \(p_1, p_2, \cdots, p_k\) 以及这条重链上挂着的虚子树的答案为

\[(\prod_{i=1}^k \begin{bmatrix} g_{p_i, 0} & g_{p_i, 0} \\ g_{p_i, 1} & -\infty \end{bmatrix}) \begin{bmatrix} 0 \\ 0 \end{bmatrix}\]

通过线段树维护每条重链以及其虚子树的答案,每次换链时计算差值并在线段树上更新即可,时间复杂度 \(O(n \log n + m \log^2 n)\)

代码
#include <bits/stdc++.h>

using namespace std;

const int N = 1e5 + 50;
const int inf = 1e9;

int n, m, a[N], f[N][2], g[N][2];
int fa[N], dep[N], siz[N], son[N];
int top[N], ed[N], tim, dfn[N], rev[N];
vector<int> e[N];

void dfs1(int u, int father, int depth) {
	fa[u] = father;
	dep[u] = depth;
	siz[u] = 1;
	for (int v : e[u]) {
		if (v == father) {
			continue;
		}
		dfs1(v, u, depth + 1);
		siz[u] += siz[v];
		if (siz[v] > siz[son[u]]) {
			son[u] = v;
		} 
	}
}

void dfs2(int u, int topf) {
	top[u] = topf;
	ed[u] = u;
	dfn[u] = ++tim;
	rev[tim] = u;
	f[u][1] = g[u][1] = a[u];
	if (son[u]) {
		dfs2(son[u], topf);
		ed[u] = ed[son[u]];
		f[u][0] += max(f[son[u]][0], f[son[u]][1]);
		f[u][1] += f[son[u]][0];
	}
	for (int v : e[u]) {
		if (v == fa[u] || v == son[u]) {
			continue;
		}
		dfs2(v, v);
		f[u][0] += max(f[v][0], f[v][1]);
		f[u][1] += f[v][0];
		g[u][0] += max(f[v][0], f[v][1]);
		g[u][1] += f[v][0];
	}
}

struct Matrix {
	int mat[2][2];
	
	Matrix() {
		mat[0][0] = mat[1][1] = 0;
		mat[0][1] = mat[1][0] = -inf;
	}
	
	int* operator [] (int i) {
		return mat[i];
	}
	const int* operator [] (int i) const {
		return mat[i]; 
	}
	
	friend Matrix operator * (const Matrix& a, const Matrix& b) {
		Matrix c;
		c[0][0] = max(a[0][0] + b[0][0], a[0][1] + b[1][0]);
		c[0][1] = max(a[0][0] + b[0][1], a[0][1] + b[1][1]);
		c[1][0] = max(a[1][0] + b[0][0], a[1][1] + b[1][0]);
		c[1][1] = max(a[1][0] + b[0][1], a[1][1] + b[1][1]);
		return c;
	}
} tr[N << 2];

#define ls cur << 1
#define rs cur << 1 | 1

void pushup(int cur) {
	tr[cur] = tr[ls] * tr[rs];
}

void build(int cur, int l, int r) {
	if (l == r) {
		tr[cur][0][0] = tr[cur][0][1] = g[rev[l]][0];
		tr[cur][1][0] = g[rev[l]][1];
		tr[cur][1][1] = -inf;
		return ;
	}
	int mid = l + r >> 1;
	build(ls, l, mid);
	build(rs, mid + 1, r);
	pushup(cur);
}

void update(int cur, int l, int r, int i) {
	if (l == r) {
		tr[cur][0][0] = tr[cur][0][1] = g[rev[l]][0];
		tr[cur][1][0] = g[rev[l]][1];
		return ;
	}
	int mid = l + r >> 1;
	if (i <= mid) {
		update(ls, l, mid, i);
	} else {
		update(rs, mid + 1, r, i);
	}
	pushup(cur);
} 

Matrix query(int cur, int l, int r, int L, int R) {
	if (L <= l && r <= R) {
		return tr[cur];
	}
	int mid = l + r >> 1;
	Matrix res;
	if (L <= mid) {
		res = query(ls, l, mid, L, R);
	}
	if (mid + 1 <= R) {
		res = res * query(rs, mid + 1, r, L, R);
	}
	return res;
}

void modify(int u, int w) {
	g[u][1] += w - a[u];
	a[u] = w;
	while (u) {
		auto old = query(1, 1, n, dfn[top[u]], dfn[ed[u]]);
		int old_f[] = {max(old[0][0], old[0][1]), max(old[1][0], old[1][1])};
		update(1, 1, n, dfn[u]);
		auto now = query(1, 1, n, dfn[top[u]], dfn[ed[u]]);
		int now_f[] = {max(now[0][0], now[0][1]), max(now[1][0], now[1][1])};
		u = fa[top[u]];
		if (u) {
			g[u][0] += max(now_f[0], now_f[1]) - max(old_f[0], old_f[1]);
			g[u][1] += now_f[0] - old_f[0];
		}
	}
}

int 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, u, v; i < n; i++) {
		cin >> u >> v;
		e[u].push_back(v);
		e[v].push_back(u);
	}
	
	dfs1(1, 0, 1);
	dfs2(1, 1);
	build(1, 1, n);
	
	for (int x, y; m--; ) {
		cin >> x >> y;
		modify(x, y);
		auto tmp = query(1, 1, n, 1, dfn[ed[1]]);
		cout << max({tmp[0][0], tmp[0][1], tmp[1][0], tmp[1][1]}) << '\n';
	}
	return 0;
}

同理,我们也可以使用 LCT 替代重链剖分,时间复杂度 \(O(n + m \log n)\)

posted @ 2026-06-02 17:00  xubaichuan  阅读(7)  评论(0)    收藏  举报