动态 DP 入门

· · 算法·理论

前置知识:矩阵加速(习题),树链剖分(习题)。

动态 DP

简介

动态动态规划(Dynamic Dynamic Programing,后文简称 DDP),一般认为是一种用于处理点权改变的树上 DP 问题。它通过将状态转移方程式变为矩阵等运算,配合树剖全局平衡二叉树等数据结构,可以将 O(n) 的修改优化到双 \log 甚至单 \log 级别。

从名字上就可以看出来,这东西和 DDT 也就只差一个字,因此它的难度可见一斑。

由于 DP 类问题不存在固定的计算方式,只能见招拆招,因此下面将会用实际例题的方式讲解 DDP。

例题:P4719 【模板】动态 DP

我们先看下面这道模板题:

给定一颗有 n 个节点的树,点带权。

m 次操作,每次修改点 x 的权为 y

你需要在每次修改结束后求出树上最大权独立集的权值。

我们先考虑静态 DP 的情况。设 dp_{u,1/0} 表示选或者不选 u 进入最大权独立集,可以很快得出方程式为:

dp_{u,0}=\sum\limits_{v \in son(u)}{\max{(dp_{v,0},dp_{v,1})}} dp_{u,1}=\sum\limits_{v \in son(u)}dp_{v,0} + val_u

最后答案为 \max(dp_{1,0}, dp_{1,1})

对于每一次修改操作,如果直接暴力修改,就可能需要重新递推,复杂度最坏可以达到 O(n),因此需要考虑优化。

广义矩阵乘法

我们注意到,整个 DP 式子只有两种运算方式:求和以及取最大值。同时,整个递推过程是线性齐次的,而我们都知道,矩阵加速就是专门用来优化这类问题的。因此可以考虑将问题转化成矩阵运算来解决。

但新的问题来了:常见的矩阵乘法都是诸如 C_{i,j}=\sum{a_{i, k} \times b_{k, j}} 的形式。而这里的运算方式很明显不适用,因此需要一种适用面更广泛的运算方式来计算。

我们都知道,矩阵乘法最重要的性质就是满足结合律,因此我们可以大胆猜想,只要满足结合律,任意的运算方式都可以用来替代加法和乘法。而显然加法和取最大值是满足结合律的,因此我们可以将乘法替换成加法加法替换成最大值,得到新的运算方式:

C_{i,j} = \max\limits_{k=1}^{n}(A_{i,k}+B_{k,j})

这就是广义矩阵乘法。它显然是满足结合律的,读者下来可以自证。

树剖优化

现在有了广义矩阵乘法,但你可以发现,原式子的递推是多叉的,而广义矩阵乘法还不足以实现它的维护,因此还需要进一步优化。如果能够把多叉求和变成单个儿子的转移就好了。

这个时候,树剖就登场了!

我们考虑对原树进行重链剖分,就可以对每一个节点分出重儿子和轻儿子。

然后我们考虑把重儿子和轻儿子的信息分开维护,我们令 g_{u,1/0} 分别表示在选或不选节点 u自己 + 轻儿子所产生的贡献,令 s 表示重儿子,于是方程可以被我们进一步简化为:

dp_{u,0} = g_{u,0}+\max(dp_{s,0},dp_{s,1}) dp_{u,1} = g_{u,1}+dp_{s,0}

对于每一次更新,类比树剖,我们就通过跳链的方式更新。根据树剖的性质(不考虑数据结构维护的复杂度)最后跳链复杂度就是单 \log

转移矩阵

通过树剖对式子变形后,接下来就可以考虑怎么构建转移矩阵了。

先再来看一遍我们刚刚得到的式子:

dp_{u,0} = g_{u,0}+\max(dp_{s,0},dp_{s,1}) dp_{u,1} = g_{u,1}+dp_{s,0}

因为我们需要将这个式子变成满足 max-plus 的标准形式,因此我们需要将这个式子改写一下,把 \max 套在最外面且不影响本身的结果。于是这个式子就可以变成这样:

dp_{u,0} = \max(g_{u,0}+dp_{s,0},g_{u,0}+dp_{s,1}) dp_{u,1} = \max(g_{u,1}+dp_{s,0},-\infty+dp_{s,1})

接下来,就可以利用矩阵的运算,找出矩阵,进行下一步的变换。于是这个式子就可以长成这样:

\begin{pmatrix} dp_{u,0} \\ dp_{u,1} \end{pmatrix} = \begin{pmatrix} g_{u,0} & g_{u,0} \\ g_{u,1} & -\infty \end{pmatrix} \otimes \begin{pmatrix} dp_{s,0} \\ dp_{s,1} \end{pmatrix}

至此,我们就列出了它的矩阵形式。

代码讲解

相信大家看肯定都是看明白了,但写不一定写的对。作为一篇写给蒻自己的文章,下面将会给出一些写代码时的讲解和细节。内容比较细且基础,码力好的大神们可以直接跳过。

树剖线段树

线段树的操作和常见的线段树没有特别大的区别,主要在于将乘的信息从元素变成了矩阵乘法。要注意的有两点:第一是矩阵乘法是不满足交换律的,由于我们的乘法是用深度浅的 u 乘深度深的 s,因此pushup的顺序一定是深度浅的 \times 深度深的而不能反过来。第二是矩阵乘法的单位元是单位矩阵,初始化的时候要注意一下。

关于树剖,除了常见的六个数组以外,由于我们需要查询一条链的 DP 值,因此在此处还需要维护一个 bottom 数组来记录每条重链的链尾。具体的维护方式就是在第二遍 DFS 的时候加一个 bottom[tp]=max(bottom[tp],dfn[u]);就好了。

然后关于 g 数组的贡献也是在这里面进行累加的。具体做法就是先加上自己的贡献,在每次递归完轻儿子后直接加上轻儿子的贡献即可。

最后为了节省码量,我们在每一个节点更新完成后就直接 Modify 它的贡献进线段树,这样就可以省掉 build 函数了。

int dfn[N], dep[N], top[N], fa[N], siz[N], heavy[N], bottom[N];//链尾 
int timi = 0;

