CC++ & Algorithm

AC自动机(Trie上的KMP)

极难2
语言版本:通用
概述:AC自动机是在字典树基础上加入失配指针的一种多模式匹配算法,可以一次性在一段文本中查找多个关键词,像超级“查找替换”工具。

一次扫描,多词命中: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(文本长度 + 模式串总长度),非常高效。

原理与核心思想

构建步骤

  1. 构建字典树:将所有模式串插入到Trie中,并标记每个单词的结尾节点。
  2. 构建失配指针(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链上的节点可能也是模式串的结尾),记录匹配到的模式串。

新手容易犯的错误

  1. 忘记构建fail指针就直接搜索:没有fail指针,AC自动机就退化成普通的字典树,只能从根开始一个个字符匹配,无法利用后缀信息,导致漏匹配或效率极低。
  2. 匹配时只检查当前节点,不检查fail链:例如模式串有"he"和"she",当文本中扫描到"she"的'e'时,当前节点是"she"的结尾,但fail链上可能还有"he"的结尾。如果只检查当前节点,就会漏掉"he"。
  3. 插入空字符串作为模式串:空字符串会导致根节点被标记为结尾,匹配时无限循环或出现奇怪行为。通常应避免插入空字符串。
  4. 字符集假设错误:代码中假设只有小写字母(26个),如果输入包含大写字母、数字或中文,需要调整字符集映射。
  5. 节点数组大小不够:实际模式串总长度可能超过预设的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算法(适用于长模式串)、后缀自动机等。

总结要点

  1. AC自动机 = Trie + 失配指针(fail指针),本质上是将KMP的思想推广到多模式匹配。
  2. 构建过程:先建Trie,再用BFS计算每个节点的fail指针。fail指针指向当前字符串的最长后缀(且是某个模式串的前缀)。
  3. 匹配过程:扫描文本,利用fail指针避免回溯,时间复杂度O(文本长度 + 模式串总长度)。
  4. 应用:敏感词过滤、代码搜索、基因序列分析等。很多搜索引擎的“查找关键词”功能底层就用到了AC自动机或其变种。
  5. 注意:实际实现中,为了提高效率,常常会将Trie的缺失边补全(类似于Trie图),从而省去匹配时的while循环,使每个字符的匹配变成O(1)。

AC自动机是个非常强大的算法,但理解它需要先掌握字典树和KMP。建议读者自己手动画一棵Trie并计算fail指针,加深理解。

例题精讲

1单选题

在AC自动机的构建过程中,失配指针(fail指针)的作用是什么?

A指向当前节点在Trie树中的父节点
B指向当前节点在Trie树中的兄弟节点
C指向当前节点匹配失败时应跳转到的节点,该节点是当前字符串的最长后缀对应的前缀节点
D指向根节点,表示重新开始匹配
2单选题

给定一个AC自动机(Trie树已构建完成),若当前匹配到节点u,且u有孩子节点c,但文本的下一个字符与c不匹配,下一步应该如何处理?

A立即将当前节点跳转到根节点,并从根节点开始尝试匹配
B将当前节点跳转到u的fail指针指向的节点,然后继续尝试匹配字符c
C将当前节点跳转到u的父节点,再尝试匹配字符c
D放弃当前文本位置,移动文本指针到下一个字符,并将当前节点置为根节点
3判断题

构建AC自动机时,所有第一层节点(深度为1的节点)的fail指针都应该指向根节点。

4判断题

在AC自动机中,如果一个节点对应一个模式串的结尾,则该节点的所有fail指针指向的节点也一定对应某个模式串的结尾。

5填空题
下面是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] = ___;
            }
        }
    }
}