ZR 集训 Day22 – 后缀数组,后缀自动机

ooliver 发布于 15 小时前 104 次阅读 OI


AI 摘要

从哈希二分的暴力美学到倍增排序的优雅加速,后缀数组 + height 数组竟能一行公式数清不同子串?单调栈一出,差异题也俯首称臣。 SAM 待续,先看这些后缀利器如何化繁为简。

待办

学习后缀自动机;
完成相关例题。

题单

P3809 【模板】后缀排序

先说一个 $O(nlog^2n)$ 的做法,那就是直接排序,重写 cmp,用哈希+二分找到两个串的最长公共前缀,比较下一位即可。

代码:

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

using ull=unsigned long long;
const int N=1e6+5;
const int p=1e5+7;
int n;
int a[N];
ull ha[N],mi[N];
char s[N];

void init(){
    mi[0]=1;
    for(int i=1;i<=n;i++) mi[i]=mi[i-1]*p,a[i]=i,ha[i]=ha[i-1]*p+s[i];
}

ull getha(int l,int r){
    if(r<l) return 0;
    return ha[r]-ha[l-1]*mi[r-l+1];
}

bool cmp(int x,int y){
    int l=0,r=min(n-y+1,n-x+1);
    while(l<r){
        int mid=(l+r+1)>>1;
        if(getha(x,x+mid-1)==getha(y,y+mid-1)) l=mid;
        else r=mid-1;
    }
    if(x+l-1==n) return 1;
    if(y+l-1==n) return 0;
    return s[x+l]<s[y+l];
}

int main(){
    scanf("%s",s+1);
    n=strlen(s+1);
    init();
    stable_sort(a+1,a+1+n,cmp);
    for(int i=1;i<=n;i++) printf("%d ",a[i]);
    return 0;
}

接下来是正解(附 $height$ 数组构建),即 $O(n\log n)$:

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

const int N=1e6+5;
int n,m,p;
int rk[N<<1],ork[N],sa[N<<1],id[N],cnt[N],ht[N];
char s[N];

int main(){
    scanf("%s",s+1);
    n=strlen(s+1);
    m=128;
    for(int i=1;i<=n;i++) cnt[rk[i]=s[i]]++;
    for(int i=1;i<=m;i++) cnt[i]+=cnt[i-1];
    for(int i=n;i>=1;i--) sa[cnt[rk[i]]--]=i;
    for(int w=1;;w<<=1,m=p){
        int cur=0;
        for(int i=n-w+1;i<=n;i++) id[++cur]=i;
        for(int i=1;i<=n;i++) if(sa[i]>w) id[++cur]=sa[i]-w;
        memset(cnt,0,sizeof(cnt));
        for(int i=1;i<=n;i++) cnt[rk[i]]++;
        for(int i=1;i<=m;i++) cnt[i]+=cnt[i-1];
        for(int i=n;i>=1;i--) sa[cnt[rk[id[i]]]--]=id[i];
        p=0;
        memcpy(ork,rk,sizeof(ork));
        for(int i=1;i<=n;i++){
            if(ork[sa[i]]==ork[sa[i-1]]&&ork[sa[i]+w]==ork[sa[i-1]+w]) rk[sa[i]]=p;
            else rk[sa[i]]=++p;
        }
        if(p==n) break;
    }
    for(int i=1,k=0;i<=n;i++){
        if(rk[i]==1) continue;
        if(k) k--;
        while(s[i+k]==s[sa[rk[i]-1]+k]) k++;
        ht[rk[i]]=k;
    }
    for(int i=1;i<=n;i++) printf("%d ",sa[i]);
    return 0;
}

P4248 [AHOI2013] 差异

我们知道:

$$
lcp(sa_i,sa_j)=min\{height_i+1,...,height_j+1\}
$$

之后此题就变得异常简单,对于前半部分直接数学拆一下,后半部分是子区间最小值之和,用两个单调栈求出每一个点作为区间最值时所能影响的区间左右端点即可。

代码:

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

#define int long long
const int N=1e6+5;
int n,m,p,ans;
int rk[N<<1],ork[N],sa[N<<1],id[N],cnt[N],ht[N],l[N],r[N];
char s[N];

