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

浙公网安备 33010602011771号