void dfs1(int u, int f){
    dep[u] = dep[f] + 1;
    fa[u] = f;
    siz[u] = 1;
    for (int v : tree[u]){
        if(v == f) continue;
        dfs1(v, u);
        siz[u] += siz[v];
        if(siz[v] > siz[heavy[u]]) heavy[u] = v;
    }
}

int g[N][2];

VEC getf(int u){//获取DP值
    MAT xx = Query(1, 1, n, dfn[u], bottom[top[u]]);
    VEC yy = VEC(0, -INF);
    return xx * yy; //乘上它相当于得到的两个元素是f[u][0]和f[u][1]
}

MAT buildmat(int u){//获取节点的转移矩阵 
    MAT res;
    res.a[0][0] = res.a[0][1] = g[u][0];
    res.a[1][0] = g[u][1], res.a[1][1] = -INF;
    return res;
}

void dfs2(int u, int tp){
    dfn[u] = ++timi;
    top[u] = tp;
    bottom[tp] = max(bottom[tp], dfn[u]);//链底 
    g[u][0] = 0;    
    g[u][1] = val[u];

    if(heavy[u]) dfs2(heavy[u], tp);

    for (int v : tree[u]){
        if(v == fa[u] || v == heavy[u]) continue;
        dfs2(v, v);
        VEC vv = getf(v);
        g[u][0] += max(vv.a[0], vv.a[1]);
        g[u][1] += vv.a[0];
    }

    Modify(1, 1, n, dfn[u], buildmat(u));
}

点权更新

首先跳链前更新自己的权值,这没什么好说的。主要考虑跳链的过程应该怎么搞。每一次循环的时候,我们需要记录一下链顶在更新前的旧值,再在当前节点做一次 Modify,然后算出链顶的新值,在链顶的父亲上更新结果就完成了一轮迭代。

其实看下来比树剖转链还好写···?

void update(int x, int v){
    g[x][1] += v - val[x];
    val[x] = v;//更新自身

    while(x){
        int t = top[x];
        VEC f1 = getf(t);

        Modify(1, 1, n, dfn[x], buildmat(x));//更新节点矩阵 

        VEC f2 = getf(t);
        if(fa[t]){
            g[fa[t]][0] += max(f2.a[0], f2.a[1]) - max(f1.a[0], f1.a[1]);
            g[fa[t]][1] += f2.a[0] - f1.a[0];
        }
        x = fa[t];//跳到上一条链 
    } 
}

完整代码

最后放一个完整的代码。

#include<bits/stdc++.h>
#define LoveCyrene ios :: sync_with_stdio(false), cin.tie(nullptr), cout.tie(nullptr)
#define int long long
using namespace std;;
constexpr int N = 1e5 + 8179;
constexpr int INF = 0x3f3f3f3f;

struct MAT{
    int a[2][2];
    MAT(){ a[0][0] = a[0][1] = a[1][0] = a[1][1] = -INF; }
    MAT friend operator *(MAT x, MAT y){
        MAT z;
        for (int i = 0; i < 2; i++){
            for (int j = 0; j < 2; j++){
                for (int k = 0; k < 2; k++){
                    z.a[i][k] = max(z.a[i][k], x.a[i][j] + y.a[j][k]);
                }
            }
        }
        return z;
    }
};

struct VEC{ 
    int a[2];
    VEC(){ a[0] = a[1] = -INF; }
    VEC(int x, int y) { a[0] = x, a[1] = y; }
    friend VEC operator * (MAT& y, VEC& x){
        VEC z;
        for (int i = 0; i < 2; i++){
            for (int j = 0; j < 2; j++){
                z.a[i] = max(z.a[i], x.a[j] + y.a[i][j]);
            }   
        }
        return z;
    }   
};

struct SEGMENT{
    MAT sgt[N << 2];

    void pushup(int x){
        sgt[x] = sgt[x << 1] * sgt[x << 1 | 1];
    }

    void modify(int x, int l, int r, int q, MAT v){
        if(l == r){ sgt[x] = v; return; }
        int mid = (l + r) >> 1;
        if(q <= mid) modify(x << 1, l, mid, q, v);
        else modify(x << 1 | 1, mid + 1, r, q, v);
        pushup(x);
    }

    MAT query(int x, int l, int r, int ql, int qr){
        if(ql <= l && r <= qr) return sgt[x];
        int mid = (l + r) >> 1;
        MAT res;
        res.a[0][0] = res.a[1][1] = 0;
        if(ql <= mid) res = res * query(x << 1, l, mid, ql, qr);
        if(qr > mid) res = res * query(x << 1 | 1, mid + 1, r, ql, qr);
        return res;
    }
    #define sgt(x) seg.sgt(x)
    #define Modify(x, l, r, q, v) seg.modify(x, l, r, q, v)
    #define Query(x, l, r, ql, qr) seg.query(x, l, r, ql, qr)
}seg;

int n, m;
int val[N];
vector<int> tree[N];

void add_edge(int u, int v){ tree[u].push_back(v), tree[v].push_back(u); }

int dfn[N], dep[N], top[N], fa[N], siz[N], heavy[N], bottom[N]; 
int timi = 0;

void dfs1(int u, int f){
    dep[u] = dep[f] + 1;
    fa[u] = f;
    siz[u] = 1;
    for (int v : tree[u]){
        if(v == f) continue;
        dfs1(v, u);
        siz[u] += siz[v];
        if(siz[v] > siz[heavy[u]]) heavy[u] = v;
    }
}

int g[N][2];

VEC getf(int u){
    MAT xx = Query(1, 1, n, dfn[u], bottom[top[u]]);
    VEC yy = VEC(0, -INF);
    return xx * yy; 
}

MAT buildmat(int u){
    MAT res;
    res.a[0][0] = res.a[0][1] = g[u][0];
    res.a[1][0] = g[u][1], res.a[1][1] = -INF;
    return res;
}

void dfs2(int u, int tp){
    dfn[u] = ++timi;
    top[u] = tp;
    bottom[tp] = max(bottom[tp], dfn[u]); 
    g[u][0] = 0;    
    g[u][1] = val[u];

    if(heavy[u]) dfs2(heavy[u], tp);

    for (int v : tree[u]){
        if(v == fa[u] || v == heavy[u]) continue;
        dfs2(v, v);
        VEC vv = getf(v);
        g[u][0] += max(vv.a[0], vv.a[1]);
        g[u][1] += vv.a[0];
    }

    Modify(1, 1, n, dfn[u], buildmat(u));
}