signed main(){
    scanf("%s",s+1);
    n=strlen(s+1),m=128;
    for(int i=1;i<=n;i++) cnt[rk[i]=s[i]]++;
    for(int i=1;i<=m;i++) cnt[i]+=cnt[i-1];
    for(int i=n;i>=1;i--) sa[cnt[rk[i]]--]=i;
    for(int w=1;;w<<=1,m=p){
        int cur=0;
        for(int i=n-w+1;i<=n;i++) id[++cur]=i;
        for(int i=1;i<=n;i++) if(sa[i]>w) id[++cur]=sa[i]-w;
        memset(cnt,0,sizeof(cnt));
        for(int i=1;i<=n;i++) cnt[rk[i]]++;
        for(int i=1;i<=m;i++) cnt[i]+=cnt[i-1];
        for(int i=n;i>=1;i--) sa[cnt[rk[id[i]]]--]=id[i];
        p=0;
        memcpy(ork,rk,sizeof(ork));
        for(int i=1;i<=n;i++){
            if(ork[sa[i]]==ork[sa[i-1]]&&ork[sa[i]+w]==ork[sa[i-1]+w]) rk[sa[i]]=p;
            else rk[sa[i]]=++p;
        }
        if(p==n) break;
    }
    for(int i=1,k=0;i<=n;i++){
        if(rk[i]==1) continue;
        if(k) k--;
        while(s[i+k]==s[sa[rk[i]-1]+k]) k++;
        ht[rk[i]]=k;
    }
    stack<int> st;
    for(int i=1;i<=n;i++){
        ans+=(n-i)*(n+3*i+1)/2;
        if(i==1) continue;
        while(st.size()&&ht[st.top()]>=ht[i]) st.pop();
        if(!st.size()) l[i]=1;
        else l[i]=st.top();
        st.push(i);
    }
    while(st.size()) st.pop();
    for(int i=n;i>=2;i--){
        while(st.size()&&ht[st.top()]>ht[i]) st.pop();
        if(!st.size()) r[i]=n+1;
        else r[i]=st.top();
        st.push(i);
    }
    for(int i=2;i<=n;i++) ans-=2*ht[i]*(i-l[i])*(r[i]-i);
    cout<<ans;
    return 0;
}

P2408 不同子串个数

根据 OI-wiki,我们有:

$$
ans=\frac{n(n+1)}{2} - \sum_{i=2}^n height_i
$$

代码:

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

#define int long long
const int N=1e6+5;
int n,m,p,ans;
int rk[N<<1],ork[N],sa[N<<1],id[N],cnt[N],ht[N];
char s[N];

signed main(){
    cin>>n;
    scanf("%s",s+1);
    m=128;
    for(int i=1;i<=n;i++) cnt[rk[i]=s[i]]++;
    for(int i=1;i<=m;i++) cnt[i]+=cnt[i-1];
    for(int i=n;i>=1;i--) sa[cnt[rk[i]]--]=i;
    for(int w=1;;w<<=1,m=p){
        int cur=0;
        for(int i=n-w+1;i<=n;i++) id[++cur]=i;
        for(int i=1;i<=n;i++) if(sa[i]>w) id[++cur]=sa[i]-w;
        memset(cnt,0,sizeof(cnt));
        for(int i=1;i<=n;i++) cnt[rk[i]]++;
        for(int i=1;i<=m;i++) cnt[i]+=cnt[i-1];
        for(int i=n;i>=1;i--) sa[cnt[rk[id[i]]]--]=id[i];
        p=0;
        memcpy(ork,rk,sizeof(ork));
        for(int i=1;i<=n;i++){
            if(ork[sa[i]]==ork[sa[i-1]]&&ork[sa[i]+w]==ork[sa[i-1]+w]) rk[sa[i]]=p;
            else rk[sa[i]]=++p;
        }
        if(p==n) break;
    }
    for(int i=1,k=0;i<=n;i++){
        if(rk[i]==1) continue;
        if(k) k--;
        while(s[i+k]==s[sa[rk[i]-1]+k]) k++;
        ht[rk[i]]=k;
    }
    int ans=n*(n+1)/2;
    for(int i=2;i<=n;i++) ans-=ht[i];
    cout<<ans;
    return 0;
}