题解:P16315 [ICPC 2023 Jinan R] 基本子串结构

· · 题解

题意简述

定义 Z_i=\operatorname{lcp}(s,\operatorname{suf}(s,i)),则 f(s)=\sum Z_i

对每个位置 t,必须把 s_t 改成另一种字符。求修改后 f(s) 的最大值 m(t),并计算所有 m(t)\oplus t 之和。

解题思路

先求原串的 Z 函数,记:

F=\sum_{i=1}^n Z_i

固定 i\ge 2。原来的 Z_i 表示如下两段逐位相等:

s_{1\dots Z_i}=s_{i\dots i+Z_i-1}

先考虑修改造成的缩短。

t\in[i,i+Z_i-1],修改的是右侧匹配段。最先受到影响的是第 t-i+1 个字符,所以新的 Z 值为 t-i

t 只在左侧匹配段中,新的 Z 值为 t-1。两段相交时,右侧出现的位置更早,故应归入上一种情况。

因此 Z_i 对不同修改位置产生的变化为:

\Delta_i(t)= \begin{cases} t-(Z_i+1), & 1\le t\le\min(Z_i,i-1) \\ t-(i+Z_i), & i\le t\le i+Z_i-1 \end{cases}

两部分都是区间上的一次函数。分别对一次项和常数项使用差分,即可求出每个 t 在不延长任何 Z 值时的答案 B_t

下面考虑延长。令原来的首个失配位置为:

p=Z_i+1,q=i+Z_i

只有 q\le n 时才存在失配。要使 Z_i 增大,必须修改 pq

若修改 p,还要满足 p<i。否则 p 已在右侧匹配段内,修改会先破坏原有匹配。此时令 s_p=s_q,增加量为:

1+\operatorname{lcp}(\operatorname{suf}(s,p+1),\operatorname{suf}(s,q+1))

若修改 q,令 s_q=s_p。这里 q 还会在左侧的后续比较中出现,需要额外分类。

d=i-1,并令:

L=\operatorname{lcp}(\operatorname{suf}(s,p+1),\operatorname{suf}(s,q+1))

位置 q 在左侧后续串中的偏移为 d-1。对应增加量为:

A_i= \begin{cases} L+1, & L<d-1 \\ d+1+\operatorname{lcp}(\operatorname{suf}(s,q+1),\operatorname{suf}(s,q+d+1)), & L=d-1\land q+d\le n\land s_p=s_{q+d} \\ d, & \text{其他情况} \end{cases}

第一种情况在再次遇到修改位置前已经失配。第二种情况恰好又修复了一次失配,可以继续比较。其余情况均在偏移 d-1 处停止。

把每次延长记录成三元组 (t,c,v),表示把 s_t 改成 c 会增加 v。同一位置选择同一字符时,多个 Z 值的增加量可以相加;不同字符不能同时选择。将三元组排序并分组,取每个位置的最大组和 C_t,最终有:

m(t)=B_t+C_t

所有一般后缀 LCP 用后缀数组、Height 数组与 ST 表求出。时间复杂度为 O(n\log n),空间复杂度为 O(n\log n)

正确性证明

若修改位置落在原有匹配段内,它首次参与的相等关系必然被破坏。

在右侧区间中,首次参与位置是第 t-i+1 个字符。在左侧独占区间中,首次参与位置是第 t 个字符。因此,两个区间更新恰好给出所有必然缩短量。

若修改位置不在原有匹配段内,要让 Z_i 超过原值,就必须修复原来的首个失配。因此修改位置只能是 pq,目标字符也分别只能是 s_qs_p。修改 p 时,后续比较不再经过 p,增加量直接由一次 LCP 得到。修改 q 时,后续比较会在偏移 d-1 处再次经过 q。三种分类覆盖了此前失配、在此处修复以及在此处停止,所以 A_i 也是准确的。

缩短量与新字符无关,延长量只与三元组中的目标字符有关。

对相同 (t,c) 的延长量求和。再在所有 c\ne s_t 中取最大值,即枚举了位置 t 的所有合法修改。因此算法求得的 m(t) 均为最优值,最终异或和正确。

参考代码

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

