AC自动机
AC自动机
前置知识:
- Trie
- KMP
简介
AC自动机,众所周知,就是Trie上跑KMP。
但是并不是就和KMP算法完全一样,只是其中一步转化用到了KMP失配时处理的思想,来加快算法的速度。所以就算不会KMP应该也是能学AC自动机的,如果会的话更好理解。
变量理解
构建Fail
AC自动机最难理解的部分。
先看这样一颗Trie树:(用蓝色标注的为单词结尾字母)

表示插入的模式串分别为:have him his is
再给定一串文本串:havehisit
我们直接模拟一次试试看。
容易发现可以匹配到2,3,4,5(have匹配完毕)。此时文本串的下一个是h,模式串下一个是空,失配了,我们回到根去重新开始找。又从2匹配起,然后匹配到6,7(his匹配完毕)。此时文本串的下一个是i,模式串下一个是空,失配,再回到根去。等等,要回到根吗?可以发现此时is这个模式串还没有匹配,但是文本串已经和his匹配完了,这不就漏掉一个吗?或者说,如果匹配到is的时候又回到根节点开始匹配,不就浪费很多时间吗?
容易发现,出现这样的情况,当且仅当一个模式串是另一个模式串的后缀,那么要怎么处理呢?想想图上这种情况,我们需要在匹配完7之后跳到10去完成匹配,那能不能构建一个指针,从7指向10呢?
这就是AC自动机最重要的思想,Fail,也就是失配指针。我通常把它理解成一条有向边,这也是后面拓扑建图优化的思想。
给出它的具体意义:令点\(u\)的Fail指针指向\(v\)。设根节点到点\(u\)的字符串为\(a\),根节点到点\(v\)的字符串为\(b\)。那么\(b\)是\(a\)的最大后缀。
最大就保证了最优。
具体实现的细节还是挺多的:
- 根节点的Fail是什么?为了方便后面的实现,我们建立一个虚拟节点0,并让0的所有儿子指向根,根的Fail指针指向0。具体原因看细节2。
- 如果我们在遍历时,遍历到空,也就是不存在节点\(v\),我们就把当前节点的\(v\)这个儿子设成当前节点的Fail边指向的节点的这个儿子。这样一定不会出错,至于为什么不会出错,建议手动画图理解,只可意会不可言传
因为我说不清楚。而且这样做可以保证存在性,就算返回值是0,0的儿子都是根,从根开始匹配当然不会出错。
Talk is cheap,show me your code.
inline void GetFail() {
for (int i=0;i<26;++i) t[0].ch[i]=1;
t[1].fail=0;
queue<int> q;
q.push(1);
while (!q.empty()) {
int u=q.front(); q.pop();
for (int i=0;i<26;++i) {
int v=t[u].ch[i],to=t[u].fail;
if (!v) {t[u].ch[i]=t[to].ch[i]; continue;}
t[v].fail=t[to].ch[i];
q.push(v);
}
}
}
例题
1.模板:P3808
显然就是把上面的搬抄一遍就好了。
#include<bits/stdc++.h>
#define ri register int
using namespace std;
const int maxn=1e6+7;
int n,tot=1;
char s[maxn];
struct node {
int ch[30],fail,flag;
}t[maxn];
inline int read() {
int s=0,w=1; char ch=getchar();
while (ch<'0' || ch>'9') w=(ch=='-')?-1:1, ch=getchar();
while (ch>='0' && ch<='9') s=((s<<1)+(s<<3)+(ch^48)), ch=getchar();
return s*w;
}
inline void Insert(char *s) {
int u=1,len=strlen(s);
for (int i=0;i<len;++i) {
int v=(s[i]-'a');
if (!t[u].ch[v]) t[u].ch[v]=(++tot);
u=t[u].ch[v];
}
t[u].flag++;
}
inline void GetFail() {
for (int i=0;i<26;++i) t[0].ch[i]=1;
t[1].fail=0;
queue<int> q;
q.push(1);
while (!q.empty()) {
int u=q.front(); q.pop();
for (int i=0;i<26;++i) {
int v=t[u].ch[i],to=t[u].fail;
if (!v) {t[u].ch[i]=t[to].ch[i]; continue;}
t[v].fail=t[to].ch[i];
q.push(v);
}
}
}
inline int Query(char *s) {
int u=1,ans=0,len=strlen(s);
for (int i=0;i<len;++i) {
int now=(s[i]-'a'),v=t[u].ch[now];
while (v>1 && t[v].flag!=-1) {
ans+=t[v].flag;
t[v].flag=-1;
v=t[v].fail;
}
u=t[u].ch[now];
}
return ans;
}
int main() {
n=read();
for (int i=1;i<=n;++i) {
scanf("%s",s);
Insert(s);
}
GetFail();
scanf("%s",s);
printf("%d",Query(s));
return 0;
}
2.P3796
加强版。比原来那个题多加一个数组,记录一下每个模式串出现的次数,排个序输出即可。
#include<bits/stdc++.h>
#define ri register int
using namespace std;
const int maxn=1e6+7;
int n,tot=1;
char s[200][100],T[maxn];
struct node {
int ch[30],fail,flag;
inline void clear() {memset(ch,0,sizeof(ch)); flag=fail=0;}
}t[maxn];
struct node2{
int ans,pos;
}a[200];
inline bool cmp(node2 x,node2 y) {
return ((x.ans>y.ans) || (x.ans==y.ans && x.pos<y.pos));
}
inline int read() {
int s=0,w=1; char ch=getchar();
while (ch<'0' || ch>'9') w=(ch=='-')?-1:1, ch=getchar();
while (ch>='0' && ch<='9') s=((s<<1)+(s<<3)+(ch^48)), ch=getchar();
return s*w;
}
inline void clr() {
for (int i=0;i<=tot;++i) t[i].clear();
tot=1;
}
inline void Insert(char *s,int id) {
int u=1,len=strlen(s);
for (int i=0;i<len;++i) {
int v=(s[i]-'a');
if (!t[u].ch[v]) t[u].ch[v]=(++tot);
u=t[u].ch[v];
}
t[u].flag=id;
}
inline void GetFail() {
for (int i=0;i<26;++i) t[0].ch[i]=1;
t[1].fail=0;
queue<int> q;
q.push(1);
while (!q.empty()) {
int u=q.front(); q.pop();
for (int i=0;i<26;++i) {
int v=t[u].ch[i],to=t[u].fail;
if (!v) {t[u].ch[i]=t[to].ch[i]; continue;}
t[v].fail=t[to].ch[i];
q.push(v);
}
}
}
inline void Query(char *s) {
int u=1,len=strlen(s);
for (int i=0;i<len;++i) {
int now=s[i]-'a',v=t[u].ch[now];
for (int j=v;j>1;j=t[j].fail) ++a[t[j].flag].ans; //记录出现次数
u=t[u].ch[now];
}
}
int main() {
while (1) {
n=read(); if (!n) return 0;
clr();
for (int i=1;i<=n;++i) {
scanf("%s",s[i]);
a[i].ans=0,a[i].pos=i;
Insert(s[i],i);
}
GetFail();
scanf("%s",T);
Query(T);
sort(a+1,a+n+1,cmp);
int maxx=a[1].ans;
cout<<maxx<<endl;
for (int i=1;i<=n;++i) {
if (a[i].ans==maxx) cout<<s[a[i].pos]<<endl;
else break;
}
}
return 0;
}
3.P5357
二次加强版。需要拓扑建图优化。
具体做法:假设每个点都向它的Fail指针连一条有向边,那么每个点的出度为1,入度不定。按照一般拓扑的思想对它们排序,并让Fail边指向的那个点继承这条边的另一个节点的答案,累加即可。这样可以保证Trie上的每个点只经过一次,复杂度就是\(O(模式串总长度)\)。
#include<bits/stdc++.h>
#define ri register int
using namespace std;
const int maxn=2e6+7;
int n,tot=1,a[maxn],ind[maxn],cnt[maxn];
char s[maxn],T[maxn];
struct node {
int ch[30],fail,flag,ans;
inline void clear() {memset(ch,0,sizeof(ch)); flag=fail=0;}
}t[maxn];
inline int read() {
int s=0,w=1; char ch=getchar();
while (ch<'0' || ch>'9') w=(ch=='-')?-1:1, ch=getchar();
while (ch>='0' && ch<='9') s=((s<<1)+(s<<3)+(ch^48)), ch=getchar();
return s*w;
}
inline void clr() {
for (ri i=0;i<=tot;++i) t[i].clear();
tot=1;
}
inline void Insert(char *s,int id) {
int u=1,len=strlen(s);
for (ri i=0;i<len;++i) {
int v=(s[i]-'a');
if (!t[u].ch[v]) t[u].ch[v]=(++tot);
u=t[u].ch[v];
}
if (!t[u].flag) t[u].flag=id;
a[id]=t[u].flag;
}
inline void GetFail() {
for (int i=0;i<26;++i) t[0].ch[i]=1;
t[1].fail=0;
queue<int> q;
q.push(1);
while (!q.empty()) {
int u=q.front(); q.pop();
for (ri i=0;i<26;++i) {
int v=t[u].ch[i],to=t[u].fail;
if (!v) {t[u].ch[i]=t[to].ch[i]; continue;}
t[v].fail=t[to].ch[i]; ++ind[t[v].fail];
q.push(v);
}
}
}
inline void Query(char *s) {
int u=1,len=strlen(s);
for (ri i=0;i<len;++i) {
u=t[u].ch[s[i]-'a'];
++t[u].ans;
}
}
inline void Topsort() {
queue<int> q;
for (int i=1;i<=tot;++i) if (!ind[i]) q.push(i);
while (!q.empty()) {
int u=q.front(); q.pop();
cnt[t[u].flag]=t[u].ans;
int v=t[u].fail; --ind[v];
t[v].ans+=t[u].ans;
if (!ind[v]) q.push(v);
}
}
int main() {
n=read();
for (ri i=1;i<=n;++i) {
scanf("%s",s);
Insert(s,i);
}
GetFail();
scanf("%s",T);
Query(T);
Topsort();
for (ri i=1;i<=n;++i) printf("%d\n",cnt[a[i]]);
return 0;
}
4.P3966
乍一看和二次加强版并无什么不同。把模式串累加即可得到文本串。双倍经验!好,快乐复制粘贴测样例。好!WA了!
好,定睛一看发现一样的字母会出现在不同的单词里。若单词首尾连接处能形成一个模式串,就会出现重复计算。
处理方法是把每个单词加进去时,在末尾加上一个特殊符号,表示单词结束。比如'0'。在查询时,如果访问到0,返回根接着查就好了。
#include<bits/stdc++.h>
#define ri register int
using namespace std;
const int maxn=2e6+7;
int n,tot=1,a[maxn],ind[maxn],cnt[maxn],lent;
char s[maxn],T[maxn];
struct node {
int ch[30],fail,flag,ans;
inline void clear() {memset(ch,0,sizeof(ch)); flag=fail=0;}
}t[maxn];
inline int read() {
int s=0,w=1; char ch=getchar();
while (ch<'0' || ch>'9') w=(ch=='-')?-1:1, ch=getchar();
while (ch>='0' && ch<='9') s=((s<<1)+(s<<3)+(ch^48)), ch=getchar();
return s*w;
}
inline void clr() {
for (ri i=0;i<=tot;++i) t[i].clear();
tot=1;
}
inline void Insert(char *s,int id) {
int u=1,len=strlen(s);
for (ri i=0;i<len;++i) {
int v=(s[i]-'a');
if (!t[u].ch[v]) t[u].ch[v]=(++tot);
u=t[u].ch[v];
}
if (!t[u].flag) t[u].flag=id;
a[id]=t[u].flag;
}
inline void GetFail() {
for (int i=0;i<26;++i) t[0].ch[i]=1;
t[1].fail=0;
queue<int> q;
q.push(1);
while (!q.empty()) {
int u=q.front(); q.pop();
for (ri i=0;i<26;++i) {
int v=t[u].ch[i],to=t[u].fail;
if (!v) {t[u].ch[i]=t[to].ch[i]; continue;}
t[v].fail=t[to].ch[i]; ++ind[t[v].fail];
q.push(v);
}
}
}
inline void Query(char *s) {
int u=1,len=strlen(s);
for (ri i=0;i<len;++i) {
if (s[i]=='0') {u=1; continue;}
u=t[u].ch[s[i]-'a'];
++t[u].ans;
}
}
inline void Topsort() {
queue<int> q;
for (int i=1;i<=tot;++i) if (!ind[i]) q.push(i);
while (!q.empty()) {
int u=q.front(); q.pop();
cnt[t[u].flag]=t[u].ans;
int v=t[u].fail; --ind[v];
t[v].ans+=t[u].ans;
if (!ind[v]) q.push(v);
}
}
int main() {
n=read();
for (ri i=1;i<=n;++i) {
scanf("%s",s);
int pre=lent,lens=strlen(s);
lent+=lens;
for (int i=pre,j=0;i<lent;++i,++j) T[i]=s[j];
++lent;
T[lent-1]='0';
Insert(s,i);
}
GetFail();
Query(T);
Topsort();
for (ri i=1;i<=n;++i) printf("%d\n",cnt[a[i]]);
return 0;
}
5.P2292
显然是个比较裸的AC自动机,先把模板抄一遍,发现过不了。这个时候我们应该打开题解区想想优化。
优化:我们开一个布尔类型的数组来记录当前串是否是一篇文章的末尾,如果是末尾就直接记录答案跳出循环,这样可以加快速度。
#include<bits/stdc++.h>
#define ri register int
using namespace std;
const int maxn=2e6+7;
int n,m,tot=1,a[maxn],cnt[maxn],lent;
char s[maxn],T[maxn];
bool book[maxn];
struct node {
int ch[30],fail,flag,ans;
inline void clear() {memset(ch,0,sizeof(ch)); flag=fail=0;}
}t[maxn];
inline int read() {
int s=0,w=1; char ch=getchar();
while (ch<'0' || ch>'9') w=(ch=='-')?-1:1, ch=getchar();
while (ch>='0' && ch<='9') s=((s<<1)+(s<<3)+(ch^48)), ch=getchar();
return s*w;
}
inline void clr() {
for (ri i=0;i<=tot;++i) t[i].clear();
tot=1;
}
inline void Insert(char *s) {
int u=1,len=strlen(s);
for (ri i=0;i<len;++i) {
int v=(s[i]-'a');
if (!t[u].ch[v]) t[u].ch[v]=(++tot);
u=t[u].ch[v];
}
cnt[u]=len;
}
inline void GetFail() {
for (int i=0;i<26;++i) t[0].ch[i]=1;
t[1].fail=0;
queue<int> q;
q.push(1);
while (!q.empty()) {
int u=q.front(); q.pop();
for (ri i=0;i<26;++i) {
int v=t[u].ch[i],to=t[u].fail;
if (!v) {t[u].ch[i]=t[to].ch[i]; continue;}
t[v].fail=t[to].ch[i];
q.push(v);
}
}
}
inline int Query(char *s) {
int u=1,ans=0,len=strlen(s);
for (ri i=0;i<len;++i) {
int now=s[i]-'a',v=t[u].ch[now];
for (ri j=v;j>1;j=t[j].fail) {
if (book[i-cnt[j]]) {
book[i]=1;
ans=max(ans,i);
break;
}
}
u=t[u].ch[now];
}
return ans;
}
int main() {
n=read(); m=read();
for (ri i=1;i<=n;++i) {
scanf("%s",s);
Insert(s);
}
GetFail();
for (ri i=1;i<=m;++i) {
scanf("%s",s+1);
memset(book,0,sizeof(book));
book[0]=1;
printf("%d\n",Query(s));
}
return 0;
}

浙公网安备 33010602011771号