peiwenjun's blog 没有知识的荒原

CF356E Xenia and String Problem 题解

题目描述

定义一个字符串 \(s\) 是好的,当且仅当:

  • \(n=|s|\) 为奇数。
  • \(s_\frac{n+1}2\)\(s\) 中仅出现一次。
  • \(n=1\)\(s_{1,\frac{n-1}2}=s_{\frac{n+3}2,n}\) 是好的。

定义一个字符串的权值为长度的平方。

定义一个字符串的美丽值为其所有好的子串的权值和。

给定长为 \(n\) 的字符串 \(s\) ,你可以修改 \(s\) 的至多一个字符,求修改后 \(s\) 的美丽值的最大值。

数据范围

  • \(1\le n\le 10^5\)

时间限制 \(\texttt{1s}\) ,空间限制 \(\texttt{256MB}\)

分析

注意到好的字符串长度为 \(2^j-1\) ,所以好的字符串至多只有 \(\mathcal O(n\log n)\) 个。

\(f_{j,i}\) 表示 \(s_{i,i+2^j-2}\) 是否为好的,使用哈希或后缀数组判断两边是否相同,即可在 \(\mathcal O(n\log n)\) 的时间内求出哪些串是好的。

再来考虑修改,由于修改后的字符串只有 \(\mathcal O(26n)\) 种,我们希望对每个串在 \(\mathcal O(\log n)\) 的时间内计算增量。

修改可以看成删除再加入。

删除会影响到所有包含它的好串,差分前缀和即可求出对每个位置,包含它的好串权值和。

加入的贡献需要换个方向计算。对每个可能成为好区间的位置 \([l,r=l+2^j-2]\) ,计算哪些修改可以让它变成好串。

  • \(f_{j-1,l}+f_{j-1,mid+1}=0\) :显然无论如何修改都不可能成为好串。
  • \(f_{j-1,l}+f_{j-1,mid+1}=1\) :左右两边必须恰好有一个字符不同,才有可能成为好串。
  • \(f_{j-1,l}+f_{j-1,mid+1}=2\) :如果两边相同,只需要枚举 \(s_{mid}\) 被修改成什么并判断第二个条件;如果两边不同,只能是两边中间的字符不同。

整个过程中会大量调用判断两个串是否相等或相差一个字符的函数,用后缀数组求 \(\texttt{lcp}\) 可以做到 \(\mathcal O(1)\) 回答。

时间复杂度 \(\mathcal O(26n\log n)\) ,但是显然跑不满。

#include<bits/stdc++.h>
#define ll long long
using namespace std;
const int maxn=1e5+5;
int n;
int c[maxn],x[maxn],y[2*maxn];
int h[maxn],rk[maxn],sa[maxn];
char s[maxn];
int f[17][maxn],g[17][maxn];
int lg[maxn],pre[maxn][26],nxt[maxn][26];
ll cur,res,sum[maxn],val[maxn][26];
void get_sa()
{
    int m=122;
    for(int i=1;i<=n;i++) x[i]=s[i],c[x[i]]++;
    for(int i=1;i<=m;i++) c[i]+=c[i-1];
    for(int i=n;i>=1;i--) sa[c[x[i]]--]=i;
    for(int k=1;k<=n;k<<=1)
    {
        int num=0;
        for(int i=n-k+1;i<=n;i++) y[++num]=i;
        for(int i=1;i<=n;i++) if(sa[i]>k) y[++num]=sa[i]-k;
        for(int i=1;i<=m;i++) c[i]=0;
        for(int i=1;i<=n;i++) c[x[i]]++;
        for(int i=1;i<=m;i++) c[i]+=c[i-1];
        for(int i=n;i>=1;i--) sa[c[x[y[i]]]--]=y[i];
        for(int i=1;i<=n;i++) y[i]=x[i],x[i]=0;
        x[sa[1]]=num=1;
        for(int i=2;i<=n;i++)
            x[sa[i]]=y[sa[i]]==y[sa[i-1]]&&y[sa[i]+k]==y[sa[i-1]+k]?num:++num;
        if(num==n) break;
        m=num;
    }
}
void get_height()
{
    for(int i=1;i<=n;i++) rk[sa[i]]=i;
    for(int i=1,k=0;i<=n;i++)
    {
        if(k) k--;
        int j=sa[rk[i]-1];
        while(i+k<=n&&j+k<=n&&s[i+k]==s[j+k]) k++;
        h[rk[i]]=k;
    }
}
int lcp(int i,int j)
{
    if(rk[i]>rk[j]) swap(i,j);
    int l=rk[i]+1,r=rk[j],k=lg[r-l+1];
    return min(g[k][l],g[k][r-(1<<k)+1]);
}
int check(int i,int j,int l)
{
    int x=lcp(i,j);
    if(x>=l) return 0;
    if(x+1+lcp(i+x+1,j+x+1)>=l) return x+1;
    return -1;
}
int main()
{
    scanf("%s",s+1),n=strlen(s+1);
    get_sa(),get_height();
    for(int i=2;i<=n;i++) lg[i]=lg[i>>1]+1;
    for(int i=1;i<=n;i++) g[0][i]=h[i];
    for(int j=1;j<=16;j++)
        for(int i=1;i+(1<<j)-1<=n;i++)
            g[j][i]=min(g[j-1][i],g[j-1][i+(1<<(j-1))]);
    for(int i=1;i<=n;i++)
    {
        memcpy(pre[i],pre[i-1],sizeof(pre[i]));
        pre[i][s[i]-'a']=i;
    }
    for(int c=0;c<26;c++) nxt[n+1][c]=n+1;
    for(int i=n;i>=1;i--)
    {
        memcpy(nxt[i],nxt[i+1],sizeof(nxt[i]));
        nxt[i][s[i]-'a']=i;
    }
    for(int j=1;j<=16;j++)
        for(int l=1,r=(1<<j)-1,mid=(l+r)>>1;r<=n;l++,mid++,r++)
        {
            int x=f[j-1][l],y=f[j-1][mid+1],tmp=check(l,mid+1,(1<<(j-1))-1);
            int c=s[mid]-'a',L=pre[mid-1][c]<l,R=nxt[mid+1][c]>r;
            ll w=1ll*(r-l+1)*(r-l+1);
            f[j][l]=j!=1?x&&y&&!tmp&&L&&R:1;
            if(f[j][l]) cur+=w,sum[l]+=w,sum[r+1]-=w;
            if(x+y==1&&tmp>0)
            {
                if(x&&L) val[mid+tmp][s[l+tmp-1]-'a']+=w;
                if(y&&R) val[l+tmp-1][s[mid+tmp]-'a']+=w;
            }
            if(j==1||(x+y==2&&!tmp))
                for(int c=0;c<26;c++)
                    if(pre[mid-1][c]<l&&nxt[mid+1][c]>r)
                        val[mid][c]+=w;
            if(x+y==2&&tmp==(1<<(j-2)))
            {
                if(L) val[mid+tmp][s[l+tmp-1]-'a']+=w;
                if(R) val[l+tmp-1][s[mid+tmp]-'a']+=w;
            }
        }
    for(int i=1;i<=n;i++) sum[i]+=sum[i-1];
    for(int i=1;i<=n;i++) for(int c=0;c<26;c++) res=max(res,cur-sum[i]+val[i][c]);
    printf("%lld\n",res);
    return 0;
}

posted on 2023-06-12 10:46  peiwenjun  阅读(8)  评论(0)    收藏  举报

导航