$O(n^2)$ 算法求助
查看原帖
$O(n^2)$ 算法求助
547908
NightTide楼主2022/10/6 16:20

RT,用 manacher + SAM 做到了 O(n2)O(n^2) 的时间复杂度,思路是用 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;
}
2022/10/6 16:20
加载中...