P9338 [JOIST 2023] 合唱 / Chorus

· · 题解

比较牛的题,场切了(如果不算没开 __int128 的话,不是我大样例过穿了都没发现这个问题)。

思路:

显然可能划分成若干个子序列满足限制的充要条件是:

必要性:

充分性:

容易发现必然存在一个最优解使得划分一个子序列时一定是 A[l, r]B[l, r] 匹配,不然如果不是 A 中区间或者 B 中不是 [l, r] 的话都可以调整法证明不比这个优秀。

于是设 w(l, r) 表示将 A[l, r]B[l, r] 匹配的最小交换次数,考虑每个 A 往前走要换多少个 B,具体的,设 h_i 表示第 iA 前面有多少个 A,那么有:

w(l, r) = \sum_{i = l}^r \max(0, h_i - (l - 1))

于是容易想到状态 f_{i, j} 表示考虑匹配前 iA, B 划分成了 j 个子序列的最小交换次数,那么有转移:

f_{i, j} = \min_{k < i} f_{k, j - 1} + w(k + 1, i)

答案是 f_{n, k},时间复杂度是 O(n^4) 的。

容易发现 f_{n, i} 是单调不增的,证明也是显然,选 i 个的限制是比选 i + 1 是严的;进一步的,w(l, r) 是满足四边形不等式的,则 f_{n, i} 是下凸的(可以参考邮局里面相关证明)。

虽然场上我是直接猜的凸性,不然肯定做不了,于是直接 wqs 二分斜率 mid,那么问题转化为划分为若干个子序列,每个子序列带一个 +mid 的贡献,求划分代价最小的情况下划分数量的最小值。

现在状态是一维的 f_i,有转移:

f_i = \min_{k < i} f_k + w(k + 1, i)

由于 w(l, r) 满足四边形,所以有决策单调性,可以二分队列解决,这样是 O(n \log^2 n) 的,可能卡常后有概率通过。

注意到 w(l, r) 可以写成一个式子,令 nxt_i 表示 i 后面第一个 h_j \ge ii,那么有:

\begin{aligned} w(l, r) &= \sum_{i = nxt_{l - 1}}^r (h_i - (l - 1)) \\ &= pre_r - pre_{nxt_{l - 1} - 1} - (r - nxt_{l - 1} + 1)(l - 1)\end{aligned}

其中 preh 的前缀和,显然当 r < nxt_{l - 1} 的时候这个式子值是有问题的,应该是 0,处理方式是类似扫描线,当 i 扫到某个 nxt_{l - 1} 的时候把对应的 l - 1 加入进来即可,但是实际实现时根本不管也可以通过,不太会证明;但是就算加入了这点代码实现也不会改很多,因为 nxt 的单调性,所以插点还是单调的。

考虑斜率优化,那么式子就变成了:

f_i = mid + f_j + pre_i - pre_{nxt_j - 1} - (i - nxt_j + 1) \times j f_i = mid + pre_i + f_j - pre_{nxt_j - 1} + (nxt_j - 1) \times j - i \times j i \times j + (f_i - mid - pre_i) = \bigl(f_j - pre_{nxt_j - 1} + (nxt_j - 1) \times j\bigr)

于是看作点 (j, f_j - pre_{nxt_j - 1} + (nxt_j - 1) \times j),查询前面点中切上一个斜率为 i 的直线后截距的最小值,于是维护前面点的下凸包即可;且查询斜率 i 单增,直接走指针即可。

时间复杂度为 O(n \log n)

完整代码:

#include<bits/stdc++.h>
#define fi first

#define se second

using namespace std;
typedef long long ll;
#define __ __int128

