AC自动机
前置知识:kmp,trie
问题场景
给你 \(n\) 个模板串和 1 个长度为 \(m\) 的文本,求有多少个模板串在文本里面出现过?
- 考虑暴力:对于 \(n\) 个模板串都在文本里面进行一次 kmp,时间复杂度 \(O(nm)\)
还是步入正题:AC自动机!
先来看一张图:

我们考虑将 \(n\) 个模板串插入到 trie 树里面(绿色点代表是个模板串的末尾)。
显然这个图片里面的模板串是:ACH,ACS,CB,CH,SX,SY。
考虑一个可能的文本:ACHACSXYCB
我们来模拟一下这个过程:
- ACH 是我们可以顺着 trie 直接走到 H 上,确定了有个 ACH 的模板串。
- 如果说是普通的匹配,我们就直接回到根节点,继续开始遍历。但是这样太慢了!我们可以借鉴 kmp 的思路:用当前已经匹配的路径上不平凡后缀在从根开始有相同的路径的最大长度来快速匹配。如果认真学习了 kmp 的人,这句话其实是很好理解的,其实就是我们已经确定了的后缀去找到相同的前缀然后这样就不用重新匹配这一段要找的前缀。
- 根据这个思路,我们会从左侧的 H 直接跳到中间的 H 上,发现这里是个绿点,所以 CH 这个模板串也在文本里面。
- 发现没有了相同的后缀对应的前缀,于是回到根节点重新开始。
- 走到了 S,是个绿点,说明有 ACS。
- 跳到右侧的 S,走到了 X,是个绿点,说明有 SX。
- 跳回根节点,发现没有 Y,继续跳回根节点。
- 走完 CB。
其实很好理解的对吧。
这个匹配过程的时间复杂度:\(O(n + k)\),\(k\) 是匹配到的个数。
发现我们从一个很暴力的匹配到现在这样线性的时间复杂度匹配,我们做的操作其实就是在匹配失败或者成功时跳一下。
在 AC自动机中,我们对于这个跳一下的操作称为:\(j = fail_j\)
怎么求 \(fail\)
其实求的过程也挺好理解的。
- 我们发现由于 \(fail\) 的定义是最长不平凡后缀对应的前缀长度,所以在 trie 树上第二层的 \(fail\) 指针都是指向根节点即为 0.
假设当前在 trie 树上的节点是 \(u\),下一个要走的节点是 \(v\):
-
如果 \(u\) 存在 \(v\) 这个儿子,\(v\) 的 \(fail\) 指针肯定就会指向 \(fail_u\) 对应的 \(v\)。理解:我们对应的前后缀在 trie 树上其实都是从上往下看是否一样,后缀的末尾增加了一个 \(v\),那么就是要在对应的前缀上增加一个 \(v\)。
-
如果 \(u\) 没有这个儿子,那么我们直接把 \(u\) 对应的儿子位置设置成 \(fail_u\) 对应的儿子 \(v\)。
我们一般写 AC自动机的时候会写成 trie 图的形式(其实就是在 trie 这个树上增加若干条边)。在建立 trie 图的时候用 bfs 来写,这个肯定是线性的。
完整代码:
#include <bits/stdc++.h>
using namespace std;
using ll = long long;
using ld = long double;
using ull = unsigned long long;
using ai2 = array<int, 2>;
const int N = 1e4 + 10, M = 1e6 + 10, S = 5e1 + 5;
int n, tr[N * S][26], tot, cnt[N * S];
int fail[N * S];
int que[N * S], hh, tt = -1;
bool is[N * S];
string str;
void insert() {
int idx = 0;
for(char c : str) {
if(tr[idx][c - 'a']) idx = tr[idx][c - 'a'];
else {
tr[idx][c - 'a'] = ++ tot;
idx = tot;
}
}
cnt[idx] ++;
}
void build() {
for(int i = 0; i < 26; i ++) {
if(tr[0][i]) que[++ tt] = tr[0][i];
}
while(hh <= tt) {
int u = que[hh ++];
for(int c = 0; c < 26; c ++) {
if(tr[u][c]) {
fail[tr[u][c]] = tr[fail[u]][c];
que[++ tt] = tr[u][c];
} else tr[u][c] = tr[fail[u]][c];
}
}
}
void solve() {
memset(tr, 0, sizeof tr);
memset(cnt, 0, sizeof cnt);
memset(fail, 0, sizeof fail);
hh = 0, tt = -1;
tot = 0;
cin >> n;
for(int i = 0; i < n; i ++) {
cin >> str;
insert();
}
build();
cin >> str;
int idx = 0, sum = 0;
for(char c : str) {
idx = tr[idx][c - 'a'];
int p = idx;
while(p) {
sum += cnt[p];
cnt[p] = 0;
p = fail[p];
}
}
cout << sum << '\n';
}
int main() {
ios::sync_with_stdio(false);
cin.tie(nullptr);
int T;
cin >> T;
while(T --) {
solve();
}
return 0;
}

浙公网安备 33010602011771号