后缀数组(SA)

给定长为 \(n\) 的字符串 \(s\),约定 \(i\) 的后缀为 \(s[i\cdots n]\)\(sa_i\) 表示第 \(i\) 小的后缀,\(rk_i\) 表示 \(i\) 的后缀的排名.

求 sa 和 rk 数组

朴素做法直接对每个后缀排序,时间复杂度 \(\mathcal{O}(n^2 \log n)\).

考虑倍增,若当前长度为 \(k\)\(i\) 的后缀表示为 \(s[i\cdots i+k-1]\)\(sa\)\(rk\) 此时维护的都是这个后缀.

考察 \(k \to 2k\) 的变化,新后缀为两端后缀的拼接,此时只需要对 \((s[i\cdots i+k-1],s[i+k\cdots i+2k-1])\) 这个二元组排序即可,可以用计数排序优化,同时因为是增量的,存在进一步优化的可能.

先对第二关键字排序,发现只需要按顺序把 \(i+k\gt n\) 的放在前面,随后按当前 \(sa\) 的顺序放 \(sa_i-k\) 即可.

按照这个索引做稳定计数排序,最后把排名离散化(本质相同的排名也应该相同),当不同的排名到达 \(n\) 时,处理完成.

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

https://www.luogu.com.cn/problem/P3809

//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;

    int m = max(255,n);
    vector<int> rk(n+1),ord(n+1),cnt(m+1),sa(n+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);
        }     

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

    for (int i=1;i<=n;i++){
        cout << sa[i] << ' ';
    }
    cout << '\n';
}

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

    return 0;
}

height 数组

定义 \(height_i = lcp(suf[sa_{i-1}],suf[sa_i])\),也就是排名为 \(i-1\) 的后缀与排名为 \(i\) 的后缀的最长公共前缀. 有一个重要的性质:

\[lcp(suf[x],suf[y]) = \min_{rk_x+1 \le i \le rk_y}{height_i} \]

这样就把求任意两后缀的 \(lcp\) 变成了区间 \(rmq\) 问题.

Kasai 算法求 height 数组

根据结论:若 \(height[rk_i]=k\),则 \(height[rk_{i+1}] \ge k-1\),也就是说 \(i+1\) 只需要从 \(k-1\) 开始拓展,这就让复杂度变成线性的了.

按原顺序遍历,\(i\) 需要匹配的后缀起点 \(j=sa[rk_i-1]\).

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

vector<int> height(n+1);
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);
}

习题

https://www.luogu.com.cn/problem/P3804

题意

给定字符串 \(s\),求 \(s\) 中所有出现次数大于 \(1\) 的子串的出现次数乘上该子串长度的最大值.

\(1\le |s| \le 10^6\).

思路

相当于在 \(height\) 数组选一个区间,最大化 \((r-l+1)\cdot \min_{l+1\le i \le r}{height_i}\),容易想到用单调栈固定最小值.

时间复杂度 \(\mathcal{O}(n \log n)\)\(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;

    int m = 255;
    vector<int> sa(n+1),rk(n+1),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;
    }

    vector<int> height(n+1);
    {
        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);
        }
    }

    stack<int> stk;
    vector<int> left(n+1),right(n+1);
    for (int i=1;i<=n;i++){
        while (!stk.empty() && height[i]<=height[stk.top()]){
            stk.pop();
        }
        left[i] = stk.empty()?0:stk.top();
        stk.push(i);
    }
    while (!stk.empty()){
        stk.pop();
    }

    for (int i=n;i>=1;i--){
        while (!stk.empty() && height[i]<=height[stk.top()]){
            stk.pop();
        }
        right[i] = stk.empty()?n+1:stk.top();
        stk.push(i);
    }

    ll mx = 0;
    for (int i=1;i<=n;i++){
        mx = max(mx,(ll)height[i]*(right[i]-left[i]));       
    }       
    cout << mx << '\n'; 
}

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

    return 0;
}

https://www.luogu.com.cn/problem/P2408

题意

给定长为 \(n\) 的字符串 \(s\),求不同的子串数量.

\(1\le n \le 10^5\).

思路

固定起点分类,起点为 \(sa_i\) 贡献是 \(n-sa_i+1\),容斥掉与前一个的重复部分,也就是 \(height_i\).

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

代码

//author:kzssCCC

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


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

    int m = 255;
    vector<int> sa(n+1),rk(n+1),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;
    }

    vector<int> height(n+1);
    {
        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);
        }
    } 

    ll res = 0;
    for (int i=1;i<=n;i++){
        res += n-sa[i]+1-height[i];
    }
    cout << res << '\n';
}

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

    return 0;
}

https://www.luogu.com.cn/problem/P2852

题意

给定长度为 \(n\) 的数组 \(a\),求出现至少 \(k\) 次的子数组的最大长度.

\(1\le n \le 2\cdot 10^4\).

思路

相当于在 \(height\) 数组选一个长度 \(\ge k-1\) 的区间,最大化区间的最小值,因为区间长度增大区间最小值不增,因此区间长度固定为 \(k-1\),使用单调队列 \(+\) \(multiset\) 维护即可.

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

代码

//author:kzssCCC

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

const int INF = 1e9;

void solve(){
    int n,k;
    cin >> n >> k;

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

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

    vector<int> height(n+1);
    {
        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 && a[i+k]==a[j+k]){
                k++;
            }
            height[rk[i]] = k;
            k = max(k-1,0);
        }
    }

    int mx = 0;
    int l=1,r=1;
    multiset<int> st;
    while (r<k-1){
        st.insert(height[r++]);
    }

    while (r<=n){
        st.insert(height[r++]);
        mx = max(mx,*st.begin());
        st.extract(height[l++]);
    }

    cout << mx << '\n';
}

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

    return 0;
}
posted @ 2026-07-23 16:32  kzssCCC  阅读(14)  评论(0)    收藏  举报