1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95
| #include<bits/stdc++.h> using namespace std; #define int long long inline int read(){int x=0,f=1;char ch=getchar(); while(ch<'0'|| ch>'9'){if(ch=='-')f=-1;ch=getchar();}while(ch>='0'&&ch<='9'){ x=(x<<1)+(x<<3)+(ch^48);ch=getchar();}return x*f;} #define Z(x) (x)*(x) #define pb push_back
#define N 2000010
void Mn(int &a, int b) { a = min(a, b); } struct F { int a, t; F operator + (const F &A) const { F B; B.a = a + A.a; B.t = t + A.t; return B; } F operator + (const int &x) const { F B; B.a=a+x; B.t=t; return B; } bool operator < (const F &A) const { if(a == A.a) return t < A.t; return a < A.a; } }f[N]; int n, m, i, j, k, T; int a[N], b[N]; int sa[N], sb[N], sum[N]; int ida[N], idb[N], h[N]; int X[N], Y[N]; int l, r, pos, ans, ja, jb; int X1, X2, X3, Y1, Y2, Y3; char s[N]; int q[N];
int calc(int x) { memset(f, 0x3f, sizeof(f)); auto suan = [&] (int j) -> int { return Y[j] - X[j] * i; }; int l, r; l = r = 0; f[0]={0, 0}; q[0]=1; X[0]=1; Y[0]=h[1]; for(i=1; i<=n; ++i) { while(r-l+1 >= 2 && suan(l) > suan(l+1)) ++l; f[i].a = suan(l); f[i].t = f[q[l]-1].t; f[i] = f[i] + sum[i] + i + x; f[i].t++; while(r-l+1 >= 2) { X1 = X[r-1]; Y1 = Y[r-1]; X2 = X[r]; Y2 = Y[r]; X3 = (i+1); Y3 = f[i].a + h[i+1]; if((Y2-Y1)*(X3-X2) > (Y3-Y2)*(X2-X1)) --r; else break; } q[++r] = i+1; X[r] = (i+1); Y[r] = f[i].a + h[i+1]; } return f[n].t; }
signed main() { freopen("easy.in", "r", stdin); freopen("easy.out", "w", stdout); n=read(); m=read(); ans=1e18; scanf("%s", s+1); for(i=1; i<=2*n; ++i) if(s[i]=='A') a[++ja]=i, sa[i]=1; else b[++jb]=i, sb[i]=1; partial_sum(sa+1, sa+2*n+1, sa+1); partial_sum(sb+1, sb+2*n+1, sb+1); for(i=1; i<=n; ++i) ida[i] = sb[a[i]]; for(i=1; i<=n; ++i) idb[i] = sa[b[i]]; for(i=1; i<=n; ++i) idb[i] = max(idb[i], i - 1); partial_sum(ida+1, ida+n+1, sum+1); for(i=1; i<=n; ++i) h[i] = idb[i]*(i-1)-sum[idb[i]]; l=0; r=1e12; while(l<r) { int mid = (l+r)>>1; if(calc(mid)<=m) r=mid; else l=mid+1; } calc(l); printf("%lld", f[n].a - m * l); return 0; }
|