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 边。可以看出,在这个例子中,刚才的方法时间复杂度会被卡到
拓扑建图优化
如果我们把刚才图中的 fail 边单独拿出,就会发现
fail 边连接起来的点构成了一棵树!
这是因为每个点都会有
因此,我们可以在这棵树上进行拓扑排序,用入度小的点的答案更新入度大的点,这样就能在跑 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;
}