void update(int x, int v){
    g[x][1] += v - val[x];
    val[x] = v;

    while(x){
        int t = top[x];
        VEC f1 = getf(t);

        Modify(1, 1, n, dfn[x], buildmat(x));

        VEC f2 = getf(t);
        if(fa[t]){
            g[fa[t]][0] += max(f2.a[0], f2.a[1]) - max(f1.a[0], f1.a[1]);
            g[fa[t]][1] += f2.a[0] - f1.a[0];
        }
        x = fa[t];
    } 
}

void init(){ dfs1(1, 0), dfs2(1, 1); }

int curans(){ return max(getf(1).a[0], getf(1).a[1]); }

signed main(){
    LoveCyrene;
    cin >> n >> m;
    for (int i = 1; i <= n; i++) cin >> val[i];
    for (int i = 1, u, v; i < n; i++) cin >> u >> v, add_edge(u, v);
    init();

    while(m--){ int x, y; cin >> x >> y, update(x, y), cout << curans() << '\n'; }
    return 0;
}

练习

由于光看模板可能还不够深入,接下来会给出一些简单 DDP 的练习题以及关键步骤的讲解。

P3097 [USACO13DEC] Optimal Milking G

link。 :::info[题意简述] 一个长为 n 的序列,有 n 次操作,每次操作修改一个点权,你需要在修改后从下标集合 \{1,2,3,4,\dots,n\} 中选择一个子集,满足相邻下标不能同时被选,最大化集合中元素对应的权值之和。 :::

解法

双倍经验说是。

把原图从 ii+1 连边,建成一条链,这道题就变成了求动态树上最大独立集。和模板直接一样了。

但话又说回来了,即使是一样的题目,也建议刚接触 DDP 的读者再重新手推一遍整个过程,思考每一个步是怎么得出的,为什么这么做,这样才能有所收获。

代码应该不需要放了吧。

P5024 [NOIP 2018 提高组] 保卫王国

link。 :::info[题意简述] 给定一颗 n 个节点的树,点带权。有 m 次询问,每次询问给定两个城市的约束,每个城市要么必须选,要么不能选。你需要求出此约束下的最小权覆盖集。如果限制条件下不存在最小权覆盖集,输出 -1。询问之间相互独立。 :::

解法

算是比较板的转化了。

首先有一个非常重要的性质:最小权覆盖集 = 全集 - 最大独立集,可以说是破题的关键。

于是问题可以被转化为求树上最大独立集。

然后考虑处理选还是不选的问题,这一步也比较简单。我们只需要在每一次询问时将 g_{u,0}(不能选时)或者 g_{u,1}(必须选时)的点权修改为 -\infty 表示永远不可能从这里转移就可以了。都是模板题的基础操作。至于无解的条件就是全集 - 最大独立集的结果非常小(说明最后被迫选了 -\infty 的点)。

这道题有一个非常非常坑人的地方,就是修改点权时的更新顺序。由于你树上修改一个点的点权再 update 的时候可能会影响整颗树上所有节点的 g 值,因此更新的时候必须按这个顺序:

  1. 记录一个点的旧值
  2. 更新一个点
  3. 记录此时另一个点的旧值**
  4. 更新另一个点
  5. 记录答案
  6. 撤销后一个点的更新操作
  7. 撤销最开始的点的更新操作

    int change(int a, int x, int b, int y){
    static int tmp[2][2];
    tmp[0][0] = g[a][0], tmp[0][1] = g[a][1];
    g[a][x] = -INF;
    update(a);//先更新a的 
    
    tmp[1][0] = g[b][0], tmp[1][1] = g[b][1];
    g[b][y] = -INF;//更新完a再更新b的状态!!!! 
    update(b);//更新b的 
    
    int ans = curans() < 0 ? -1 : sum - curans();
    
    g[b][0] = tmp[1][0], g[b][1] = tmp[1][1];
    update(b);
    g[a][0] = tmp[0][0], g[a][1] = tmp[0][1];
    update(a);//撤销 
    return ans;
    }

    ::::success[完整代码]

    
    #include<bits/stdc++.h>
    #define LoveCyrene ios :: sync_with_stdio(false), cin.tie(nullptr), cout.tie(nullptr)
    #define int long long int
    using namespace std;;
    constexpr int N = 1e5 + 7891;
    constexpr int INF = 1e18;

int n, m, sum = 0; int val[N]; string ty;

vector<int> tree[N]; void add_edge(int u, int v) { tree[u].push_back(v), tree[v].push_back(u); }

struct MAT{ int a[2][2]; MAT() { a[0][0] = a[0][1] = a[1][0] = a[1][1] = -INF; } friend MAT operator * (const MAT &x, const MAT &y){ MAT z; for (int i = 0; i < 2; i++){ for (int k = 0; k < 2; k++){ for (int j = 0; j < 2; j++){ z.a[i][k] = max(z.a[i][k], x.a[i][j] + y.a[j][k]); } } } return z; } };

struct VEC{ int a[2]; VEC() { a[0] = a[1] = -INF; } VEC(int x, int y) { a[0] = x, a[1] = y; } friend VEC operator * (const MAT & y, const VEC & x){ VEC z; for (int i = 0; i < 2; i++){ for (int j = 0; j < 2; j++){ z.a[i] = max(z.a[i], x.a[j] + y.a[i][j]); } } return z; } };

struct SEG{ MAT sgt[N << 2];

void pushup(int x){
    sgt[x] = sgt[x << 1] * sgt[x << 1 | 1];
}

void modify(int x, int l, int r, int p, MAT v){
    if(l == r) { sgt[x] = v; return; }
    int mid = (l + r) >> 1;
    if(p <= mid) modify(x << 1, l, mid, p, v);
    else modify(x << 1 | 1, mid + 1, r, p, v);
    pushup(x);
}

MAT query(int x, int l, int r, int ql, int qr){
    if(ql <= l && r <= qr) return sgt[x];
    int mid = (l + r) >> 1;
    MAT res;
    res.a[0][0] = res.a[1][1] = 0;
    if(ql <= mid) res = res * query(x << 1, l, mid, ql, qr);
    if(qr > mid) res = res * query(x << 1 | 1, mid + 1, r, ql, qr);
    return res; 
}
#define Modify(x, l, r, p, v) seg.modify(x, l, r, p, v)
#define Query(x, l, r, ql, qr) seg.query(x, l, r, ql, qr)

}seg;

