AC自动机(Trie上的KMP)
极难2一次扫描,多词命中:AC自动机(Trie上的KMP)
你正在开发一个聊天软件,用户发消息时需要实时过滤几十个敏感词,比如"暴力"、"色情"、"赌博"、"诈骗"……如果每次用KMP一个词一个词地查,1000个关键词、10万字的文章就要做1000次查找,太慢了!有没有办法像网一样一次性把所有关键词都捞出来?有,这就是AC自动机(Aho-Corasick Automaton)。它把多个关键词组织成一棵字典树,再添上特殊的“失配指针”,在一遍扫描文本的过程中同时匹配所有关键词,时间复杂度只有 O(文本长度 + 模式串总长度)。它就像给字典树装上了KMP的灵魂,所以也叫“Trie上的KMP”。
生活中的例子:文章的敏感词过滤
假设你运营一个社交平台,需要在一篇用户发布的文章中找出所有违禁词(比如“暴力”、“色情”、“赌博”等)。如果用一个一个关键词去文章中查找,关键词很多时效率很低。比如有1000个关键词,文章有10万字,就需要做1000次KMP或字符串查找,非常慢。
有没有一种方法可以“同时”查找所有关键词呢?AC自动机(Aho-Corasick Automaton)就是为此设计的。它把所有的关键词构建成一棵字典树,然后在树上添加一种叫“失配指针(fail)”的边,当匹配失败时,可以快速跳到另一个分支继续匹配,就像KMP算法中的next数组一样,所以被称为“Trie上的KMP”。
AC自动机可以在一遍扫描文本的过程中,同时匹配多个模式串,时间复杂度为O(文本长度 + 模式串总长度),非常高效。
原理与核心思想
构建步骤:
- 构建字典树:将所有模式串插入到Trie中,并标记每个单词的结尾节点。
- 构建失配指针(fail指针):这是AC自动机的核心。对于Trie中的每个节点u,它的fail指针指向另一个节点v,使得从根到v的字符串是从根到u的字符串的最长后缀(且v不能是u本身,并且这个后缀必须是某个模式串的前缀)。换句话说,如果在节点u处匹配失败,我们就通过fail指针跳转到另一个分支,继续尝试匹配,以避免从头开始。
构建fail指针的方法:使用BFS(广度优先搜索)。根节点的fail设为空(或自己),根的子节点的fail设为根。然后对于每个节点u,遍历它的子节点v(对应字符c),假设u的fail是f,如果f有字符c的子节点,则v的fail指向f的那个子节点;否则v的fail指向根。
匹配过程:
- 从根节点开始,依次遍历文本中的每个字符。
- 如果当前节点有字符c的子节点,则移动到该子节点。
- 如果没有,则沿着fail指针回退,直到找到一个有c子节点的节点,或者到达根。
- 移动到相应节点后,检查该节点以及它通过fail链能到达的所有节点(因为fail链上的节点可能也是模式串的结尾),记录匹配到的模式串。
新手容易犯的错误
- 忘记构建fail指针就直接搜索:没有fail指针,AC自动机就退化成普通的字典树,只能从根开始一个个字符匹配,无法利用后缀信息,导致漏匹配或效率极低。
- 匹配时只检查当前节点,不检查fail链:例如模式串有"he"和"she",当文本中扫描到"she"的'e'时,当前节点是"she"的结尾,但fail链上可能还有"he"的结尾。如果只检查当前节点,就会漏掉"he"。
- 插入空字符串作为模式串:空字符串会导致根节点被标记为结尾,匹配时无限循环或出现奇怪行为。通常应避免插入空字符串。
- 字符集假设错误:代码中假设只有小写字母(26个),如果输入包含大写字母、数字或中文,需要调整字符集映射。
- 节点数组大小不够:实际模式串总长度可能超过预设的MAXN,导致越界。建议使用动态数组(vector)。
示意图
假设模式串有:he, she, his, hers。构建的Trie(只显示部分关键节点)和fail指针如下(简化ASCII):
root(0)
/ | \
h s ...
/ |
e(1) h(4)
/ \ |
r i(3) e(5)
/ \ \
s(2) s(6) ...
- 节点1 (he)的fail?根节点没有以'e'开头的子节点?实际上,节点1的父亲是'h',根节点有子节点's',但's'不是'e'。最长的后缀是'e',但根没有'e'的子节点,所以fail指向根。
- 节点2 (hers)的父亲是'r',但'r'的fail是根,所以最终fail指向根。
- 节点3 (his)的父亲是'i',i的fail是根?实际上需要更详细,但这里不展开。
为了清晰,我们直接通过代码来理解。
C++完整代码实现
下面实现一个简单的AC自动机,支持多个模式串插入和文本匹配,输出所有匹配的模式串及其位置(起始下标)。为了简化,我们只记录匹配到哪些模式串,不记录位置。
#include <iostream>
#include <queue>
#include <vector>
#include <string>
#include <cstring>
using namespace std;
const int MAXN = 1000; // 假设总节点数不超过1000
const int CHAR_SIZE = 26; // 小写字母
struct ACNode {
int next[CHAR_SIZE]; // 子节点索引,-1表示不存在
int fail; // 失配指针
bool isEnd; // 是否是某个模式串的结尾
int len; // 记录该节点代表的前缀长度(方便输出单词)
// 可以根据需要增加一个vector<int> output; 存储以该节点结尾的模式串编号
};
class ACAutomaton {
private:
vector<ACNode> nodes; // 所有节点
int sz; // 当前节点数
int newNode() {
ACNode node;
memset(node.next, -1, sizeof(node.next)); // 初始化为-1
node.fail = 0;
node.isEnd = false;
node.len = 0;
nodes.push_back(node);
return sz++;
}
public:
ACAutomaton() {
nodes.clear();
sz = 0;
newNode(); // 根节点为0
}
// 插入模式串
void insert(const string& pattern) {
int cur = 0; // 从根开始
for (char ch : pattern) {
int idx = ch - 'a'; // 字符转数字
if (nodes[cur].next[idx] == -1) {
int newId = newNode();
nodes[cur].next[idx] = newId;
nodes[newId].len = nodes[cur].len + 1; // 前缀长度+1
}
cur = nodes[cur].next[idx];
}
nodes[cur].isEnd = true; // 标记结尾
}
// 构建fail指针(BFS)
void buildFail() {
queue<int> q;
// 初始化第一层节点的fail为根
for (int i = 0; i < CHAR_SIZE; i++) {
if (nodes[0].next[i] != -1) {
int child = nodes[0].next[i];
nodes[child].fail = 0;
q.push(child);
} else {
// 为了后续方便,将不存在的边指向根(虚拟节点),实际也可以不处理
nodes[0].next[i] = 0;
}
}
while (!q.empty()) {
int u = q.front(); q.pop();
for (int i = 0; i < CHAR_SIZE; i++) {
int v = nodes[u].next[i];
if (v != -1) {
// 计算v的fail指针
int f = nodes[u].fail;
while (f != 0 && nodes[f].next[i] == -1) {
f = nodes[f].fail;
}
if (nodes[f].next[i] != -1 && nodes[f].next[i] != v) {
// 注意:这里要检查实际存在,因为我们已经把不存在的边指向根,所以实际上会更简单
// 下面用标准写法:
nodes[v].fail = nodes[f].next[i];
} else {
nodes[v].fail = 0;
}
// 如果fail节点也是结尾,也要标记(因为匹配时需传递)
if (nodes[nodes[v].fail].isEnd) {
nodes[v].isEnd = true; // 简化处理:只要fail链上有结尾,当前节点也视为结尾
}
q.push(v);
}
}
}
}
// 在文本text中查找所有模式串出现的位置,返回匹配到的模式串列表(每个模式串重复出现只记录一次?这里不具体输出位置)
void search(const string& text) {
int cur = 0; // 当前节点
for (int i = 0; i < text.size(); i++) {
int idx = text[i] - 'a';
// 如果当前节点没有该字符的子节点,就沿着fail跳
while (cur != 0 && nodes[cur].next[idx] == -1) {
cur = nodes[cur].fail;
}
if (nodes[cur].next[idx] != -1) {
cur = nodes[cur].next[idx];
}
// 检查cur节点及其fail链上是否有结尾
int temp = cur;
while (temp != 0) {
if (nodes[temp].isEnd) {
// 可以在这里记录匹配成功,比如输出单词(从根到temp的路径)
// 为了演示,我们输出匹配到的单词长度(或单词本身)
string word = "";
// 可以根据len回溯获取单词,但这里省略
cout << "在位置 " << i - nodes[temp].len + 1 << " 匹配到单词(长度" << nodes[temp].len << ")" << endl;
}
temp = nodes[temp].fail;
}
}
}
};
int main() {
ACAutomaton ac;
ac.insert("he");
ac.insert("she");
ac.insert("his");
ac.insert("hers");
ac.buildFail();
string text = "ushers";
cout << "文本: " << text << endl;
ac.search(text);
// 期望输出:在位置1匹配到"she",位置2匹配到"he",位置2匹配到"hers"等
return 0;
}
Python完整代码实现
Python实现AC自动机,使用类和队列。
from collections import deque
class ACNode:
def __init__(self):
self.next = {} # 子节点字典,键为字符,值为节点
self.fail = None # 失配指针
self.is_end = False # 是否为模式串结尾
self.length = 0 # 该节点代表的前缀长度
class ACAutomaton:
def __init__(self):
self.root = ACNode()
self.root.length = 0
def insert(self, pattern: str) -> None:
cur = self.root
for ch in pattern:
if ch not in cur.next:
new_node = ACNode()
new_node.length = cur.length + 1
cur.next[ch] = new_node
cur = cur.next[ch]
cur.is_end = True
def build_fail(self) -> None:
"""使用BFS构建失配指针"""
q = deque()
# 第一层节点的fail指向根
for child in self.root.next.values():
child.fail = self.root
q.append(child)
# 为了简化,将根节点没有的子节点视为指向根(自动处理)
while q:
u = q.popleft()
for ch, v in u.next.items():
# 计算v的fail
f = u.fail
while f is not None and ch not in f.next:
f = f.fail
if f is not None and ch in f.next:
v.fail = f.next[ch]
else:
v.fail = self.root
# 如果fail节点是结尾,当前节点也视为实际上可匹配(传递)
if v.fail.is_end:
v.is_end = True # 简化处理
q.append(v)
def search(self, text: str) -> list:
"""返回所有匹配的模式串及其起始位置(起始位置从0开始)"""
result = []
cur = self.root
for i, ch in enumerate(text):
# 沿fail指针回退直到能找到ch子节点或到根
while cur is not self.root and ch not in cur.next:
cur = cur.fail
if ch in cur.next:
cur = cur.next[ch]
# 检查cur及其fail链上的所有节点
temp = cur
while temp is not None:
if temp.is_end:
# 匹配到一个模式串,其长度为temp.length
start = i - temp.length + 1
# 由于我们简化了is_end传递,这里可能重复匹配,但实际每个模式串只会在其结尾节点被记录一次
# 如果需要记录单词内容,可以存储从根到temp的字符串(但需要额外信息)
result.append((start, temp.length)) # 起始位置和长度
temp = temp.fail
return result
if __name__ == "__main__":
ac = ACAutomaton()
ac.insert("he")
ac.insert("she")
ac.insert("his")
ac.insert("hers")
ac.build_fail()
text = "ushers"
matches = ac.search(text)
print("文本:", text)
for start, length in matches:
print(f"位置{start} 匹配到长度为{length}的单词")
运行结果示例
文本: ushers
输出(C++类似):
位置1 匹配到长度为3的单词 # "she"
位置2 匹配到长度为2的单词 # "he"
位置2 匹配到长度为4的单词 # "hers"
注意:位置从0开始或1开始取决于实现。上述Python代码中位置从0开始,所以位置1对应第二个字符(s),位置2对应第三个字符(h)。
应用场景
- 敏感词过滤:聊天、评论系统实时检测违禁词。
- 病毒特征码扫描:杀毒软件用AC自动机在一堆文件中快速查找已知病毒特征。
- 基因序列匹配:在一段DNA序列中同时查找多种基因片段。
- 搜索引擎关键词高亮:当用户搜索多个词时,快速在文本中标记所有匹配位置。
相关知识点指引
- 字典树(Trie):AC自动机的基础,先掌握Trie的构建和遍历。
- KMP算法:理解单模式串匹配中的next数组,有助于理解fail指针的本质。
- 双数组Trie:一种高效存储Trie的方法,可以进一步优化AC自动机的空间和速度。
- 多模式匹配的其他算法:如Wu-Manber算法(适用于长模式串)、后缀自动机等。
总结要点
- AC自动机 = Trie + 失配指针(fail指针),本质上是将KMP的思想推广到多模式匹配。
- 构建过程:先建Trie,再用BFS计算每个节点的fail指针。fail指针指向当前字符串的最长后缀(且是某个模式串的前缀)。
- 匹配过程:扫描文本,利用fail指针避免回溯,时间复杂度O(文本长度 + 模式串总长度)。
- 应用:敏感词过滤、代码搜索、基因序列分析等。很多搜索引擎的“查找关键词”功能底层就用到了AC自动机或其变种。
- 注意:实际实现中,为了提高效率,常常会将Trie的缺失边补全(类似于Trie图),从而省去匹配时的while循环,使每个字符的匹配变成O(1)。
AC自动机是个非常强大的算法,但理解它需要先掌握字典树和KMP。建议读者自己手动画一棵Trie并计算fail指针,加深理解。
例题精讲
在AC自动机的构建过程中,失配指针(fail指针)的作用是什么?
给定一个AC自动机(Trie树已构建完成),若当前匹配到节点u,且u有孩子节点c,但文本的下一个字符与c不匹配,下一步应该如何处理?
构建AC自动机时,所有第一层节点(深度为1的节点)的fail指针都应该指向根节点。
在AC自动机中,如果一个节点对应一个模式串的结尾,则该节点的所有fail指针指向的节点也一定对应某个模式串的结尾。
下面是AC自动机构建过程中BFS设置fail指针的代码片段,请补充完整。假设Trie树使用数组存储,ch[u][c]表示节点u的字符c子节点,fail[u]表示节点u的失配指针,字符集大小为26(小写字母)。queue用于BFS。
void build() {
queue<int> q;
for (int c = 0; c < 26; c++) {
if (ch[0][c]) {
fail[ch[0][c]] = 0;
q.push(ch[0][c]);
}
}
while (!q.empty()) {
int u = q.front(); q.pop();
for (int c = 0; c < 26; c++) {
if (ch[u][c]) {
int v = ch[u][c];
fail[v] = ch[fail[u]][c];
q.push(v);
} else {
ch[u][c] = ___;
}
}
}
}