SA&SAM学习笔记
后缀数组SA
定义 \(sa_i\) 表示排名为 \(i\) 的后缀第几个,\(rk_i\) 表示后缀 \(i\) 排名多少。
P10469 后缀数组
#include <iostream>
#include <cstdio>
#include <cstring>
#include <algorithm>
#include <stdlib.h>
#include <vector>
#define int long long
#define uint unsigned long long
#define N 300005
using namespace std;
char s[N];
uint seed = 231;
uint h[N],Pow[N];
int rk[N],sa[N],height[N];
signed main(){
scanf("%s",s + 1);
int n = strlen(s + 1);
Pow[0] = 1;
s[n + 1] = '$';
for (int i = 1;i <= n;i ++) Pow[i] = Pow[i - 1] * seed;
for (int i = 1;i <= n;i ++) h[i] = h[i - 1] * seed + s[i];
for (int i = 1;i <= n;i ++) sa[i] = i;
auto gethash = [&](int l,int r) {
return h[r] - h[l - 1] * Pow[r - l + 1];
};
sort(sa + 1,sa + 1 + n,[n,gethash](const int &a,const int &b) {
int lena = n - a + 1,lenb = n - b + 1;
int l = 1,r = min(lena,lenb),res = r + 1;
while(l <= r) {
int mid = l + r >> 1;
if (gethash(a,a + mid - 1) != gethash(b,b + mid - 1)) res = mid,r = mid - 1;
else l = mid + 1;
}
if (res == min(lena,lenb) + 1) return lena < lenb;
return s[a + res - 1] < s[b + res - 1];
});
auto print = [](vector<int> a) -> void {
for (auto i : a) printf("%lld ",i);
putchar('\n');
};
vector<int> ans;
ans.reserve(n);
for (int i = 1;i <= n;i ++) ans.push_back(sa[i] - 1);
print(move(ans));
for (int i = 1;i <= n;i ++) rk[sa[i]] = i;
for (int i = 1,k = 0;i <= n;i ++) {
if (k) k --;
while(s[i + k] == s[sa[rk[i] - 1] + k]) k ++;
height[rk[i]] = k;
}
ans.reserve(n);
for (int i = 1;i <= n;i ++) ans.push_back(height[i]);
print(move(ans));
return 0;
}
目前只会哈希+二分的方法,还不会其他的。
现在需要讲解利用Kasai 算法来 \(\mathcal{O}(n)\) 求解 \(height\) 数组(排名为 \(i\) 的后缀和排名为 \(i-1\) 的后缀的最长公共子串(LCP))。
重要性质:对于任意后缀 \(i(1\leq i\leq n)\),设 \(h_i=height_{rk_i}\)(即后缀 \(i\) 它前一名后缀的 LCP)。
则 \(h_i\geq h_{i-1}-1\)。简要证明:令 \(L=height_{rk_{i-1}}\)(后缀 \(i-1\) 与那个前一名后缀),令 \(p=rk_{i-1},j=sa_{p-1}\),那么 \(L\) 为后缀 \(i-1\) 和后缀 \(j\) 的LCP,因此后缀 \(i\) 和后缀 \(j+1\) 的 LCP 为 \(L-1\),因为 \(j\) 排在 \(i\) 前面,由字典序比较,\(j+1\) 也排在 \(i\) 前面,那么所以至少有 \(L-1\) 可以配对。因此:\(height[rk[i]] ≥ L - 1 = height[rk[i-1]] - 1\)。
因此在实现中,令 \(k=height_{rk_{i-1}}\),然后暴力去找 \(height_{rk_i}\),因为最多是 \(\mathcal{O}(n)\) 的。
应用
最小循环移动位置(JSOI2007 字符加密)
好简单。
#include <iostream>
#include <cstdio>
#include <cstring>
#include <algorithm>
#include <stdlib.h>
#include <vector>
#define int long long
#define uint unsigned long long
#define N 100005
using namespace std;
char s[N << 1];
int rk[N],sa[N];
uint seed = 231,h[N << 1],Pow[N << 1];
signed main(){
scanf("%s",s + 1);
int n = strlen(s + 1);
for (int i = 1;i <= n;i ++) s[i + n] = s[i],sa[i] = i;
Pow[0] = 1;
for (int i = 1;i <= 2 * n;i ++) Pow[i] = Pow[i - 1] * seed,h[i] = h[i - 1] * seed + s[i];
auto gethash = [h](int l,int r) {
return h[r] - Pow[r - l + 1] * h[l - 1];
};
stable_sort(sa + 1,sa + 1 + n,[gethash,n](int x,int y) {
int l = 1,r = n,res = r;
while(l <= r) {
int mid = l + r >> 1;
if (gethash(x,x + mid - 1) != gethash(y,y + mid - 1)) r = mid - 1,res = mid;
else l = mid + 1;
}
return s[x + res - 1] < s[y + res - 1];
});
// for (int i = 1;i <= n;i ++) {
// for (int j = 1;j <= n;j ++) putchar(s[sa[i] + j - 1]);
// putchar('\n');
// }
for (int i = 1;i <= n;i ++) putchar(s[sa[i] + n - 1]);
return 0;
}
至少重复K次子串最长(可重叠)
方法一:直接使用二分答案+hash,用map统计,时间复杂度 \(\mathcal{O}(n\log^2n)\)
方法二:使用后缀数组,发现实际上就是一段height,用单调队列维护这个东西即可,时间复杂度 \(\mathal{O}(n\log^2n)\)(计数排序优化)。
#include <iostream>
#include <cstdio>
#include <cstring>
#include <algorithm>
#include <stdlib.h>
#include <vector>
#define int long long
#define uint unsigned long long
#define N 20005
using namespace std;
int n,k,sa[N],height[N],rk[N],a[N];
uint h[N],seed = 231,Pow[N];
int q[N];
signed main(){
cin >> n >> k;k --;
Pow[0] = 1;
for (int i = 1;i <= n;i ++) {
int x;
scanf("%lld",&x);
a[i] = x;
sa[i] = i;
Pow[i] = Pow[i - 1] * seed;
h[i] = h[i - 1] * seed + x;
}
auto gethash = [](int l,int r) {
return h[r] - h[l - 1] * Pow[r - l + 1];
};
stable_sort(sa + 1,sa + 1 + n,[gethash](int x,int y) {
int l = 1,r = min(n - x + 1,n - y + 1),res = r;
while(l <= r) {
int mid = l + r >> 1;
if (gethash(x,x + mid - 1) != gethash(y,y + mid - 1)) res = mid,r = mid - 1;
else l = mid + 1;
}
if (res == min(n - x + 1,n - y + 1)) return n - x + 1 < n - y + 1;
return a[x + res - 1] < a[y + res - 1];
});
for (int i = 1;i <= n;i ++) rk[sa[i]] = i;
for (int i = 1,k = 0;i <= n;i ++) {
if (k) k --;
while(i + k <= n && sa[rk[i] - 1] + k <= n && a[i + k] == a[sa[rk[i] - 1] + k]) k ++;
height[rk[i]] = k;
}
int head = 1,tail = 0,ans = 0;
for (int i = 2;i <= n;i ++) {
while(head <= tail && q[head] <= i - k) head ++;
while(head <= tail && height[q[tail]] > height[i]) tail --;
q[++tail] = i;
if (i >= k) ans = max(ans,height[q[head]]);
}
cout << ans;
return 0;
}
部分子串问题
在主串 \(T\) 寻找模式串 \(S\),要求在线。
考虑在 \(sa\) 上二分出一个大于等于 \(S\) 的后缀,这样就可以判断其是否存在 \(S\),然后在 \(sa\) 中二分出第一个长度为 \(|S|\) 且大于 \(S\) 的就可以求出来有多少个了。
不同子串个数
考虑到后缀数组有极强的前缀性和高重复性,因此对于求不同子串的个数,我们考虑每一个后缀的贡献就是:
加和即可。
部分比较问题
P2870 [USACO07DEC] Best Cow Line G
方法一,考虑贪心,两边只会取字典序最小的那一个,如果相同,比较后面一位,以此类推,直接二分+hash直接通过。
#include <iostream>
#include <cstdio>
#include <cstring>
#include <algorithm>
#include <stdlib.h>
#include <vector>
#define int long long
#define uint unsigned long long
#define N 500005
using namespace std;
char s[N],t[N];
uint seed = 231,h[2][N],Pow[N];
signed main(){
// scanf("%s",s + 1);
// int n = strlen(s + 1);
int n;
cin >> n;
Pow[0] = 1;
for (int i = 1;i <= n;i ++) {
char x;
Pow[i] = Pow[i - 1] * seed;
cin >> x;
s[i] = x;
t[n - i + 1] = x;
}
for (int i = 1;i <= n;i ++) {
h[0][i] = h[0][i - 1] * seed + s[i];
h[1][i] = h[1][i - 1] * seed + t[i];
}
auto gethash = [](int l,int r,int id) {
return h[id][r] - Pow[r - l + 1] * h[id][l - 1];
};
int i = 1,j = 1,cnt = 0;
while(i <= n - j + 1) {
int l = 1,r = n - j - i + 2,res = r;
while(l <= r) {
int mid = l + r >> 1;
if (gethash(i,i + mid - 1,0) != gethash(j,j + mid - 1,1)) res = mid,r = mid - 1;
else l = mid + 1;
}
if (s[i + res - 1] <= t[j + res - 1]) putchar(s[i]),i ++;
else putchar(t[j]),j ++;
cnt ++;
if (cnt % 80 == 0) putchar('\n');
}
return 0;
}
方法二:同理,只是判断不一样,考虑搞一个反串然后它们拼在一起,然后你这个最后比较rk就可以了(自己推导),那不是会有多的串吗?在这种比较下面是没有任何影响的,你可以观察。
#include <bits/stdc++.h>
using namespace std;
const int MAX = 110000;
int n, m;
char str[MAX];
int SA[MAX], rnk[MAX], height[MAX], tax[MAX], tp[MAX], a[MAX];
void Sort()
{
for (int i = 0; i <= m; ++i)
tax[i] = 0;
for (int i = 1; i <= n; ++i)
tax[rnk[tp[i]]] ++;
for (int i = 1; i <= m; ++i)
tax[i] += tax[i - 1];
for (int i = n; i >= 1; --i)
SA[tax[rnk[tp[i]]] --] = tp[i];
}
bool cmp(int *f, int x, int y, int w)
{
return (f[x] == f[y] && f[x + w] == f[y + w]);
}
void getSA()
{
for (int i = 1; i <= n; ++i)
rnk[i] = a[i], tp[i] = i;
m = 127, Sort();
for (int w = 1, p = 1, i; p < n; w += w, m = p)
{
for (p = 0, i = n - w + 1; i <= n; ++i)
tp[++p] = i;
for (i = 1; i <= n; ++i)
if (SA[i] > w)
tp[++p] = SA[i] - w;
Sort(); swap(rnk, tp); rnk[SA[1]] = p = 1;
for (int i = 2; i <= n; ++i)
rnk[SA[i]] = cmp(tp, SA[i], SA[i - 1], w) ? p : ++p;
}
}
void init()
{
scanf("%d", &n);
for (int i = 0; i < n; ++i)
scanf("%s", str + i);
for (int i = 0; i < n; ++i)
a[2 * (n + 1) - i - 1] = a[i + 1] = str[i];
n = 2 * n + 1;//中间字符可以加可以不加,看个人习惯,一般也可以用#。
getSA();
}
char ans[MAX];
int cnt;
int main()
{
init();
getSA();
n = (n - 1) / 2;
int L = 1, R = n;
while (L < R)
{
if (a[L] < a[R])
ans[++cnt] = (char)(a[L++]);
else if (a[L] > a[R])
ans[++cnt] = (char)(a[R--]);
else
{
if (rnk[L] < rnk[2 * (n + 1) - R])
ans[++cnt] = (char)(a[L++]);
else
ans[++cnt] = (char)(a[R--]);
}
}
ans[++cnt] = (char)(a[L]);
n = strlen(ans + 1);
for (int i = 1; i <= n; ++i)
{
printf("%c", ans[i]);
if (i % 80 == 0) putchar('\n');
}
return 0;
}

浙公网安备 33010602011771号