RT,用 manacher + SAM 做到了 O(n2) 的时间复杂度,思路是用 manacher 找回文串,一旦找到就在 SAM 上跑,并给沿途结点打上标记。如果沿途结点有没有被标记到的,说明当前的子串之前并没有被统计过,就 ans++,否则不加,思路类似 P3649 的 manacher + SAM 没有倍增优化的做法。这里过了,但是貌似被我们教练卡掉了,请大佬们帮忙看看哪里有问题。
#include<bits/stdc++.h>
#define MAXN 8010
#define int long long
using namespace std;
struct node{
int len, link;
// unordered_map<char, int> nxt;
int nxt[30];
bool vis;
};
node sam[MAXN << 1];
int n, m, tot, last, ans;
int len[MAXN << 1];
char s[MAXN], t[MAXN << 1];
void sam_init(){
sam[0].len = 0;
sam[0].link = -1;
sam[0].vis = true;
last = 0;
tot = 0;
}
void sam_extend(char c){
int now = ++tot;
sam[now].len = sam[last].len + 1;
int p = last;
while(p != -1 && !sam[p].nxt[c - 'a']){
sam[p].nxt[c - 'a'] = now;
p = sam[p].link;
}
if(p == -1) sam[now].link = 0;
else{
int q = sam[p].nxt[c - 'a'];
if(sam[p].len + 1 == sam[q].len) sam[now].link = q;
else{
int clone = ++tot;
// sam[clone] = sam[q];
sam[clone].len = sam[p].len + 1;
sam[clone].link = sam[q].link;
for(int i = 0; i < 30; i++) sam[clone].nxt[i] = sam[q].nxt[i];
// sam[clone] = (node){, sam[q].link, sam[q].nxt, sam[q].vis};
while(p != -1 && sam[p].nxt[c - 'a'] == q){
sam[p].nxt[c - 'a'] = clone;
p = sam[p].link;
}
sam[q].link = sam[now].link = clone;
}
}
last = now;
}
void check(int l, int r){
// printf("%d %d\n",l,r);
int now = 0; bool flag = true;
for(int i = l; i <= r; i++){
if(t[i] == '*') continue;
now = sam[now].nxt[t[i] - 'a'];
flag &= sam[now].vis; sam[now].vis = true;
}
if(!flag) ans++;
}
signed main(){
while(~scanf("%s",s + 1)){
ans = 0;
memset(t, 0, sizeof(t));
memset(len, 0, sizeof(len));
sam_init(); n = strlen(s + 1);
for(int i = 1; i <= n; i++) sam_extend(s[i]);
for(int i = 1; i <= tot; i++) sam[i].vis = false;
m = n * 2 + 2;
t[1] = '*'; t[0] = '$'; t[m] = '~';
for(int i = 1; i <= m; i++) t[i << 1] = s[i], t[i << 1 | 1] = '*';
int maxr = 0, mid = 0;
for(int i = 1; i < m; i++){
if(i < maxr) len[i] = min(maxr - i, len[(mid << 1) - i]);
else len[i] = 1;
check(i - len[i] + 1, i + len[i] - 1);
while(t[i - len[i]] == t[i + len[i]]){
len[i]++;
check(i - len[i] + 1, i + len[i] - 1);
}
if(i + len[i] > maxr){
mid = i;
maxr = i + len[i];
}
}
printf("The string '%s' contains %lld palindromes.\n",s + 1, ans);
for(int i = 0; i <= tot; i++){
memset(sam[i].nxt, 0, sizeof(sam[i].nxt));
sam[i].len = sam[i].link = sam[i].vis = 0;
}
}
return 0;
}