待办
学习后缀自动机;
完成相关例题。
题单
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;
}


Comments NOTHING