int dfn[N], top[N], siz[N], heavy[N], bottom[N], fa[N], dep[N]; int g[N][2]; int timi = 0;

VEC getf(int x){ MAT xx = Query(1, 1, n, dfn[x], bottom[top[x]]); VEC yy = VEC(0, -INF); return xx * yy; }

MAT mat(int u){ MAT res; res.a[0][0] = res.a[0][1] = g[u][0]; res.a[1][0] = g[u][1], res.a[1][1] = -INF; return res; }

void dfs1(int u, int f){ fa[u] = f; dep[u] = dep[f] + 1; siz[u] = 1; for (int v : tree[u]){ if(v == f) continue; dfs1(v, u); siz[u] += siz[v]; if(siz[v] > siz[heavy[u]]) heavy[u] = v; } }

void dfs2(int u, int tp){ dfn[u] = ++timi; top[u] = tp; bottom[tp] = max(bottom[tp], dfn[u]); if(heavy[u]) dfs2(heavy[u], tp);

g[u][0] = 0;
if(!g[u][1]) g[u][1] = val[u];

for (int v : tree[u]){
    if(v == fa[u] || v == heavy[u]) continue;
    dfs2(v, v);
    VEC vv = getf(v);
    g[u][0] += max(vv.a[0], vv.a[1]);
    g[u][1] += vv.a[0];
}

Modify(1, 1, n, dfn[u], mat(u));

}

void update(int x){ while(x){ int t = top[x]; VEC f1 = getf(t);

    Modify(1, 1, n, dfn[x], mat(x));

    VEC f2 = getf(t);
    if(fa[t]){
        g[fa[t]][0] += max(f2.a[0], f2.a[1]) - max(f1.a[0], f1.a[1]);
        g[fa[t]][1] += f2.a[0] - f1.a[0];
    }
    x = fa[t];
}

}

int curans(){ return max(getf(1).a[0], getf(1).a[1]); }

int change(int a, int x, int b, int y){ static int tmp[2][2]; tmp[0][0] = g[a][0], tmp[0][1] = g[a][1]; g[a][x] = -INF; update(a);//先更新a的

tmp[1][0] = g[b][0], tmp[1][1] = g[b][1];
g[b][y] = -INF;//更新完a再更新b的状态!!!! 
update(b);//更新b的 

int ans = curans() < 0 ? -1 : sum - curans();

g[b][0] = tmp[1][0], g[b][1] = tmp[1][1];
update(b);
g[a][0] = tmp[0][0], g[a][1] = tmp[0][1];
update(a);//撤销 
return ans;

}

signed main(){ LoveCyrene; cin >> n >> m >> ty; for (int i = 1; i <= n; i++) cin >> val[i], sum += val[i]; for (int i = 1, u, v; i < n; i++) cin >> u >> v, add_edge(u, v); dfs1(1, 0), dfs2(1, 1); while(m--){ int a, x, b, y; cin >> a >> x >> b >> y; cout << change(a, x, b, y) << '\n'; } return 0; }

::::
### P6021 洪水
前面的题充其量是一个简单的转化,我们并没有去真正推一个新的 DP 式子,而下面这道题就需要自己重新手推一个矩阵出来了。

