题解:CF2201E ABBA Counting

· · 题解

观察 1:设 S 的前半和后半为 P,Q,则 S 能被表示为 A+B+B+A 等价于 QP 的循环移位。

发现直接枚举位移计算会带上一个最小整周期的权,于是考虑枚举最小整周期并把每个剩余系合成一个考虑,最后对最小整周期做一点容斥即可。

变为给定两个串计算每个循环移位匹配的方案数之和,需要对每个移位求出 ab?? 交的次数,使用 ntt 优化字符串匹配的 trick 即可。

大概就是给每种字符设计一个权值,再设计一个二元式子求和如 \sum ab(a-b)^2 表示匹配,最后用 ntt 快速计算所有移位的结果。

时间复杂度一个松的界是 O(n\sqrt n+n\log n\log\log n),其中后面一坨带有 ntt 大常数。

:::success[点击查看参考代码]

#include<bits/stdc++.h>
#define TIME chrono::duration_cast<chrono::milliseconds>(chrono::high_resolution_clock::now().time_since_epoch()).count()
#define rep(i,l,r) for(int qwp=(r),i=(l);i<=qwp;i++)
#define per(i,r,l) for(int qwp=(l),i=(r);i>=qwp;i--)
using namespace std;
namespace c0dE1ng{
typedef long long ll;
constexpr ll mod=998244353;
constexpr ll G=3,IG=332748118;
constexpr int N=4e5+5;
inline int frint(){int n=0;char c=getchar();while(!isdigit(c))c=getchar();while(isdigit(c))n=n*10+c-48,c=getchar();return n;}
inline int frstr(char *s){int n=0;char c=getchar();while(c!='a'&&c!='b'&&c!='?')c=getchar();while(c=='a'||c=='b'||c=='?')s[++n]=c,c=getchar();return n;}
void wrll(ll x){if(x>9)wrll(x/10);putchar(x%10+48);}
inline void inc(ll &x,ll y){(x+=y)>=mod&&(x-=mod);}
inline void dec(ll &x,ll y){(x-=y)<0&&(x+=mod);}
ll pw2[N];inline void Init(int n){pw2[0]=1;rep(i,1,n)pw2[i]=pw2[i-1]*2%mod;}
inline int F(char x){return x=='?'?0:1<<x-'a';}
inline ll ksm(ll a,ll b,ll p){a%=p;ll r=1;while(b){if(b&1)r=r*a%p;a=a*a%p,b>>=1;}return r%p;}
inline ll inv(ll x){return ksm(x,mod-2,mod);}
int to[N<<2];
void ntt(int g,ll *a,int m){
    rep(i,0,g-1)if(i<(to[i]=(to[i>>1]>>1)|((g>>1)*(i&1))))swap(a[i],a[to[i]]);
    for(int k=1;k<g;k<<=1){
        ll W=ksm(m==1?G:IG,(mod-1)/(k<<1),mod);for(int i=0;i<g;i+=k<<1){
            ll w=1;rep(j,0,k-1){
                const ll p=a[i|j],q=a[i|j|k]*w%mod;
                a[i|j]=(p+q)%mod,a[i|j|k]=(p+mod-q)%mod,w=w*W%mod;
            }
        }
    }
    if(m==-1){ll w=inv(g);rep(i,0,g-1)a[i]=a[i]*w%mod;}
}
ll A[N<<2],B[N<<2],C[N<<2];
inline void calc(int n){
    int g=1;while(g<=(n<<1))g<<=1;
    rep(i,0,g-1)C[i]=0;
    ntt(g,A,1),ntt(g,B,1);rep(i,0,g-1)C[i]=A[i]*B[i]%mod;ntt(g,C,-1);
    rep(i,0,g-1)A[i]=B[i]=0;
}
int n;char a[N];ll f[N];
int x[N],y[N];ll p[N],q[N];
inline ll cal(int m){
    rep(i,1,m){
        x[i]=y[i]=0;for(int j=i;j<=n;j+=m)x[i]|=F(a[j]),y[i]|=F(a[j+n]);
        if(x[i]==3||y[i]==3)return 0;
    }copy_n(y+1,m,y+1+m);
    rep(i,0,m*2)p[i]=q[i]=0;
    rep(i,0,m*2)A[i]=1<=i&&i<=m?x[i]*x[i]*x[i]:0,B[i]=1<=i&&i<=m*2?y[i]:0;
    reverse(A,A+m*2+1),calc(m*2);rep(s,0,m-1)inc(p[s],C[m*2+s]);
    rep(i,0,m*2)A[i]=1<=i&&i<=m?x[i]*x[i]:0,B[i]=1<=i&&i<=m*2?y[i]*y[i]:0;
    reverse(A,A+m*2+1),calc(m*2);rep(s,0,m-1)dec(p[s],C[m*2+s]*2);
    rep(i,0,m*2)A[i]=1<=i&&i<=m?x[i]:0,B[i]=1<=i&&i<=m*2?y[i]*y[i]*y[i]:0;
    reverse(A,A+m*2+1),calc(m*2);rep(s,0,m-1)inc(p[s],C[m*2+s]);
    rep(i,0,m*2)A[i]=1<=i&&i<=m?!x[i]:0,B[i]=1<=i&&i<=m*2?!y[i]:0;
    reverse(A,A+m*2+1),calc(m*2);rep(s,0,m-1)inc(q[s],C[m*2+s]);
    ll res=0;rep(s,0,m-1)if(!p[s])inc(res,pw2[q[s]]);return res;
}
void slv(){
    frint(),n=frstr(a)>>1;rep(i,1,n)f[i]=0;
    rep(i,1,n)if(!(n%i)){inc(f[i],cal(i));rep(j,2,n/i)if(!(n%(i*j)))dec(f[i*j],f[i]*j%mod);}
    ll ans=0;rep(i,1,n)inc(ans,f[i]);wrll(ans),putchar('\n');
}
void main(){
    Init(N-5);
    int T;T=frint();
    while(T--)slv();
}
}
int main(){
    auto _Tbe=TIME;
    c0dE1ng::main();
    auto _Ted=TIME;
    cerr<<"\nTIME:"<<_Ted-_Tbe<<'\n';
    return 0;
}
/*
ulimit -s 1048576
g++ -O2 -std=c++14 -static A.cpp -o A.exe;size A.exe;./A.exe < A.in > A.out
*/

:::