【学习笔记】斜率优化 DP

· · 算法·理论

前言

更好的阅读体验?

本文介绍了斜率优化 DP 的推导过程及其基础应用,如果想要有更多了解请自行查阅其它文章。

算法介绍

在平时的 DP 过程中,我们可能会遇到这样的 DP 式子:

dp_{i}=\min_{j=i-x}^{i-y}(dp_j+a_i\times b_j+c_i+d_j+M)

单调队列优化 DP 拼劲全力无法战胜,只能请它大哥出场了。

推导一

动态规划当然要考虑最优决策点的位置呀!所以我们假设现在有 j_1,j_2(j_1<j_2) 两个决策点,如果 j_2 更优秀,应该满足的条件是酱紫的:

\begin{aligned} dp_{j_2}+a_i\times b_{j_2}+c_i+d_{j_2}+M&\le dp_{j_1}+a_i\times b_{j_1}+c_i+d_{j_1}+M\\ dp_{j_2}+a_i\times b_{j_2}+d_{j_2}&\le dp_{j_1}+a_i\times b_{j_1}+d_{j_1}\\ -a_i\times(b_{j_1}-b_{j_2})&\le (dp_{j_1}+d_{j_1})-(dp_{j_2}+d_{j_2})\\ \end{aligned}

如果 b_{j_1}-b_{j_2}\ne 0,那么我们就会得到这个不等式:

a_i\ge \frac{(dp_{j_1}+d_{j_1})-(dp_{j_2}+d_{j_2})}{b_{j_1}-b_{j_2}}

如果 b_{j_1}-b_{j_2}=0,你可以认为上面那个值是 \infty

这里我们令 X(i)=b_i,Y(i)=dp_{i}+d_{i},那么不等式变为:

a_i\ge \frac{Y(j_1)-Y(j_2)}{X(j_1)-X(j_2)}

(X(j),Y(j)) 作为 j 的对应点放入平面。如果 j_1j_2 构成的直线斜率小于等于 a_i(这里 j_1,j_2 作为决策点出现),那么 j_2 优于 j_1;否则 j_1 优于 j_2

假设 j\in [i-x,i-y],那么下图的点就是对应到平面上的散点:

我们把目光放到 A,B,C 三个点上,这里假设 A,B 构成的直线斜率是 k_1B,C 构成的直线斜率是 k_2,此处我们钦定 k_1>k_2

也就是说,B 永远不会存为最优点,就要淘汰掉(太馋人,哦不太残忍了)。

同样的道理,这里看似有这么多点,但实则只有下凸包上的点会用到:

又因为下凸包的斜率是单调递增的,如果 j 前面的斜率都是小于等于 a_i,后面的斜率都是 >a_i,那么 j 就是当前的最优决策点。

那不对呀!上凸包怎么能被忽略呢?!

所以如果上面不等式是 a_i\le \frac{Y(j_1)-Y(j_2)}{X(j_1)-X(j_2)} 的话,那么决策点就在上凸包上啦!

推导二

TA 回来了:

dp_{i}=\min_{j=i-x}^{i-y}(dp_j+a_i\times b_j+c_i+d_j+M)

先假装看不见 \min

dp_{i}=dp_j+a_i\times b_j+c_i+d_j+M

进一步:

dp_j+d_j=-a_i\times b_j+dp_i-c_i-m

y=dp_j+d_j,k=-a_i,x=b_j,b=dp_i-c_i-m,那么上式就变成了 y=kx+b,此时最小化 dp_i 就相当于最小化 b

怎样移向呢?看法则:

我们把每一个 x,y 当成一个点对应到平面上,这就是我们的决策点。接着我们用 k=-a[i] 的直线去对准每一个点(如图),看看哪条直线的截距(也就是 b)最小就好了。

你看,最优决策点的位置不变,说明最优决策点还是在下凸包的斜率单峰处。

当然,最大化 b 就是上凸包。

总结

其实它们本质是相同的,可以相辅相成地来理解。

接下来,明白了最优决策点在什么位置,那该怎么快速地找最优决策点呢?

例题讲解

P3195 [HNOI2008] 玩具装箱

答案就是 dp_n

如果当前状态为 dp_i,若存在决策点 j_1,j_2(j_1<j_2),且 j_2 优于 j_1,则:

