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\)最大后缀。

最大就保证了最优。

具体实现的细节还是挺多的:

  1. 根节点的Fail是什么?为了方便后面的实现,我们建立一个虚拟节点0,并让0的所有儿子指向根,根的Fail指针指向0。具体原因看细节2。
  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;
}

参考资料

AC自动机 by hyfhaha

posted @ 2021-11-11 14:50  Carlotta24  阅读(90)  评论(0)    收藏  举报