CodeForces 204E Little Elephant and Strings
http://codeforces.com/contest/204/problem/E
输入n个字符串,由小写字母组成。
对于每一个字符串,求出这个字符串,有多少个区间[L,R],使得子串L..R在至少K个字符串中出现,包括它本身。
SAM竟有如此美妙的构造方法。
首先用一个Trie来构造SAM。然后考虑par树,我们要求的是每个结点所代表的子串,包含了多少个字符串。
如果沿着Trie上的结点对应的SAM结点来访问SAM的话,可以发现对于每个结点它以及它沿着par树上的结点都在Trie上出现过一次,将这个结点的访问次数加一,那么对par树用一遍dfs使cnt[u]+=cnt[v]似乎就可以求出所有结点的访问次数了。但是这样会有一个重复计数,两个结点的公共祖先以上的结点重复计数过了,因此把它们的公共结点访问次数减一就对了。用Tarjan做一遍离线的LCA可以解决,记录Trie对应的每个串所包含的点的当前的公共祖先,如果有一个新的结点出现,就处理旧的公共祖先与当前点的公共祖先。
最后保留下所有出现次数>=k的结点。统计这个结点的len[u]-len[par[u]]即包含的子串数。我们通过Trie知道了某个结点是否是某个串的前缀,那么这个结点的后缀也应该在这个串上,统计par树上的父结点到当前结点的cnt,累加到该串的答案中。
1 #include <iostream> 2 #include <cstring> 3 #include <cstdio> 4 #include <vector> 5 #include <queue> 6 7 using namespace std; 8 typedef long long LL; 9 typedef vector<int> VI; 10 typedef vector<int>::iterator Iter; 11 const int maxn=210000; 12 const int MAXS=26; 13 char buf[maxn]; 14 LL ans[maxn]; 15 int TtoM[maxn],MtoT[maxn]; 16 int n,K; 17 int cnt[maxn]; 18 namespace Trie{ 19 int trans[maxn][26]; 20 VI spb[maxn]; 21 int tot; 22 int NewNode(){return tot++;} 23 void Insert(char s[],int id){ 24 int u=0; 25 for (int i=0;s[i];i++){ 26 int c=s[i]-'a'; 27 if (!trans[u][c]) trans[u][c]=NewNode(); 28 u=trans[u][c]; 29 spb[u].push_back(id); 30 } 31 } 32 void Init(){ 33 tot=0; 34 NewNode(); 35 } 36 } 37 38 namespace SAM{ 39 int trans[maxn][26],par[maxn],len[maxn],tot; 40 VI adj[maxn]; 41 inline int NewNode(){ 42 memset(trans[tot],0,sizeof(trans[tot])); 43 return tot++; 44 } 45 inline int NewNode(int u){ 46 memcpy(trans[tot],trans[u],sizeof(trans[u])); 47 par[tot]=par[u]; 48 return tot++; 49 } 50 inline int h(int u){ 51 return len[u]-len[par[u]]; 52 } 53 int Extend(int last,int w){ 54 int p=last,np=NewNode(); 55 len[np]=len[p]+1; 56 while (p&&!trans[p][w]){ 57 trans[p][w]=np; 58 p=par[p]; 59 } 60 if (!p&&!trans[p][w]) { 61 trans[p][w]=np; 62 par[np]=0; 63 } 64 else{ 65 int q=trans[p][w]; 66 if (len[p]+1==len[q]) par[np]=q; 67 else{ 68 int nq=NewNode(q); 69 len[nq]=len[p]+1; 70 par[q]=par[np]=nq; 71 while (p&&trans[p][w]==q){ 72 trans[p][w]=nq; 73 p=par[p]; 74 } 75 if (!p&&trans[p][w]==q) trans[p][w]=nq; 76 } 77 } 78 return np; 79 } 80 void Gao(){ 81 tot=0; 82 queue<int>que; 83 que.push(0); 84 TtoM[0]=NewNode(); 85 while (que.size()){ 86 int u=que.front(); 87 que.pop(); 88 for (int w=0;w<26;w++){ 89 int v=Trie::trans[u][w]; 90 if (v){ 91 TtoM[v]=Extend(TtoM[u],w); 92 MtoT[TtoM[v]]=v; 93 que.push(v); 94 } 95 } 96 } 97 for (int i=1;i<tot;i++) adj[par[i]].push_back(i); 98 } 99 void Check(){ 100 for (int u=1;u<tot;u++){ 101 cnt[u]=cnt[u]>=K?h(u):0; 102 } 103 } 104 } 105 106 namespace Tarjan{ 107 int pa[maxn],r[maxn],a[maxn],l[maxn]; 108 void Make(int x){ 109 r[x] = 0; 110 pa[x] = x; 111 a[x] = x; 112 } 113 int Find(int x){ 114 if (x!=pa[x]) pa[x]=Find(pa[x]); 115 return pa[x]; 116 } 117 void Union(int _x, int _y){ 118 int x=Find(_x); 119 int y=Find(_y); 120 if (r[x]<r[y]) pa[y]=x; 121 else{ 122 pa[x]=y; 123 a[y]=_x; 124 if (r[x]==r[y]) r[x]++; 125 } 126 } 127 void Tarjan(int u=0){ 128 Make(u); 129 for (Iter it=SAM::adj[u].begin();it!=SAM::adj[u].end();it++){ 130 Tarjan(*it); 131 Union(u,*it); 132 } 133 if (MtoT[u]){ 134 for (Iter it=Trie::spb[MtoT[u]].begin();it!=Trie::spb[MtoT[u]].end();it++){ 135 int v=*it; 136 --cnt[a[Find(l[v])]]; 137 ++cnt[u]; 138 l[v]=u; 139 140 } 141 } 142 } 143 void dfs1(int u=0){ 144 int sz=SAM::adj[u].size(); 145 for (int i=0;i<sz;i++){ 146 int v=SAM::adj[u][i]; 147 dfs1(v); 148 cnt[u]+=cnt[v]; 149 } 150 } 151 void dfs2(int u=0){ 152 int sz=SAM::adj[u].size(); 153 for (int i=0;i<sz;i++){ 154 int v=SAM::adj[u][i]; 155 cnt[v]+=cnt[u]; 156 dfs2(v); 157 } 158 sz=Trie::spb[MtoT[u]].size(); 159 for (int i=0;i<sz;i++){ 160 int v=Trie::spb[MtoT[u]][i]; 161 ans[v]+=cnt[u]; 162 } 163 } 164 void Init(){ 165 memset(l,0,sizeof(l)); 166 } 167 } 168 int main() 169 { 170 scanf("%d%d",&n,&K); 171 Trie::Init(); 172 for (int i=0;i<n;i++){ 173 scanf("%s",buf); 174 Trie::Insert(buf,i); 175 } 176 SAM::Gao(); 177 Tarjan::Init(); 178 Tarjan::Tarjan(); 179 Tarjan::dfs1(); 180 SAM::Check(); 181 Tarjan::dfs2(); 182 for (int i=0;i<n;i++) printf("%I64d ",ans[i]);puts(""); 183 return 0; 184 }