P9338 [JOIST 2023] 合唱 / Chorus
Genius_Star · · 题解
比较牛的题,场切了(如果不算没开 __int128 的话,不是我大样例过穿了都没发现这个问题)。
思路:
显然可能划分成若干个子序列满足限制的充要条件是:
- 第
i 个A 的位置小于第i 个B 的位置。
必要性:
- 考虑第
i 个B ,如果前面A 的个数<i ,那么这个B 就没有和它匹配的,就寄了.
充分性:
- 显然可以构造成划分为
n 个子序列的方案。
容易发现必然存在一个最优解使得划分一个子序列时一定是
于是设
于是容易想到状态
答案是
容易发现
虽然场上我是直接猜的凸性,不然肯定做不了,于是直接 wqs 二分斜率
现在状态是一维的
由于
注意到
其中
考虑斜率优化,那么式子就变成了:
于是看作点
时间复杂度为
完整代码:
#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;
}