const int N = 2e6 + 10;
inline ll read(){
    ll x = 0, f = 1;
    char c = getchar();
    while(c < '0' || c > '9'){
        if(c == '-')
          f = -1;
        c = getchar();
    }
    while(c >= '0' && c <= '9'){
        x = (x << 1) + (x << 3) + (c ^ 48);
        c = getchar();
    }
    return x * f;
}
inline void write(ll x){
    if(x < 0){
        putchar('-');
        x = -x;
    }
    if(x > 9)
      write(x / 10);
    putchar(x % 10 + '0');
}
int n, k, ca, cb, head, top;
int posa[N], posb[N], h[N], sb[N], nxt[N], stk[N];
//ll w[N >> 1][N >> 1];
//ll dp[N][N];
__ f[N], pre[N];
__ dx[N], dy[N];
int g[N];
char s[N];
inline ll getw(int l, int r){
    if(nxt[l - 1] > r){
//      cerr << pre[r] - pre[nxt[l - 1] - 1] - 1ll * (r - nxt[l - 1] + 1) * (l - 1) << '\n';
        return 0;
    }
    return pre[r] - pre[nxt[l - 1] - 1] - 1ll * (r - nxt[l - 1] + 1) * (l - 1);
}
// f[i] = mid + f[j] + getw(j + 1, i)
// f[i] = mid + f[j] + pre[i] - pre[nxt[j] - 1] - (i - nxt[j] + 1) * j
// f[i] = mid + pre[i] + f[j] - pre[nxt[j] - 1] + (nxt[j] - 1) * j - i * j
// i * j + (f[i] - mid - pre[i]) = (f[j] - pre[nxt[j] - 1] + (nxt[j] - 1) * j)
// j -> (j, f[j] - pre[nxt[j] - 1] + (nxt[j] - 1) * j)
// 维护下凸包 
// 查询斜率为 i 的直线在这个凸包上的切点即可
// i 是单增的 
inline int get(ll mid, bool flag = 0){
    dx[0] = 0, dy[0] = f[0] - pre[nxt[0] - 1];
    head = 1, top = 0;
    stk[++top] = 0;
    for(int i = 1; i <= n; ++i){
        while(head < top && (dy[stk[head + 1]] - dy[stk[head]]) < (__)1ll * i * (dx[stk[head + 1]] - dx[stk[head]]))
          ++head;
        int j = stk[head];
        f[i] = mid + f[j] + pre[i] - pre[nxt[j] - 1] - 1ll * (i - nxt[j] + 1) * j;
        g[i] = g[j] + 1;
        dx[i] = i, dy[i] = f[i] - pre[nxt[i] - 1] + 1ll * (nxt[i] - 1) * i;
        while(head < top && (((__)(dy[i] - dy[stk[top]]) * (dx[stk[top]] - dx[stk[top - 1]]) < (__)(dy[stk[top]] - dy[stk[top - 1]]) * (dx[i] - dx[stk[top]])) || ((__)(dy[i] - dy[stk[top]]) * (dx[stk[top]] - dx[stk[top - 1]]) == (__)(dy[stk[top]] - dy[stk[top - 1]]) * (dx[i] - dx[stk[top]]) && g[i] <= g[stk[top]])))
          --top;
        stk[++top] = i;
//      if(flag)
//        cerr << dx[i] << ' ' << dy[i] << ' ' << id << ' ' << i * dx[id] + (f[i] - mid - pre[i]) << ' ' << dy[id] << '\n';
    }
    return g[n];
}
int main(){
    // freopen("purple.in", "r", stdin);
    // freopen("purple.out", "w", stdout);
    n = read(), k = read();
    scanf("%s", s + 1);
    for(int i = 1; i <= (n << 1); ++i){
        if(s[i] == 'A')
          posa[++ca] = i;
        else
          posb[++cb] = i;
        sb[i] = sb[i - 1] + (s[i] == 'B');
    }
    for(int i = 1; i <= n; ++i)
      h[i] = sb[posa[i]], pre[i] = pre[i - 1] + h[i];
    nxt[0] = 1;
    for(int i = 1; i <= n; ++i){
        int l = i + 1, r = n;
        while(l <= r){
            int mid = (l + r) >> 1;
            if(h[mid] >= i)
              nxt[i] = mid, r = mid - 1;
            else
              l = mid + 1;
        }
        if(!nxt[i])
          nxt[i] = n + 1;
    }
//  for(int l = 1; l <= n; ++l)
//    for(int r = l; r <= n; ++r)
//      w[l][r] = max(0, h[r] - (l - 1)) + w[l][r - 1];
//  for(int l = 1; l < n; ++l){
//      for(int r = l + 1; r < n; ++r){
//          int a = w[l][r] + w[l + 1][r - 1], b = w[l][r + 1] + w[l + 1][r];
////            if(a > b)
////              cerr << l << ' ' << r << ' ' << a << ' ' << b << '\n';
//          assert(a <= b);
//      }
//  }
//  for(int l = 1; l <= n; ++l){
//      for(int r = l; r <= n; ++r){
//          cerr << l << ' ' << r << ' ' << w(l, r) << '\n';
//      }
//  }
    /*
    dp[0][0] = 0;
    for(int i = 1; i <= n; ++i)
      dp[0][i] = 1e18;
    for(int i = 1; i <= n; ++i){
        dp[i][1] = w(1, i);
        for(int j = 2; j <= i; ++j){
            dp[i][j] = 1e18;    
            for(int k = 1; k < i; ++k)
              dp[i][j] = min(dp[i][j], dp[k][j - 1] + w(k + 1, i));
//          cerr << "dp: " << i << ' ' << j << ' ' << dp[i][j] << '\n';
        }
    }
//  for(int i = 1; i <= n; ++i)
//    cerr << dp[n][i] << ' ';
//  cerr << '\n';
    write(dp[n][k]);*/
    ll l = 0, r = 1e12, ans = 0;
    while(l <= r){
        ll mid = (l + r) >> 1;
        if(get(mid) <= k){
            ans = mid;
            r = mid - 1;
        }
        else
          l = mid + 1;
    }
//  cerr << ans << ' ' << get(ans) << '\n';
    get(ans, 1);
    write(f[n] - 1ll * k * ans);
    return 0;
}