[link](https://www.luogu.com.cn/problem/P6021)。
:::info[题意简述]
给定一颗 $n$ 个节点的树,根为 $1$,点带权。支持两种操作:
> `Q x` 求 $x$ 子树内的**最小割**。
> 
> `C x t` 给 $x$ 的点权增加 $t$,保证 $t$ 非负。
:::
### 解法
还是老样子,先考虑静态 DP。

设 $f_u$ 表示 $u$ 子树内的最小割,可以很快地推出静态 DP 式子为:

$$
f_u = min(w_u, \sum\limits_{v \in son(u)}f_v)
$$

下一步就是考虑树剖优化。令 $g_u$ 表示 $u$ 子树内轻儿子的最小割,$s$ 表示重儿子,于是简化式子得到:
$$
f_u = min(w_u, g_u+f_s)
$$

类比上面,可以发现此时的运算无非就是从最大值变成了最小值,因此我们依着葫芦画一个瓢,尝试构建一个 **min-plus** 矩阵。于是你可以高兴地得出:

$$
f_u = \begin{pmatrix} g_u & w_u \end{pmatrix} \otimes \begin{pmatrix} f_s \\ 0 \end{pmatrix}
$$

但光是这样就够了吗?还不行。

由于此时的 $f$ 数组只有一个维度,因此我们得到的方程是一个行向量左乘列向量的形式,但由于两个行向量之间是不能做乘法的,因此没法直接用线段树合并。怎么办呢?

解决方式也很简单,既然我们习惯矩阵乘列向量,那么把向量变成矩阵,加一个无关紧要的维度就可以了!

于是我们在行向量的下面加点货,使得无论怎么乘,第二维都不会干扰结果,我们可以始终放心地取第一维。基于此,就可以进一步地写出我们熟悉的形式:

$$
\begin{pmatrix} f_u \\ 0 \end{pmatrix}= \begin{pmatrix} g_u & w_u \\ +\infty &  0\end{pmatrix} \otimes \begin{pmatrix} f_s \\ 0 \end{pmatrix}
$$

然后我们就又可以愉快地写个模板了。
::::success[完整代码]
```cpp
#include<bits/stdc++.h>
#define LoveCyrene ios :: sync_with_stdio(false), cin.tie(nullptr), cout.tie(nullptr)
#define int long long int
using namespace std;;
constexpr int N = 2e5 + 7918;
constexpr int INF = 1e18;

int n, m;
vector<int> tree[N];
int val[N];

void add_edge(int u, int v){ tree[u].push_back(v), tree[v].push_back(u); }

struct MAT{
    int a[2][2];
    MAT(){ a[0][0] = a[0][1] = a[1][0] = a[1][1] = INF; }
    friend MAT operator * (const MAT& x, const MAT& y){
        MAT z;
        for (int i = 0; i < 2; i++){
            for (int k = 0; k < 2; k++){
                for (int j = 0; j < 2; j++){
                    z.a[i][k] = min(z.a[i][k], x.a[i][j] + y.a[j][k]);
                }
            }
        }
        return z;
    }
};

struct VEC{
    int a[2];
    VEC(){ a[0] = a[1] = INF; }
    VEC(int x, int y) { a[0] = x, a[1] = y; }
    friend VEC operator * (const MAT& x, const VEC& y){
        VEC z;
        for (int i = 0; i < 2; i++){
            for (int j = 0; j < 2; j++){
                z.a[i] = min(z.a[i], x.a[i][j] + y.a[j]);
            }
        }
        return z;
    }   
};

struct SEG{
    MAT sgt[N << 2];

    void pushup(int x){
        sgt[x] = sgt[x << 1] * sgt[x << 1 | 1];
    }

    void modify(int x, int l, int r, int p, MAT v){
        if(l == r) { sgt[x] = v; return; }
        int mid = (l + r) >> 1;
        if(p <= mid) modify(x << 1, l, mid, p, v);
        if(p > mid) modify(x << 1 | 1, mid + 1, r, p, v);
        pushup(x);
    }

    MAT query(int x, int l, int r, int ql, int qr){
        if(ql <= l && r <= qr) return sgt[x];
        int mid = (l + r) >> 1;
        MAT res;
        res.a[0][0] = res.a[1][1] = 0;
        if(ql <= mid) res = res * query(x << 1, l, mid, ql, qr);
        if(qr > mid) res = res * query(x << 1 | 1, mid + 1, r, ql, qr);
        return res;
    }
    #define Modify(x, l, r, p, v) seg.modify(x, l, r, p, v)
    #define Query(x, l, r, ql, qr) seg.query(x, l, r, ql, qr)
}seg; 

int dfn[N], top[N], bottom[N], dep[N], siz[N], heavy[N], fa[N];
int timi = 0;
int g[N];

VEC getf(int x){
    MAT xx = Query(1, 1, n, dfn[x], bottom[top[x]]);
    VEC yy = VEC(INF, 0);
    return xx * yy;
}

MAT mat(int x){
    MAT res;
    res.a[0][0] = g[x], res.a[0][1] = val[x];
    res.a[1][0] = INF, res.a[1][1] = 0;
    return res;
}

void dfs1(int u, int f){
    fa[u] = f;
    dep[u] = dep[f] + 1;
    siz[u] = 1;
    for (int v : tree[u]){
        if(v == f) continue;
        dfs1(v, u);
        siz[u] += siz[v];
        if(siz[v] > siz[heavy[u]]) heavy[u] = v;
    }
}

void dfs2(int u, int tp){
    dfn[u] = ++timi;
    top[u] = tp;
    bottom[tp] = max(bottom[tp], dfn[u]);
    if(heavy[u]) dfs2(heavy[u], tp);

    for (int v : tree[u]){
        if(v == fa[u] || v == heavy[u]) continue;
        dfs2(v, v);

        VEC vv = getf(v);
        g[u] += vv.a[0];
    }

    Modify(1, 1, n, dfn[u], mat(u)); 
}

void update(int x, int v){
    val[x] += v;
    while(x){
        int t = top[x];
        VEC f1 = getf(t);

        Modify(1, 1, n, dfn[x], mat(x));

        VEC f2 = getf(t);
        if(fa[t]) g[fa[t]] += f2.a[0] - f1.a[0];
        x = fa[t];
    }
}

signed main(){
    LoveCyrene;
    cin >> n;
    for (int i = 1; i <= n; i++) cin >> val[i];
    for (int i = 1, u, v; i < n; i++) cin >> u >> v, add_edge(u, v);

    dfs1(1, 0), dfs2(1, 1);
    cin >> m; while(m--){
        char op; cin >> op;
        if(op == 'C'){ int x, t; cin >> x >> t, update(x, t); }
        else if(op == 'Q'){ int x; cin >> x, cout << getf(x).a[0] << '\n'; }
        else cout << "Love Cyrene. \n";
    }
    return 0;
}

::::

P4115 Qtree4

下面这道题则更有趣,它的转移不再拘泥于矩阵这一种单一的形式,而是变成了其它形式的运算,体现了 DDP 的灵活多变。

link。 ::::info[题意简述] 给出一棵边带权的节点数量为 n 的树,初始树上所有节点都是白色。有两种操作:

C x,改变节点 x 的颜色,即白变黑,黑变白。

A,询问树中最远的两个白色节点的距离,这两个白色节点可以重合(此时距离为 0)。 ::::

解法

这道题的正解是动态树分治,但 DDP 也可以做。

如果我们用和刚才一样的方式去套模板,你会发现一件事:似乎很难写出转移方程的矩阵形式。

我们知道,矩阵作为一种满足结合律的运算,用它合并信息可以有效加速转移的过程。但换个角度来看,也可以理解为:只要满足结合律,配合数据结构,都可以实现和矩阵等价的效果。

于是我们可以大胆将矩阵换成另一种可合并运算并使用线段树维护。怎么换呢?

还是列出静态 DP 式子。令 f_u 表示 u 子树中离 u 最远的白点,ans_u 表示 u 子树中最远的白点之间的距离,则 DP 式子为:

f_u = \max\limits_{v \in son(u)}(f_u,f_v + w_{u,v}) ans_u = \max(\max\limits_{v}ans_v,\max\limits_{v1 \ne v2}(f_{v_1} + w_{u, v_1} + f_{v_2} + w_{u, v_2}),[u\text{为白点}]\cdot f_u)

初始化 f_uans_u0(白点)或 -\infty(黑点)。

这个式子实在太抽象了,还是考虑分离轻儿子。对于每一个节点 u,设它和它的轻儿子构成的子树中的 f_uX,它和它的轻儿子构成的子树中的 ans_uY,重儿子为 s,于是乎转移变成:

f_u = \max(X, f_s + w_{u,s}) ans_u = \max(Y,ans_s, X + f_s + w_{u,s})

先别管轻儿子,上面这一步可以说明,只要我们知道重儿子的 fans,也就是对重链维护 fans,就可以得出想要的信息。那么怎么维护这俩信息呢?假设已知两段重链区间 A(位置更浅)和 B(位置更深),现在要合并出一个新的区间 C

先来看 f 值的变化,有这两种情况:

也就是说,我们需要维护两个字段,一个是 f 本身,一个是区间长度 len

再来看 ans 的变化,要分成三种情况:

于是我们可以将上面的整个过程封装成一个多元组为 (len,f,g,ans,rr)。每次合并的规律就是进行如下形式的运算:

len_C = len_A + w + len_B f_C = \max(f_A, len_A + w + f_B) g_C = \max(g_A, len_A + w + g_B) ans_C = \max(ans_A, ans_B, g_C + w + f_B) rr_C = rr_B

于是整个过程就完美替代了上面的矩阵运算。

现在重儿子处理了,轻儿子呢?

其实很简单,开俩 multiset,一个存从 u 出发到轻儿子 v,以及它的 f_v 之和;一个存从一个存轻儿子 vans。预处理 DFS 完轻儿子或者每次 update 之后在 multiset 内部更新一下就完了。需要的时候刨一下里面的最大值或者次大值什么的就行了。

确实有点难,慢慢理解。 ::::success[完整代码]

#include<bits/stdc++.h>
#define LoveCyrene ios :: sync_with_stdio(false), cin.tie(nullptr), cout.tie(nullptr)
#define int long long
using namespace std;

constexpr int N = 2e5 + 7918;
constexpr int INF = 1e18;

int n, m;
struct Edge{ int v, w;};
vector<Edge> tree[N];
bool white[N];
int hw[N], fw[N];//到重儿子/父节点的边权

void add_edge(int u, int v, int w){
    tree[u].push_back({v, w});
    tree[v].push_back({u, w});
}

struct Node{
    int len, f, g, ans, rr;
    //f:从区间最左端,向右能走到的最大距离 
    //g:从区间最右端,向左能走到的最大距离 
    Node(){
        len = rr = 0;
        f = g = ans = -INF;
    }
};

struct SEG{
    Node sgt[N << 2];

    //合并两个区间,中间边权为 w
    Node merge(const Node& X, const Node& Y, int w){//X浅Y深 
        Node Z;
        Z.len = X.len + w + Y.len;
        Z.f = max(X.f, X.len + w + Y.f);
        Z.g = max(Y.g, Y.len + w + X.g);
        Z.ans = max({X.ans, Y.ans, X.g + w + Y.f});
        Z.rr = Y.rr;
        return Z;
    }

    void pushup(int x){
        sgt[x] = merge(sgt[x << 1], sgt[x << 1 | 1], hw[ sgt[x << 1].rr ]);
    }

    void modify(int x, int l, int r, int p, Node v){
        if(l == r){ sgt[x] = v; return; }
        int mid = (l + r) >> 1;
        if(p <= mid) modify(x << 1, l, mid, p, v);
        if(p > mid) modify(x << 1 | 1, mid + 1, r, p, v);
        pushup(x);
    }

    Node query(int x, int l, int r, int ql, int qr){
        if(ql <= l && r <= qr) return sgt[x];
        int mid = (l + r) >> 1;
        if(qr <= mid) return query(x << 1, l, mid, ql, qr);
        if(ql > mid) return query(x << 1 | 1, mid + 1, r, ql, qr);
        Node L = query(x << 1, l, mid, ql, qr);
        Node R = query(x << 1 | 1, mid + 1, r, ql, qr);
        return merge(L, R, hw[L.rr]);
    }
    #define Modify(x, l, r, p, v) seg.modify(x, l, r, p, v)
    #define Query(x, l, r, ql, qr) seg.query(x, l, r, ql, qr)
}seg;

int dfn[N], top[N], bottom[N], dep[N], siz[N], heavy[N], fa[N];
int timi = 0;

multiset<int, greater<int>> lp[N], ld[N];
//lp:轻儿子的f+w ld:轻儿子的ans 

void del(multiset<int, greater<int>>& ms, int val){
    auto it = ms.find(val);
    if(it != ms.end()) ms.erase(it);
}

int maxx(multiset<int, greater<int>>& ms){//轻儿子的最大值 
    //相当于 max(ansv) 
    return ms.empty() ? -INF : *ms.begin();
}

int summ(multiset<int, greater<int>>& ms){//轻儿子的最大+次大 
    //相当于 max(fv1 + w(u, v1) + fv2 + w(u,v2) 
    if(ms.empty()) return -INF;
    auto it = ms.begin();
    int mx1 = *it; ++it;
    if(it == ms.end()) return -INF;
    return mx1 + *it;
}

//计算节点u的信息
Node mat(int u){
    Node cur;
    cur.rr = u;
    int x = white[u] ? 0 : -INF;
    x = max(x, maxx(lp[u]));

    int y = white[u] ? 0 : -INF;
    y = max(y, max(maxx(ld[u]), summ(lp[u])));

    cur.f = cur.g = x;
    cur.ans = y;
    cur.len = 0;
    return cur;
}

void dfs1(int u, int f){
    fa[u] = f;
    dep[u] = dep[f] + 1;
    siz[u] = 1;
    for(auto e : tree[u]){
        if(e.v == f) continue;
        fw[e.v] = e.w;
        dfs1(e.v, u);
        siz[u] += siz[e.v];
        if(siz[e.v] > siz[heavy[u]]){
            heavy[u] = e.v;
            hw[u] = e.w;
        }
    }
}

void dfs2(int u, int tp){
    dfn[u] = ++timi;
    top[u] = tp;
    bottom[tp] = max(bottom[tp], dfn[u]);

    if(heavy[u]) dfs2(heavy[u], tp);

    for(auto e : tree[u]){
        int v = e.v;
        if(v == fa[u] || v == heavy[u]) continue;
        dfs2(v, v);

        Node vc = Query(1, 1, n, dfn[v], bottom[v]);
        if(vc.f > -INF) lp[u].insert(vc.f + e.w);
        if(vc.ans > -INF) ld[u].insert(vc.ans);
    }

    Modify(1, 1, n, dfn[u], mat(u));
}

void update(int x){
    white[x] xor_eq 1;

    while(x){
        int t = top[x];
        Node f1 = Query(1, 1, n, dfn[t], bottom[t]);
        Modify(1, 1, n, dfn[x], mat(x));
        Node f2 = Query(1, 1, n, dfn[t], bottom[t]);

        if(fa[t]){
            int w = fw[t];
            if(f1.f > -INF) del(lp[fa[t]], f1.f + w);
            if(f1.ans > -INF) del(ld[fa[t]], f1.ans);
            if(f2.f > -INF) lp[fa[t]].insert(f2.f + w);
            if(f2.ans > -INF) ld[fa[t]].insert(f2.ans);
        }
        x = fa[t];
    }
}

int curans(){ return Query(1, 1, n, dfn[1], bottom[1]).ans; }

signed main(){
    LoveCyrene;
    cin >> n;
    for(int i = 1, u, v, w; i < n; i++) cin >> u >> v >> w, add_edge(u, v, w);
    fill(white + 1, white + n + 1, true);

    dfs1(1, 0), dfs2(1, 1);

    cin >> m; while(m--){
        char op; cin >> op;
        if(op == 'C'){ int x; cin >> x; update(x); }
        else if (op == 'A'){
            if(curans() < 0) cout << "They have disappeared.\n";
            else cout << curans() << '\n';
        }
        else cout << "Love Cyrene. \n";
    }
    return 0;
}

::::

P2056 [ZJOI2007] 捉迷藏

link。

上面那道题难度跨度有点大,因此给个福利。双倍经验,和上一道题完全一样(甚至还简化了,因为边权都是 1),题意和代码就都不放了。

动态 DP 的优化

简介

最后简单提一点 DDP 的优化问题。其实 DDP 优化这个说法不是很好,因为 DDP 本身已经是一种加速 DP 的优化技巧了,再套一层优化,语法上感觉有点奇怪。DDP 优化的本质,不过是将树剖线段树替换成另一种数据结构,就跟倍增 LCA 与树剖 LCA 的关系一样,个人更喜欢看作是同一种问题的不同求解过程,只是性能差异罢了。

常见的 DDP 优化有三种:LCT 优化,全局平衡二叉树优化,以及静态 Top Tree 优化。其中 LCT 常数太大在三者中最劣,静态 Top Tree 太难,因此本文着重讲解第二种,也就是全局平衡二叉树。

由于 DDP 优化的是数据结构,DP 转移的逻辑和过程与普通 DDP 完全一样,因此这部分内容可以看作是纯数据结构的内容。

全局平衡二叉树

先来认识一下这个东西。

全局平衡二叉树(Global Balanced Tree,后文简称 GBT)本质上是一种由多颗重链剖分形成的二叉树构成的数据结构。它基于树剖实现,借用了 LCT 的思想,使得构建出来的树满足根节点到任意节点的距离都不超过 O(\log{n})。这也是它名字中“全局平衡”的由来。

原理

比如有这么一棵树,红色的边表示重链。 GBT 的原理就是,对每一条重链开一个加权平衡二叉树,权重为轻儿子的子树大小之和 +1。这棵树满足:

  1. 对于任意子树,左右子树的权值和都不超过该子树总权值的一半。
  2. 中序遍历顺序为原重链顺序。

比如,下面就是根据上图构建出来的一种 GBT,蓝色的边表示虚边。

可以发现,基于此构建出来的 GBT 的树高是不会超过 O(\log{n}) 的。 ::::info[简单证明] 对于链内二叉树,由于我们每次选择的是带权中位数做当前的根,因此每往下走一次,权值至少减半,也就意味着链内二叉树的树高不会超过 O(\log{n})

对于轻儿子,由于每跳一次虚边,子树大小至少会加倍(重剖的性质),因此从一个轻儿子跳到根,最多也只会跳 O(\log{n}) 次虚边。

也就是说,任意节点的深度至多为链内二叉树深度 + 虚边数量 = O(\log{n}) 级别。 :::: 利用这个性质,我们就可以在每次修改时,直接暴力往上跳,复杂度至多只有 O(\log{n}),比树剖线段树直接少了一个 \log,简直太强了。

建立

由此,我们也可以得出 GBT 的构建逻辑。

首先还是剖一遍原树,然后预处理出每个节点的权值。对于每一条重链,我们将它提取出来,用类似线段树的方式,找出当前区间的带权中位数,这个点就作为当前的节点,左右部分用同样的方式递归地去找就可以了。

int build(const vector<int> & arr, int l, int r){
    if(l > r) return 0;
    int tot = 0, cur = 0, mid;
    for (int i = l; i <= r; i++) tot += w(arr[i]);
    for (int i = l; i <= r; i++){
        cur += w(arr[i]);
        if(cur * 2 >= tot) {
            mid = i; break;
        }
    }
    int u = arr[mid];
    ls(u) = build(arr, l, mid - 1);
    rs(u) = build(arr, mid + 1, r);
    if(ls(u)) fa(ls(u)) = u;
    if(rs(u)) fa(rs(u)) = u;
    return u;
}

void buildGBT(int n){
    for (int i = 1; i <= n; i++){
        w(i) = 1;
        for (int v : tree[i]){
            if(v == pa(i) || v == heavy(i)) continue;
            w(i) += siz(v);
        }
    }
    for (int i = 1; i <= n; i++){
        if(top(i) != i) continue;
        vector<int> chain;
        for (int x = i; x; x = heavy(x)) chain.push_back(x);
        int rt = build(chain, 0, chain.size() - 1);
        hvtop(i) = rt;//hvtop表示当前重链的GBT根 
        fa(rt) = pa(i);//将当前链内二叉树根节点的父亲设为原树上的父亲,作为一个虚儿子 
    }
}

剩下的处理部分就和树剖线段树的逻辑大差不差了。

练习:P4751 【模板】动态 DP(加强版)

套一下维护的信息即可,不再赘述。 ::::success[完整代码]

#include<bits/stdc++.h>
#define LoveCyrene ios :: sync_with_stdio(false), cin.tie(nullptr), cout.tie(nullptr)
using namespace std;;
constexpr int N = 1e6 + 7891;
constexpr int INF = 1e9;

vector<int> tree[N];

void add_edge(int u, int v){
    tree[u].push_back(v);
    tree[v].push_back(u);
}

struct MAT{//代码太长懒得分两个结构体存矩阵和向量了 
    int a[2][2];
    MAT(){ a[0][0] = a[0][1] = a[1][0] = a[1][1] = -INF; }
    MAT(int g0, int g1){
        a[0][0] = a[0][1] = g0;
        a[1][0] = g1, a[1][1] = -INF;
    }
    friend MAT operator * (const MAT &x, const MAT &y){
        MAT z;
        for (int i = 0; i < 2; i++){
            for (int k = 0; k < 2; k++){
                for (int j = 0; j < 2; j++){
                    z.a[i][j] = max(z.a[i][j], x.a[i][k] + y.a[k][j]);
                }
            }
        }
        return z;
    }
};

struct TreeChainSeg{
    int dfn[N], pa[N], top[N], heavy[N], siz[N];
    int tim = 0;
    void dfs1(int u, int f){
        pa[u] = f;
        siz[u] = 1;
        for (int v : tree[u]){
            if (v == f) continue;
            dfs1(v, u);
            siz[u] += siz[v];
            if (siz[v] > siz[heavy[u]]) heavy[u] = v;
        }
    }
    void dfs2(int u, int tp){
        dfn[u] = ++tim;
        top[u] = tp;
        if (heavy[u]) dfs2(heavy[u], tp);
        for (int v : tree[u]){
            if (v == pa[u] || v == heavy[u]) continue;
            dfs2(v, v);
        }
    }
    #define dfn(x) TCS.dfn[x]
    #define pa(x) TCS.pa[x]
    #define top(x) TCS.top[x]
    #define dep(x) TCS.dep[x]
    #define siz(x) TCS.siz[x]
    #define heavy(x) TCS.heavy[x]
    #define rnk(x) TCS.rnk[x]
}TCS;

struct GloBalTree{
    int fa, son[2];
    int val;//点权 
    MAT sum;//子树乘积
    MAT g;//转移矩阵
    int hvtop;//重链的GBT根
    int w;//GBT权值 
    int f0, f1;//轻儿子的贡献 
    #define fa(x) GBT[x].fa
    #define ls(x) GBT[x].son[0]
    #define rs(x) GBT[x].son[1]
    #define val(x) GBT[x].val
    #define sz(x) GBT[x].sz
    #define sum(x) GBT[x].sum
    #define g(x) GBT[x].g
    #define hvtop(x) GBT[x].hvtop
    #define w(x) GBT[x].w
    #define f0(x) GBT[x].f0
    #define f1(x) GBT[x].f1
}GBT[N];

vector<int> light[N];

void pushup(int x){
    sum(x) = sum(ls(x)) * g(x) * sum(rs(x));
} 

int build(const vector<int> & arr, int l, int r){
    if(l > r) return 0;
    int tot = 0, cur = 0, mid;
    for (int i = l; i <= r; i++) tot += w(arr[i]);
    for (int i = l; i <= r; i++){
        cur += w(arr[i]);
        if(cur * 2 >= tot) {
            mid = i; break;
        }
    }
    int u = arr[mid];
    ls(u) = build(arr, l, mid - 1);
    rs(u) = build(arr, mid + 1, r);
    if(ls(u)) fa(ls(u)) = u;
    if(rs(u)) fa(rs(u)) = u;
    return u;
}

void buildGBT(int n){
    for (int i = 1; i <= n; i++){
        w(i) = 1;
        for (int v : tree[i]){
            if(v == pa(i) || v == heavy(i)) continue;
            w(i) += siz(v);
        }
    }
    for (int i = 1; i <= n; i++){
        if(top(i) != i) continue;
        vector<int> chain;
        for (int x = i; x; x = heavy(x)) chain.push_back(x);
        int rt = build(chain, 0, chain.size() - 1);
        hvtop(i) = rt;//hvtop表示当前重链的GBT根 
        fa(rt) = pa(i);//将当前链内二叉树根节点的父亲设为原树上的父亲,作为一个虚儿子 
    }
}

void buildDP(int u){//初始化DP 
    if(!u) return;
    buildDP(ls(u)), buildDP(rs(u));
    for (int v : light[u]) buildDP(hvtop(v));

    f0(u) = f1(u) = 0;
    for (int v : light[u]){
        int ff0 = sum(hvtop(v)).a[0][0];
        int ff1 = sum(hvtop(v)).a[1][0];
        f0(u) += max(ff0, ff1);
        f1(u) += ff0;
    }
    g(u) = MAT(f0(u), f1(u) + val(u));
    pushup(u);
}

void change(int x, int v){
    val(x) = v;
    while(x){
        g(x) = MAT(f0(x), f1(x) + val(x));
        MAT y = sum(x);
        pushup(x);

        if(x == hvtop(top(x)) && fa(x)){
            f0(fa(x)) += max(sum(x).a[0][0], sum(x).a[1][0]) - max(y.a[0][0], y.a[1][0]);
            f1(fa(x)) += sum(x).a[0][0] - y.a[0][0];
        }
        x = fa(x);
    }
}

void pre(int n){
    TCS.dfs1(1, 0), TCS.dfs2(1, 1);
    for (int i = 1; i <= n; i++){
        for (int v : tree[i]){
            if(v == pa(i) || v == heavy(i)) continue;
            light[i].push_back(v);
        }
    }//重剖+预处理轻儿子

    sum(0).a[0][0] = sum(0).a[1][1] = 0;//设置空节点矩阵为单位元 
    buildGBT(n);//建出GBT 
    buildDP(hvtop(1));//初始化DP 
}

int curans(int &lst){ return lst = max(sum(hvtop(1)).a[0][0], sum(hvtop(1)).a[1][0]); }

signed main(){
    LoveCyrene;
    int n, m, lst = 0; cin >> n >> m;
    for (int i = 1; i <= n; i++) cin >> val(i);
    for (int i = 1, u, v; i < n; i++) cin >> u >> v, add_edge(u, v);
    pre(n);
    for (int u, v; m; m--) cin >> u >> v, change(u xor lst, v), cout << curans(lst) << '\n';
    return 0;
}

::::