[Str记录]CF932G Palindrome Partition

· · 个人记录

题意 : 给定串 s ,将其划分为若干字符串 s_1s_2s_3...s_k

求满足 2|k,\ s_1=s_k,s_2=s_{k-1}...s_{i}=s_{k-i+1} 的方案数。

------------ 旧文分档。配合 [回文自动机小记](https://www.luogu.com.cn/blog/command-block/hui-wen-zi-dong-ji-xiao-ji) 食用。 - **结论①** : 一个回文串的回文后缀与其 $\rm Bd$ 一一对应。 **推论** : 串 $s$ 的回文后(前)缀的长度可以被划分为 $O(\log |s|)$ 个等差序列。 - **结论②** : 对于两个回文串 $u,v$ ,$v$ 是 $u$ 的最长回文严格后(前)缀,且 $2|v|>u$ ,则 $v$ 在 $u$ 中只会匹配两次,分别为前缀和后缀。 结论的证明请见上面拉链接的文章。 回到本题,条件需要判断串相等,并不方便。 考虑将串对折,使得需要相等的每一个小串“重叠”。 构造 $S'=S[n]S[1]S[n-1]S[2]S[n-2]...$ ,不难发现问题等价于把 $S'$ 划分成若干偶回文串。 设 $f[i]$ 为把 $S'[1,i]$ 划分的方案数,显然有如下转移 : $f[n]=\sum\limits_{i=0}^n\big[S'[i+1,n]\text{是偶回文串}\big]f[i]

我们按照 i=1\sim n 的顺序转移 f[i] ,那么每次需要查看的就是在 i 结尾的回文串。

可以在 \rm PAM 上直接跳 fail 暴力转移,复杂度是 O(n^2) ,无法通过。

根据 结论①\rm PAM 上一条链的 len 可以划分成 O(\log n) 个等差序列。

对于一个等差序列,如果中间某两项不满足 结论② 则将其强行断开,不难发现,等差数列的个数仍然是 O(\log n)

这样,每条等差序列相邻两个串之间都满足 结论② ,便于我们确定匹配位置。

对各个等差序列分别转移。

如图,绿红蓝为同一等差序列中的三个串,绿色虚线表示绿色串上一次出现的位置,蓝色虚线同理。

设当前处理 f[i],且上图所示等差序列 A 的公差为 d

(该等差链中)能转移到 f[i-d] 的位置由空心紫色圆标出,能转移到 f[i] 的位置由实心紫色圆标出。

红色串的出现位置我们还不能确知。但能够肯定的是,其一定不会给 f[i-d] 贡献,否则可以得到更长的能加入该等差序列的回文串,与红串最长矛盾。

不难发现,A\rightarrow f[i-d]A\rightarrow f[i] 的贡献只差了一个位置(即等差链顶),单独加上就好了。注意在本题中,需要特判是否为回文串。

需要给每条回文链记录上一次转移的贡献,注意这不是链剖分,各个等差序列可能重叠,所以需要把这个信息记录在等差序列的最深处。

具体如何处理等差链请见代码。

复杂度 O(n\log n)

#include<algorithm>
#include<cstring>
#include<cstdio>
#define MaxN 1000500
using namespace std;
const int mod=1000000007;
struct Node
{int t[26],len,f,tf,d,s;}a[MaxN];
int tn,las;
void linkd(int u)
{
  int fa=a[u].f;
  a[u].d=a[u].len-a[fa].len;
  a[u].tf=(a[u].d==a[fa].d&&a[u].d*2<=a[u].len) ? a[fa].tf : u;
}
void ins(int k,int c,char *str)
{
  int p=las;
  while(str[k-a[p].len-1]!=c)p=a[p].f;
  if (!a[p].t[c]){
    int np=las=++tn,v;
    for (v=a[p].f;str[k-a[v].len-1]!=c;v=a[v].f);
    if (!a[v].t[c])a[np].f=2;
    else a[np].f=a[v].t[c];
    a[a[p].t[c]=np].len=a[p].len+2;
    linkd(np);
  }else las=a[p].t[c];
}
int f[MaxN];
void dp(int k)
{
  for (int p=las;p>2;p=a[a[p].tf].f){
    int tl=a[a[p].tf].len;
    a[p].s=f[k-tl];
    if (a[p].tf!=p)a[p].s=(a[a[p].f].s+a[p].s)%mod;
    if (!(k&1))f[k]=(f[k]+a[p].s)%mod;
  }
}
void Init()
{a[1].f=a[2].f=1;a[1].len=-1;las=tn=2;}
int n;
char s[MaxN],s2[MaxN];
int main()
{
  scanf("%s",s2+1);
  n=strlen(s2+1);
  if (n&1){puts("0");return 0;}
  for (int i=1;i+i<=n;i++){
    s[i*2-1]=s2[n-i+1];
    s[i*2]=s2[i];
  }s[0]=-1;f[0]=1;Init();
  for (int i=1;i<=n;i++){
    ins(i,s[i]-='a',s);
    dp(i);
  }printf("%d",f[n]);
  return 0;
}