P5357 【模板】AC 自动机 题解

· · 题解

建议大家先通过AC 自动机(简单版 II),了解 AC 自动机原理。本题方法就是对简单版的优化。

加强版通常都是优化时间或内存,既然本题中内存按照简单版建数组没有问题,那就是优化时间了。但简单版中哪部分代码需要优化呢?

如果我们计算代码的时间复杂度,就会发现这段循环跳 fail 边效率极低:

void query (string x) {
    int len = x.length (), now = 0;
    for (int i = 0; i < len; i++) {
        int c = x[i] - 'a';
        now = trie[now][c];
        for (int j = now; j; j = fail[j])
            vis[j]++;
    }
}

看一下下面的例子(手画勿喷

红色的边是 fail 边。可以看出,在这个例子中,刚才的方法时间复杂度会被卡到 O(n \lvert S \rvert),算法就退化为了 n 次 KMP 算法。我们需要一种避免循环的方法,在询问结束后再处理答案。这就是——

拓扑建图优化

如果我们把刚才图中的 fail 边单独拿出,就会发现

fail 边连接起来的点构成了一棵树!

这是因为每个点都会有 1 个 fail,trie 树的根节点没有 fail,边数比点数少 1,就构成了一棵树。

因此,我们可以在这棵树上进行拓扑排序,用入度小的点的答案更新入度大的点,这样就能在跑 AC 自动机时不跳 fail 边,跑完后统一更新答案,顺便把答案赋给原字符串。

拓扑排序代码:

void topu () {
    queue <int> q;
    for (int i = 1; i <= cnt; i++)
        if (!in[i])
            q.push (i); //入度为0入队
    while (!q.empty ()) {
        int u = q.front ();
        q.pop ();
        for (auto it = flag[u].begin (); it != flag[u].end (); it++)
            ans[*it] = vis[u]; //u点对应的所有字符串
        int v = fail[u];
        vis[v] += vis[u];//更新u点fail边指向的点答案
        in[v]--; 
        if (!in[v])
            q.push (v); //入度为0入队
    }
}

这里的 in 数组需要在添加 fail 边时维护,那么 flag 数组为什么会有“begin”“end”呢?题目中提到数据不保证任意两个模式串不相同,因此我们把 flag 数组设为 vector,把一个点对应的字符串都加进去。其他的细节代码中有注释。

最终代码

#include <bits/stdc++.h>

#define N 200003

using namespace std;

int n, trie[N][26], fail[N], cnt, in[N], vis[N], ans[N];

string t, s;

vector <int> flag[N]; //点对应的字符串

void add (string x, int id) { //建trie树
    int len = x.length (), now = 0;
    for (int i = 0; i < len; i++) {
        int c = x[i] - 'a';
        if (!trie[now][c])
            trie[now][c] = ++cnt;
        now = trie[now][c];
    }
    flag[now].push_back (id);
}

void get_fail () { //添加fail边
    queue <int> q;
    for (int i = 0; i < 26; i++)
        if (trie[0][i])
            q.push (trie[0][i]);
    while (!q.empty ()) {
        int u = q.front ();
        q.pop ();
        for (int i = 0; i < 26; i++) {
            if (trie[u][i]) {
                fail[trie[u][i]] = trie[fail[u]][i];
                in[fail[trie[u][i]]]++; //fail边指向的点入度+1
                q.push (trie[u][i]);
            }
            else trie[u][i] = trie[fail[u]][i];
        }
    }
}

void query (string x) { //查询答案
    int len = x.length (), now = 0;
    for (int i = 0; i < len; i++) {
        int c = x[i] - 'a';
        now = trie[now][c];
        vis[now]++; //不需跳fail边
    }
}

void topu () { //拓扑排序
    queue <int> q;
    for (int i = 1; i <= cnt; i++)
        if (!in[i])
            q.push (i);
    while (!q.empty ()) {
        int u = q.front ();
        q.pop ();
        for (auto it = flag[u].begin (); it != flag[u].end (); it++)
            ans[*it] = vis[u];
        int v = fail[u];
        vis[v] += vis[u];
        in[v]--;
        if (!in[v])
            q.push (v);
    }
}

int main () {
    ios::sync_with_stdio (false);
    cin.tie (nullptr);
    cout.tie (nullptr);
    cin >> n;
    for (int i = 1; i <= n; i++) {
        cin >> t;
        add (t, i);
    }
    cin >> s;
    get_fail ();
    query (s);
    topu ();
    for (int i = 1; i <= n; i++)
        cout << ans[i] << "\n";
    return 0;
}