\begin{aligned} dp_{j_2}+S_i^2+S_{j_2}^2+M^2-2S_iS_{j_2}-2S_iM+2S_{j_2}M &\le dp_{j_1}+S_i^2+S_{j_1}^2+M^2-2S_iS_{j_1}-2S_iM+2S_{j_1}M \\ dp_{j_2}+S_{j_2}^2-2S_iS_{j_2}+2S_{j_2}M &\le dp_{j_1}+S_{j_1}^2-2S_iS_{j_1}+2S_{j_1}M \\ 2S_i(S_{j_1}-S_{j_2}) &\ge \big(dp_{j_1}+S_{j_1}^2+2S_{j_1}M\big)-\big(dp_{j_2}+S_{j_2}^2+2S_{j_2}M\big) \\ &\because C_i\ge 1,\ j_1\ne j_2 \\ &\therefore S_{j_1}-S_{j_2}\ne 0 \\ 2S_i &\ge \frac{\big(dp_{j_1}+S_{j_1}^2+2S_{j_1}M\big)-\big(dp_{j_2}+S_{j_2}^2+2S_{j_2}M\big)}{S_{j_1}-S_{j_2}} \end{aligned}

X(j) = S_j,\ Y(j) = dp[j] + S_j^2 + 2S_jM,得:

2S_i \ge \frac{Y(j_1) - Y(j_2)}{X(j_1) - X(j_2)}

满足此条件时有 j_2 优于 j_1,则根据原理处的推导,只需要维护 j \in [1,i-1] 对应的点集 (X(j),Y(j)) 的下凸包即可。

又由于 S_i 单调递增,所以状态转移具有决策单调性,可以用单调队列维护,即不是 i 的最优决策点的点不会再是 i+1 的最优决策点。

注意:

:::success[Code]

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

const int N = 5e4 + 10;
int n, L, c[N];
int dp[N], s[N], a[N], b[N];
int q[N], h = 1, t = 1;
int X(int p) { return b[p]; }
int Y(int p) { return dp[p] + b[p] * b[p]; }
double slope(int a, int b) {
    if (X(a) == X(b)) {
        if (Y(a) == Y(b)) return 0; // 斜率无法比较
        if (Y(a) > Y(b)) return -1e18; // 斜率为负无穷
        else return 1e18; // 斜率为正无穷
    }
    return (Y(a) - Y(b)) * 1.00 / (X(a) - X(b));
}

signed main() {
    scanf("%lld%lld", &n, &L);
    for (int i = 1; i <= n; i++) scanf("%lld", &c[i]), s[i] = s[i - 1] + c[i];
    for (int i = 0; i <= n; i++) a[i] = s[i] + i, b[i] = s[i] + i + L + 1;
    q[1] = 0;
    for (int i = 1; i <= n; i++) {
        // 由于凸包的斜率一定单调递增,于是把队头斜率小于 2 * a[i] 的点删除
        while (h < t && slope(q[h], q[h + 1]) <= 2 * a[i]) h++;
        int j = q[h];
        dp[i] = dp[j] + (a[i] - b[j]) * (a[i] - b[j]); // 状态转移方程
        // 把队尾的斜率小于当前点的坐标删除,放入当前点
        while (h < t && slope(q[t - 1], q[t]) >= slope(q[t], i)) t--;
        q[++t] = i;
    }
    printf("%lld\n", dp[n]);
    return 0;
}

:::

P2365 [IOI 2002] 任务安排

答案就是 \min\{dp_{n,j}(j\ge 1)\}

一个完美的 \mathcal{O}(n^3) 的超时做法。

我们发现为了知道当前是第几组,我们耗费了一维空间去存储这个数据,所以可以考虑把这一维扔掉。

但是这样我们就不知道机器启动过几次了!

别慌,冷静分析一下你会发现,如果我们要从 dp_j 转移到 dp_i 的话,由于第 j+1\sim i 都是在同一组内完成的,我们只需要把 sj+1 后的影响补充到费用中就可以了!

答案就是 \min\{dp_{n}\}

:::success[Code]

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