using ll=long long;
const int N=200005;
const int M=400005;
const int K=19;
struct Node
{
    int p,c;
    ll v;
    bool operator<(const Node &x)const
    {
        if(p!=x.p)return p<x.p;
        return c<x.c;
    }
};
int n;
int s[N],z[N],sa[N],rk[N],tmp[N],nrk[N],cnt[N],h[N],lg[N],st[K][N];
ll da[N],db[N],best[N];
Node e[M];
void make_z()
{
    z[1]=n;
    int l=1,r=1;
    for(int i=2;i<=n;i++)
    {
        z[i]=0;
        if(i<=r)z[i]=min(r-i+1,z[i-l+1]);
        while(i+z[i]<=n&&s[z[i]+1]==s[i+z[i]])z[i]++;
        if(i+z[i]-1>r)
        {
            l=i;
            r=i+z[i]-1;
        }
    }
}
void make_sa()
{
    int m=n;
    fill(cnt+1,cnt+m+1,0);
    for(int i=1;i<=n;i++)
    {
        rk[i]=s[i];
        cnt[rk[i]]++;
    }
    for(int i=2;i<=m;i++)cnt[i]+=cnt[i-1];
    for(int i=n;i>=1;i--)sa[cnt[rk[i]]--]=i;
    int p=0;
    for(int k=1;p<n;k<<=1)
    {
        p=0;
        for(int i=max(1,n-k+1);i<=n;i++)tmp[++p]=i;
        for(int i=1;i<=n;i++)if(sa[i]>k)tmp[++p]=sa[i]-k;
        fill(cnt+1,cnt+m+1,0);
        for(int i=1;i<=n;i++)cnt[rk[i]]++;
        for(int i=2;i<=m;i++)cnt[i]+=cnt[i-1];
        for(int i=n;i>=1;i--)sa[cnt[rk[tmp[i]]]--]=tmp[i];
        p=1;
        nrk[sa[1]]=1;
        for(int i=2;i<=n;i++)
        {
            int x=sa[i],y=sa[i-1];
            if(rk[x]!=rk[y]||(x+k<=n?rk[x+k]:0)!=(y+k<=n?rk[y+k]:0))p++;
            nrk[x]=p;
        }
        for(int i=1;i<=n;i++)rk[i]=nrk[i];
        m=p;
    }
    int len=0;
    for(int i=1;i<=n;i++)
    {
        if(rk[i]==1)
        {
            len=0;
            continue;
        }
        if(len)len--;
        int j=sa[rk[i]-1];
        while(i+len<=n&&j+len<=n&&s[i+len]==s[j+len])len++;
        h[rk[i]]=len;
    }
    h[1]=0;
    lg[1]=0;
    for(int i=2;i<=n;i++)lg[i]=lg[i/2]+1;
    for(int i=1;i<=n;i++)st[0][i]=h[i];
    for(int k=1;(1<<k)<=n;k++)
        for(int i=1;i+(1<<k)-1<=n;i++)
            st[k][i]=min(st[k-1][i],st[k-1][i+(1<<(k-1))]);
}
int lcp(int x,int y)
{
    if(x>n||y>n)return 0;
    if(x==y)return n-x+1;
    x=rk[x];
    y=rk[y];
    if(x>y)swap(x,y);
    x++;
    int k=lg[y-x+1];
    return min(st[k][x],st[k][y-(1<<k)+1]);
}
void add(int l,int r,ll a,ll b)
{
    if(l>r)return;
    da[l]+=a;
    da[r+1]-=a;
    db[l]+=b;
    db[r+1]-=b;
}
void solve()
{
    cin>>n;
    for(int i=1;i<=n;i++)cin>>s[i];
    make_z();
    make_sa();
    ll sum=0;
    for(int i=1;i<=n;i++)sum+=z[i];
    for(int i=1;i<=n+1;i++)
    {
        da[i]=0;
        db[i]=0;
        best[i]=0;
    }
    int tot=0;
    for(int i=2;i<=n;i++)
    {
        if(z[i])
        {
            add(1,min(z[i],i-1),1,-z[i]-1);
            add(i,i+z[i]-1,1,-i-z[i]);
        }
        int p=z[i]+1,q=i+z[i];
        if(q>n)continue;
        int d=i-1,l=lcp(p+1,q+1),v;
        if(l<d-1)v=l+1;
        else if(l>d-1)v=d;
        else if(q+d<=n&&s[p]==s[q+d])v=d+1+lcp(q+1,q+d+1);
        else v=d;
        e[++tot]={q,s[p],v};
        if(p<i)e[++tot]={p,s[q],l+1};
    }
    sort(e+1,e+tot+1);
    for(int i=1;i<=tot;)
    {
        int j=i;
        ll v=0;
        while(j<=tot&&e[j].p==e[i].p&&e[j].c==e[i].c)v+=e[j++].v;
        best[e[i].p]=max(best[e[i].p],v);
        i=j;
    }
    ll a=0,b=0,ans=0;
    for(int i=1;i<=n;i++)
    {
        a+=da[i];
        b+=db[i];
        ll v=sum+a*i+b+best[i];
        ans+=(v^i);
    }
    cout<<ans<<'\n';
}
int main()
{
    ios::sync_with_stdio(false);
    cin.tie(nullptr);
    int t;
    cin>>t;
    while(t--)solve();
    return 0;
}