题解 P6292 区间本质不同子串个数
这是一个使用后缀数组来替代后缀自动机的做法。
以前我也曾认为区间本质不同子串个数这类问题,SA 是完全无法解决的;直到某次做题时偶然受到 SAM 的 parent tree 启发,才发现了这种非常强的做法。
考虑这样一个问题:我们定义
你一定可以发现:这就是 SAM 中反串的
我们知道每一个等价类在 sa 数组上都会对应一个区间,直观的想法是枚举每个左端点,再计算
观察这棵笛卡尔树,我们可以发现其具备很多优美的性质:每个等价类都在其上对应了一棵子树,也对应了 sa 数组的一段区间;其先序遍历的结果就是后缀数组;任意两个叶子节点的 lca 就是它们的 lcp …… 可以发现,它几乎就是一棵后缀树!
现在,我们终于可以借助这棵笛卡尔树,把 SA 与线段树合并或是 LCT 等数据结构结合起来,这棵树完全可以发挥与 SAM 的 parent tree 相同的效力,而且代码实现也并不复杂。另外值得一提的是,我们甚至可以在这棵树上实现在线的匹配。
这样一来,剩下的过程就与其它的题解无异了。模仿区间数颜色的做法,我们首先将离线询问,再从右往左扫描线。每插入一个新的后缀,其在笛卡尔树上涉及的仅是一条由叶子到根的链,且操作过程形似 LCT 的 access 操作,所以可以使用 LCT 优化。我们只需再维护区间加法和区间求和即可,我的实现是
#include<bits/stdc++.h>
#define rep(i,a,b) for(int i=a;i<=b;i++)
#define endl putchar('\n')
const int N=200005;
#define int long long
using namespace std;
int n,Q,sa[N],rk[N],h[N],root,val[N],pre[N],ans[N];
char s[N];
struct ques {
int l,r,id;
bool operator < (const ques &t) const { return l<t.l; }
} q[N];
struct suffix_array {
int k1[N],k2[N],cnt[N],mx,num;
void radix_sort() {
rep(i,1,mx) cnt[i]=0;
rep(i,1,n) cnt[k1[i]]++;
rep(i,1,mx) cnt[i]+=cnt[i-1];
for(int i=n;i>=1;i--) sa[cnt[k1[k2[i]]]--]=k2[i];
}
void sort() {
rep(i,1,n) k1[i]=s[i],k2[i]=i; mx='z';
radix_sort();
for(int j=1;j<=n;j<<=1,num=0) {
rep(i,n-j+1,n) k2[++num]=i;
rep(i,1,n) if(sa[i]-j>=1) k2[++num]=sa[i]-j;
radix_sort(),swap(k2,k1);
num=k1[sa[1]]=1;
rep(i,2,n) k1[sa[i]]=k2[sa[i]]==k2[sa[i-1]]&&k2[sa[i]+j]==k2[sa[i-1]+j]?num:++num;
if(num==n) break;
mx=num;
}
}
void height() {
rep(i,1,n) rk[sa[i]]=i;
int k=0;
rep(i,1,n) {
if(rk[i]==1) continue;
if(k) k--;
int j=sa[rk[i]-1];
while(max(i,j)+k<=n&&s[i+k]==s[j+k]) k++;
h[rk[i]]=k;
}
}
} SA;
struct SMT {
int sm[N<<1],ad[N<<1];
#define ls (k<<1)
#define rs (k<<1|1)
#define mid ((l+r)>>1)
void add(int k,int l,int r,int v) { sm[k]+=(r-l+1)*v,ad[k]+=v; }
void pushdown(int k,int l,int r) {
if(ad[k]) add(ls,l,mid,ad[k]),add(rs,mid+1,r,ad[k]),ad[k]=0;
}
void pushup(int k) { sm[k]=sm[ls]+sm[rs]; }
void add(int k,int l,int r,int x,int y,int v) {
if(x<=l&&r<=y) return add(k,l,r,v);
pushdown(k,l,r);
if(x<=mid) add(ls,l,mid,x,y,v);
if(y>mid) add(rs,mid+1,r,x,y,v);
pushup(k);
}
int query(int k,int l,int r,int x,int y) {
if(x<=l&&r<=y) return sm[k];
pushdown(k,l,r);
int res=0;
if(x<=mid) res+=query(ls,l,mid,x,y);
if(y>mid) res+=query(rs,mid+1,r,x,y);
return res;
}
#undef ls
#undef rs
#undef mid
} smt;
struct LCT {
struct node {
int sm,top,bel,fa,ch[2];
#define ls(x) nod[x].ch[0]
#define rs(x) nod[x].ch[1]
#define fa(x) nod[x].fa
#define sm(x) nod[x].sm
#define top(x) nod[x].top
#define bel(x) nod[x].bel
} nod[N];
bool cmp(int x) { return x==rs(fa(x)); }
bool isroot(int x) { return nod[fa(x)].ch[cmp(x)]!=x; }
void pushdown(int x) { bel(ls(x))=bel(rs(x))=bel(x); }
void pushup(int x) {
sm(x)=sm(ls(x))+val[x]+sm(rs(x));
top(x)=ls(x)?top(ls(x)):pre[x];
}
void connect(int x,int fa,int son) { nod[fa].ch[son]=x,fa(x)=fa; }
void rotate(int x) {
int y=fa(x),z=fa(y),ys=cmp(x),zs=cmp(y);
if(isroot(y)) fa(x)=z; else connect(x,z,zs);
connect(nod[x].ch[ys^1],y,ys),connect(y,x,ys^1),pushup(y),pushup(x);
}
void pushall(int x) { if(!isroot(x)) pushall(fa(x)); pushdown(x); }
void splay(int x) {
pushall(x);
while(!isroot(x)) {
if(!isroot(fa(x))) rotate(cmp(x)^cmp(fa(x))?x:fa(x));
rotate(x);
}
}
void access(int x) {
int pos=x; x=rk[x]+n;
for(int y=0;x;y=x,x=fa(x)) {
splay(x),rs(x)=0,pushup(x);
int l=top(x),r=top(x)+sm(x)-1;
if(l<=r&&bel(x)) smt.add(1,1,n,bel(x)+l,bel(x)+r,-1);
bel(x)=pos;
if(l<=r) smt.add(1,1,n,bel(x)+l,bel(x)+r,1);
rs(x)=y,pushup(x);
}
}
} lct;
struct treap {
int ls[N],rs[N],s[N],top,fa[N];
void build() {
if(n==1) return root=2,void();
ls[2]=n+1,rs[2]=n+2,fa[ls[2]]=fa[rs[2]]=root=s[++top]=2;
rep(i,3,n) {
while(top&&h[s[top]]>h[i]) top--;
if(top) ls[i]=rs[s[top]],fa[rs[s[top]]]=i,rs[s[top]]=i,fa[i]=s[top];
else ls[i]=root,fa[root]=i,root=i;
rs[i]=n+i,fa[rs[i]]=i,s[++top]=i;
}
rep(i,2,n+n) val[i]=(i>n?n-sa[i-n]+1:h[i])-h[fa[i]];
}
void prep(int x) {
lct.sm(x)=val[x],lct.top(x)=pre[x];
if(x>n) return;
pre[ls[x]]=pre[rs[x]]=pre[x]+val[x];
lct.fa(ls[x])=lct.fa(rs[x])=x;
prep(ls[x]),prep(rs[x]);
}
} trp;
void init() {
SA.sort(),SA.height();
trp.build(),trp.prep(root);
}
void solve() {
sort(q+1,q+Q+1);
int j=Q;
for(int i=n;i>=1;i--) {
lct.access(i);
while(j>=1&&q[j].l==i) {
ans[q[j].id]=smt.query(1,1,n,q[j].l,q[j].r);
j--;
}
}
}
signed main() {
scanf("%s",s+1),n=strlen(s+1);
init();
scanf("%lld",&Q);
rep(i,1,Q) scanf("%lld%lld",&q[i].l,&q[i].r),q[i].id=i;
solve();
rep(i,1,Q) printf("%lld\n",ans[i]);
return 0;
}