题解 P6292 区间本质不同子串个数

· · 题解

这是一个使用后缀数组来替代后缀自动机的做法。

以前我也曾认为区间本质不同子串个数这类问题,SA 是完全无法解决的;直到某次做题时偶然受到 SAM 的 parent tree 启发,才发现了这种非常强的做法。

考虑这样一个问题:我们定义 {\rm beginpos}(t) 表示字符串 s 的子串 ts 中的所有起始位置所构成的集合。显然,所有这样的 t 都可以根据 {\rm beginpos}(t) 而被划分为若干个等价类;我们的目标则是对于每一个这样的等价类,求出其所包含的本质不同的 t 的数量。

你一定可以发现:这就是 SAM 中反串的 {\rm endpos} 集合。但现在我们不妨站在 SA 的角度来思考问题:我们知道如上定义的本质不同的等价类的数量是 O(n) 的,那么我们能否使用 SA 来求解呢?

我们知道每一个等价类在 sa 数组上都会对应一个区间,直观的想法是枚举每个左端点,再计算 |t|=1 时的右端点。然后我们确定这个区间所允许的最长的 t ,再直接计算允许更大的 |t| 的右端点,直到这样的区间不存在为止,时间复杂度 O(n\log n) 。这是一个笨拙的方法,也不具有可扩展性,但其暗示我们考虑按最小值来划分区间的过程,所以我们不妨对 height 数组建立笛卡尔树,并将原串的后缀作为其叶子节点。

观察这棵笛卡尔树,我们可以发现其具备很多优美的性质:每个等价类都在其上对应了一棵子树,也对应了 sa 数组的一段区间;其先序遍历的结果就是后缀数组;任意两个叶子节点的 lca 就是它们的 lcp …… 可以发现,它几乎就是一棵后缀树!

现在,我们终于可以借助这棵笛卡尔树,把 SA 与线段树合并或是 LCT 等数据结构结合起来,这棵树完全可以发挥与 SAM 的 parent tree 相同的效力,而且代码实现也并不复杂。另外值得一提的是,我们甚至可以在这棵树上实现在线的匹配。

这样一来,剩下的过程就与其它的题解无异了。模仿区间数颜色的做法,我们首先将离线询问,再从右往左扫描线。每插入一个新的后缀,其在笛卡尔树上涉及的仅是一条由叶子到根的链,且操作过程形似 LCT 的 access 操作,所以可以使用 LCT 优化。我们只需再维护区间加法和区间求和即可,我的实现是 O(n\log ^2 n) 的,如下:

#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;
}