2025 ICPC武汉邀请赛 J 题思路分享(SA,并查集,ST 表)

题意

给定长度为 \(n\) 的字符串,\(q\) 个查询,每个查询学习 \(s\) 中所有 \(s[l\cdots r]\) 开头的子串,求到该查询为止学习的不同字符串的数量.

\(1\le n,q \le 2\cdot 10^5\).

思路

\(P = s[l\cdots r]\),借助 \(sa\) 对本质不同子串的划分,即:当 \(L=sa_i\) 时,\(R\in[sa_i+height_i,n]\).

考虑维护 \(sa\) 每个位置的 \(R\) 的左边界 \(last_i\),贡献为 \(n-last_i+1\),每次增量更新实际上就是对 \(last_i\)\(min\).

找到 \(rk_l\) 左边最后一个以 \(P\) 开头的后缀 \(A\),以及右边最后一个以 \(P\) 开头的后缀 \(B\),只对 \([A\cdots B]\) 产生更新,这可以通过预处理 \(height\)\(ST\) 表,二分实现.

对于 \(A\)\(height_A \lt r-l+1\),因此 \(R\) 覆盖范围为 \([sa_i+r-l,n]\);对于 \(A+1\cdots B\),完全覆盖,即 \([sa_i+height_i,n]\),直接对 \([A\cdots B]\) 暴力更新,用并查集快速跳过被完全覆盖的点.

时间复杂度 \(\mathcal{O}(n\log n)\).

代码

//author:kzssCCC

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


void solve(){
    string s;
    cin >> s;
    int n = s.size();
    s = ' '+s;

    vector<int> sa(n+1),rk(n+1),height(n+1);
    {
        int m = 255;
        vector<int> ord(n+1),cnt(max(n,m)+1);
        for (int i=1;i<=n;i++){
            rk[i] = (unsigned char)s[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[i]]--] = i;
        }

        for (int k=1;;k<<=1){
            int p = 1;
            for (int i=n-k+1;i<=n;i++){
                ord[p++] = i;
            }
            for (int i=1;i<=n;i++){
                if (sa[i]>k){
                    ord[p++] = sa[i]-k;
                }
            }

            fill(cnt.begin(),cnt.begin()+m+1,0);
            for (int i=1;i<=n;i++){
                cnt[rk[ord[i]]]++;
            }
            for (int i=1;i<=m;i++){
                cnt[i] += cnt[i-1];
            }
            for (int i=n;i>=1;i--){
                sa[cnt[rk[ord[i]]]--] = ord[i];     
            }

            swap(rk,ord);
            rk[sa[1]] = 1;
            for (int i=2;i<=n;i++){
                pair<int,int> p1 = {ord[sa[i-1]],sa[i-1]+k<=n?ord[sa[i-1]+k]:-1};
                pair<int,int> p2 = {ord[sa[i]],sa[i]+k<=n?ord[sa[i]+k]:-1};
                rk[sa[i]] = rk[sa[i-1]]+(p1<p2);
            }

            m = rk[sa[n]];
            if (m==n) break;
        }

        int k = 0;
        for (int i=1;i<=n;i++){
            if (rk[i]==1) continue;
            int j = sa[rk[i]-1];
            while (i+k<=n && j+k<=n && s[i+k]==s[j+k]){
                k++;
            }
            height[rk[i]] = k;
            k = max(k-1,0);
        }
    }

    vector<int> lg(n+1);
    for (int i=2;i<=n;i++){
        lg[i] = lg[i>>1]+1;
    } 

    int K = lg[n];
    vector<vector<int>> st(n+1,vector<int>(K+1));
    for (int i=1;i<=n;i++){
        st[i][0] = height[i];
    }
    for (int k=1;k<=K;k++){
        int len = 1<<k;
        int half = len>>1;
        for (int i=1;i+len-1<=n;i++){
            st[i][k] = min(st[i][k-1],st[i+half][k-1]);
        }
    }

    auto query = [&](int l,int r){
        int len = r-l+1;
        int k = lg[len];
        return min(st[l][k],st[r-(1<<k)+1][k]);
    };

    vector<int> last(n+1),p(n+1);
    for (int i=1;i<=n;i++){
        last[i] = n-sa[i]+2;
        p[i] = i+1;
    }

    auto find = [&](int u){
        int v = p[u];
        while (v<=n && last[v]==height[v]+1){
            v = p[v];
        }
        while (u!=v){
            int next = p[u];
            p[u] = v;
            u = next;
        }
        return v;
    };

    ll res = 0;
    int q;
    cin >> q;

    while (q--){
        int l,r;
        cin >> l >> r;
        int p = rk[l];

        int L,R;
        {
            int left=1,right=p-1;
            while (left<=right){
                int mid = left+right >> 1;
                if (query(mid+1,p)>=r-l+1){
                    right = mid-1;
                }
                else{
                    left = mid+1;
                }
            }
            L = left;
        }
        {
            int left=p+1,right=n;
            while (left<=right){
                int mid = left+right >> 1;
                if (query(p+1,mid)<r-l+1){
                    right = mid-1;
                }
                else{
                    left = mid+1;
                }
            }            
            R = right;   
        }

        int j = L+1;
        while (j<=R){
            int nstu = height[j]+1;
            res += last[j]-nstu;
            last[j] = nstu;
            j = find(j); 
        }

        int nstu = min(last[L],r-l+1);
        res += last[L]-nstu;
        last[L] = nstu;

        cout << res << ' ';
    }
    cout << '\n';
}

int main(){
    ios::sync_with_stdio(false);
    cin.tie(0);
    
    int t = 1;
    cin >> t;
    while (t--) solve();

    return 0;
}
posted @ 2026-07-24 20:39  kzssCCC  阅读(8)  评论(0)    收藏  举报