动态 DP
创建时间:2026-04-07
广义矩阵乘法
我们知道,一般矩阵乘法的定义如下:设有 \(m \times n\) 的矩阵 \(A\) 与 \(n \times s\) 的矩阵,它们的乘积 \(AB\) 被定义一个 \(m \times s\) 的矩阵 \(C\),满足
一般矩阵乘法满足结合律,但不满足交换律。
定义广义矩阵乘法 \((\max, +)\),表示将一般矩阵乘法中的 \(+\) 替换为 \(\max\),\(\times\) 替换为 \(+\)。形式化的,广义矩阵乘法 \((\max, +)\) 中 \(C = AB\) 满足
显然,广义矩阵乘法同样 \((\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\) 的答案,状态转移如下:
答案即为 \(\max(f_{1, 0}, f_{1, 1})\)。到现在,我们可以以 \(O(nm)\) 的时间复杂度解决问题。
上面算法的瓶颈在于每次修改都重算了所有节点的 \(f\),这看起来太傻了。注意到修改 \(x\) 的点权只对 \(x\) 和 \(x\) 的祖先的 \(f\) 有影响,故考虑重链剖分。
令 \(son_u\) 表示 \(u\) 的重儿子,设
则 \(f_u\) 可表示为
这像极了广义矩阵乘法 \((\max, +)\),具体的,我们有:
不难推出,一条重链 \(p_1, p_2, \cdots, p_k\) 以及这条重链上挂着的虚子树的答案为
通过线段树维护每条重链以及其虚子树的答案,每次换链时计算差值并在线段树上更新即可,时间复杂度 \(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)\)。

浙公网安备 33010602011771号