洛谷P17372 [ECNA 2023] Double Up 题解 区间dp+二分

题目链接:https://www.luogu.com.cn/problem/P17372

题解区目前给出的都是 \(O(n^3)\) 的区间dp解法。

但是其实这些题解里面的 \(k\) 是可以二分得到的。

因为区间 \([i, j]\) 对应的状态 \(f_{i, j}\),我们枚举 \(k\) 的主要目的是为了找到一个满足 \(f_{i,k} = f_{k+1, j}\) 的下标 \(k\)。

而,随着 \(k\) 的增大:

  • \(f_{i,k}\) 是(非严格)单调递增的;
  • \(f_{k+1, j}\) 是(非严格)单调递减的

所以我们可以二分 \(k\)。


题意是给你一排 \(2\) 的幂,你能随便删数,也能把相邻且相等的两个数合并成它们的两倍。问最后只剩一个数时,最大能是多少。

因为所有数都是 \(2\) 的幂,合并就是乘 \(2\),所以我们只关心能凑出多大的数。一个很自然的想法是 区间DP。

我们定义 f[i][j] 表示:只考虑区间 [i, j] 里的数,通过删数和合并,最后能搞出来的最大数值。注意这个“最后”不是说区间里所有数都必须用上,而是我们可以任意删除中间的元素,只要顺序不变就行。

初始化很简单:如果区间长度是 1,也就是 i == j,那没得选,最大就是 a[i] 本身。

接下来考虑长度大于 1 的区间怎么转移。对于 [i, j],我们有两种操作:

  1. 删数:我们可以把最左边或最右边的数删掉。所以 f[i][j] 至少可以继承 f[i+1][j](删掉 i)和 f[i][j-1](删掉 j)中的较大值。代码里就是 f[i][j] = max(f[i][j-1], f[i+1][j])。

  2. 合并:如果能把 [i, j] 切成左右两半 [i, k] 和 [k+1, j],并且左边能搞出的最大数和右边能搞出的最大数恰好相等,那就可以把这两个数合并成两倍大的数。也就是如果 f[i][k] == f[k+1][j],那就能得到 2 * f[i][k],更新 f[i][j]。

问题来了:怎么找这个分割点 k?总不能每个 k 都试一遍吧,那样太慢。这里有个很妙的性质:f[i][k] 随着 k 增大是单调不降的——因为左区间变长了,能用的数更多,最大结果只会变大不会变小。而 f[k+1][j] 随着 k 增大是单调不增的——右区间变短了,最大结果只会变小不会变大。一个单增,一个单减,那它们最多只有一个交点。如果存在相等,二分就能找到;如果不存在,二分也会很快结束。

所以对于每个区间 [i, j],我们在 i 到 j-1 之间二分找 k:如果 f[i][k] < f[k+1][j],说明左边还不够大,得往右找;如果 f[i][k] > f[k+1][j],说明左边太大了,得往左找;如果相等,那就合并,更新答案,然后 break 就行。

最后输出 f[1][n] 就是答案。

因为输入的数可能大到 \(2^{100}\),超过了 long long,所以得用 __int128 来存,读写也得手写一下。

时间复杂度,区间数量是 \(O(n^2)\),每个区间二分 \(O(log n)\),共 \(O(n^2 \log n)\)


示例程序:

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

__int128 read() {
    __int128 x = 0;
    string s;
    cin >> s;
    for (auto c : s)
        x = x * 10 + c - '0';
    return x;
}

void write(__int128 a) {
    int x = a % 10;
    if (a / 10 > 0)
        write(a / 10);
    cout << x;
}

int n;
__int128 a[1005], f[1005][1005];

int main() {
    cin >> n;
    for (int i = 1; i <= n; i++) a[i] = read();
    for (int l = 1; l <= n; l++) {
        for (int i = 1; i+l-1 <= n; i++) {
            int j = i+l-1;
            if (l == 1) f[i][j] = a[i];
            else {
                f[i][j] = max(f[i][j-1], f[i+1][j]);
                int L = i, R = j-1;
                while (L <= R) {
                    int k = (L + R) >> 1;
                    if (f[i][k] < f[k+1][j])
                        L = k + 1;
                    else if (f[i][k] > f[k+1][j])
                        R = k - 1;
                    else {
                        f[i][j] = max(f[i][j], f[i][k] * 2);
                        break;
                    }
                }
            }
        }
    }
    write(f[1][n]);
    return 0;
}
posted @ 2026-09-13 14:40  quanjun  阅读(6)  评论(0)    收藏  举报