int n, s, sc[maxn], st[maxn], f[maxn];
signed main() {
    cin >> n >> s;
    for (int i = 1; i <= n; i++) cin >> st[i] >> sc[i];
    for (int i = 1; i <= n; i++) st[i] += st[i - 1], sc[i] += sc[i - 1];
    memset(f, 0x3f, sizeof(f));
    f[0] = 0;
    for (int i = 1; i <= n; i++)
        for (int j = 0; j < i; j++)
            f[i] = min(f[i], f[j] + s * (sc[n] - sc[j]) + st[i] * (sc[i] - sc[j]));
    cout << f[n] << endl;
    return 0;
}

:::

P10979 任务安排 2

状态转移方程中有含 i 项与含 j 项的乘积,斜率优化没跑了。

依旧扔掉 \min 并且拆开括号:

\begin{aligned} dp_i = dp_j + T_i \times F_i - T_i \times F_j + s \times F_n - s \times F_j\\ dp_j = (T_i + s) \times F_j + dp_i - T_i \times F_i - s \times F_n \end{aligned}

y = dp_j, k = T_i + s, x = F_j, b = dp_i - T_i \times F_i - s \times F_n,一定要保证 b 中不含关于 j 的项,因为这样才能将问题转化为最小化 b

如果将 j 在平面中对应的点设为 (x,y),问题就转化成了拿斜率为 k 的线去对应每个点,最小化截距 b

易知最优决策点是下凸包上最后一个斜率小于等于 k 的右端点。

在这道题中,tim_i \ge 1,所以 T_i 单调递增,进而 k = T_i + s 单调递增,如果下凸包上某两个点构成的直线的斜率已经不是 i 前面最大的小于等于 T_i + s 的斜率,由于 T_{i+1} + s > st[i],所以该直线的右端点也不可能作为 i + 1 的最优决策点。这种特性使得本题可以使用单调队列在维护下凸包的同时维护 i 的最优决策点。

:::success[Code]

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

int n, s, l = 1, r = 1, q[N];
int c[N], tim[N], f[N];
int dy(int j, int k) { return f[j] - f[k]; }
int dx(int j, int k) { return c[j] - c[k]; }

signed main() {
    scanf("%lld%lld", &n, &s);
    for (int i = 1; i <= n; i++) {
        scanf("%lld%lld", tim + i, c + i);
        tim[i] += tim[i - 1];
        c[i] += c[i - 1];
    }

    q[1] = 0;
    for (int i = 1; i <= n; i++) {
        while (l < r && dy(q[l + 1], q[l]) <= (tim[i] + s) * dx(q[l + 1], q[l])) l++;

        int j = q[l];
        f[i] = f[j] + tim[i] * (c[i] - c[j]) + s * (c[n] - c[j]);

        while (l < r && dy(i, q[r]) * dx(q[r], q[r - 1]) <= dy(q[r], q[r - 1]) * dx(i, q[r])) r--;
        q[++r] = i;
    }

    printf("%lld\n", f[n]);
    return 0;
}

:::

P5785 [SDOI2012] 任务安排

在这道题中,tim 可能为负,所以 T[i]+s 不具有单调性。我们就需要记下来 i 前面的点的下凸包上的所有点,根据凸包上斜率的单调性进行二分。

:::success[Code]

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

int n, s, h = 1, t = 0, q[N];
int c[N], tim[N], f[N];
int dy(int j, int k) { return f[j] - f[k]; }
int dx(int j, int k) { return c[j] - c[k]; }
int find(int i) {
    if (h == t) return q[h];
    int l = h, r = t, res = h;
    while (l <= r) {
        int mid = (l + r) >> 1;
        if (dy(q[mid + 1], q[mid]) <= dx(q[mid + 1], q[mid]) * (tim[i] + s)) l = mid + 1;
        else r = mid - 1, res = mid;
    }
    return q[res];
}

signed main() {
    scanf("%lld%lld", &n, &s);
    for (int i = 1; i <= n; i++) {
        scanf("%lld%lld", tim + i, c + i);
        tim[i] += tim[i - 1];
        c[i] += c[i - 1];
    }

    for (int i = 1; i <= n; i++) {
        while (h < t && dy(i - 1, q[t]) * dx(q[t], q[t - 1]) <= dx(i - 1, q[t]) * dy(q[t], q[t - 1])) t--;
        q[++t] = i - 1;
        int j = find(i);
        f[i] = f[j] + tim[i] * (c[i] - c[j]) + s * (c[n] - c[j]);
    }
    printf("%lld\n", f[n]);
    return 0;
}

:::

参考资料