题解:P16157 [ICPC 2016 NAIPC] K-Inversions

· · 题解

题意简述

给定只含字符 \texttt{A}\texttt{B} 的字符串。 对每个 1\le k<n, 求满足左侧字符为 \texttt{B}、右侧字符为 \texttt{A}, 且两个位置相距 k 的下标对数量。

解题思路

把字符串下标改成从 0 开始。 需要计算:

\operatorname{ans}_k =\sum_{i=0}^{n-k-1}[s_i=\texttt{B}][s_{i+k}=\texttt{A}]

这是两个 0/1 序列按下标差进行匹配, 可以把其中一个序列翻转,将差转化为卷积中的下标和。

构造多项式:

F(x)=\sum_{\substack{0\le i<n\\s_i=\texttt{B}}}x^i

再把 \texttt{A} 的位置翻转:

G(x)=\sum_{\substack{0\le j<n\\s_j=\texttt{A}}}x^{n-1-j}

一对位置 i<j 在乘积中贡献的次数为:

i+n-1-j=n-1-(j-i)

因此,当 j-i=k 时, 这一对位置恰好给 F(x)G(x)n-1-k 次项贡献 1。 反过来,只有满足同样下标差的 \texttt{B}\texttt{A} 位置对 才会贡献到该次数。 所以:

\operatorname{ans}_k=[x^{n-1-k}]F(x)G(x)

使用 NTT 求出两个多项式的卷积。 最高次数为 2n-2, 变换长度取不小于 2n-1 的最小二次幂。 由于 n\le10^6,最大长度为 2^{21}, 模数 998244353 支持这一长度的 NTT,原根可取 3

对固定的 k,答案不超过 n-k, 而 n<998244353。 所以卷积系数在模数下的值就是实际整数答案,不会发生信息损失。

时间复杂度为 O(n\log n),空间复杂度为 O(n)

参考代码

#include <bits/stdc++.h>
using namespace std;

using ll=long long;
const int N=1<<21;
const int mod=998244353;
const int g=3;
int a[N],b[N];
ll Pow(ll x,ll y)
{
    x%=mod;
    ll res=1;
    while(y)
    {
        if(y&1)res=res*x%mod;
        x=x*x%mod;
        y>>=1;
    }
    return res;
}
void ntt(int a[],int n,bool op)
{
    for(int i=1,j=0;i<n;i++)
    {
        int k=n>>1;
        while(j>=k)
        {
            j-=k;
            k>>=1;
        }
        j+=k;
        if(i<j)swap(a[i],a[j]);
    }
    for(int i=2;i<=n;i<<=1)
    {
        int w=Pow(g,(mod-1)/i);
        if(!op)w=Pow(w,mod-2);
        for(int j=0;j<n;j+=i)
        {
            ll x=1;
            for(int k=0;k<i/2;k++)
            {
                int u=a[j+k];
                int v=int(x*a[j+k+i/2]%mod);
                a[j+k]=u+v<mod?u+v:u+v-mod;
                a[j+k+i/2]=u-v<0?u-v+mod:u-v;
                x=x*w%mod;
            }
        }
    }
    if(!op)
    {
        int inv=Pow(n,mod-2);
        for(int i=0;i<n;i++)a[i]=int((ll)a[i]*inv%mod);
    }
}
int main()
{
    ios::sync_with_stdio(false);
    cin.tie(nullptr);
    string s;
    cin>>s;
    int n=s.size();
    for(int i=0;i<n;i++)
    {
        if(s[i]=='B')a[i]=1;
        else b[n-1-i]=1;
    }
    int len=1;
    while(len<n+n-1)len<<=1;
    ntt(a,len,1);
    ntt(b,len,1);
    for(int i=0;i<len;i++)a[i]=int((ll)a[i]*b[i]%mod);
    ntt(a,len,0);
    for(int i=1;i<n;i++)cout<<a[n-1-i]<<'\n';
    return 0;
}