多项式快速插值题解+学习笔记

多项式快速插值题解+学习笔记

知识目录

  1. 拉格朗日插值:模板题链接

  2. 多项式乘法(FFT):模板题链接

  3. 多项式乘法(NTT):模板题链接

  4. 多项式乘法逆:模板题链接

  5. 多项式除法&多项式取模:模板题链接

  6. 多项式多点求值:模板题链接

  7. 多项式快速插值:模板题链接

拉格朗日插值讲解

洛谷题目题意和数据范围

对于 \(n\) 个点 \((x_i,y_i)\),如果满足 \(\forall i\neq j, x_i\neq x_j\),那么经过这 \(n\) 个点可以唯一地确定一个 \(n-1\) 次多项式 \(y = f(x)\)

现在,给定这样 \(n\) 个点,请你确定这个 \(n-1\) 次多项式,并求出 \(f(k) \bmod 998244353\) 的值。

\(1 \le n \leq 2\times 10^3\)\(1 \le x_i,y_i,k < 998244353\)\(x_i\) 两两不同。

题目思路

首先,我们可以发现,如果我们定义: \(f(x)=\sum_{i=1}^{n} y_i\prod_{j\neq i}^{n} \frac{x-x_j}{x_i-x_j}\)。那么,这个函数可以发现是能拟合所有的 \(n\) 个点。具体的,如果将 \(x_k\) 带入,则推导:

\(f(x_k)=\sum_{i=1}^{n} y_i\prod_{j\neq i}^{n} \frac{x_k-x_j}{x_i-x_j}\)

\(f(x_k)=\sum_{i\neq k}^{} (y_i\cdot \frac{x_k-x_k}{x_i-x_k} \prod_{j\neq i;j\neq k}^{n} \frac{x_k-x_j}{x_i-x_j})+y_k\prod_{j\neq k}^{n} \frac{x_k-x_j}{x_k-x_j}\)

\(f(x_k)=\sum_{i\neq k}^{} (y_i\cdot 0 \prod_{j\neq i;j\neq k}^{n} \frac{x_k-x_j}{x_i-x_j})+y_k\prod_{j\neq k}^{n} 1=y_k\)

那么,我们可以发现,该函数是可行的。接下来,我们要证明这个函数是唯一符合要求的函数,证明过程如下:

假设有 \(g(x)\) 不等于 \(f(x)\) 也同样符合要求。那么,我们设 \(h(x)=g(x)-f(x)\),则对于每个 \(x_i\),其都是 \(h(x)\) 的零点,且其两两不同。根据代数基本定理,对于一个 \(n-1\) 次函数,其最多有 \(n\) 个不重复的零点。所以 \(h(x)=0\) 或其为更高阶的函数,显然其不符合 \(g(x)\) 定义,所以证明 \(f(x)\) 唯一。

这样的话,我们就证明了此方法的正确性和唯一性,而观察到其时间复杂度为 \(n^2\) 可过,所以解决此题。

时间复杂度 \(O(n^2)\) .

题目代码

#include<bits/stdc++.h>
using namespace std;
long long n,m,op=0,os=1;
const long long mod=998244353;
struct one{
    long long x,y;
}a[2005];
inline long long pow2(long long a1,long long b1){
    long long c1=1;
    a1%=mod;
    while(b1!=0){
        if(b1%2==1){
            c1=c1*a1%mod;
        }
        a1=a1*a1%mod;
        b1/=2;
    }
    return c1;
}
int main(){
	cin>>n>>m;
    for(int i=1;i<=n;i++){
        cin>>a[i].x>>a[i].y;
    }
    for(int i=1;i<=n;i++){
        os=a[i].y;
        for(int j=1;j<=n;j++){
            if(i!=j){
                os=os*((m-a[j].x+mod)*pow2(a[i].x-a[j].x+mod,mod-2)%mod)%mod;   
            }
        }
        op=(op+os)%mod;
    }
    cout<<op<<endl;
	return 0;
} 

多项式乘法(FFT)

洛谷题目题意和数据范围

给定一个 \(n\) 次多项式 \(F(x)\),和一个 \(m\) 次多项式 \(G(x)\)

请求出 \(F(x)\)\(G(x)\) 的乘积。

保证输入中的系数均为大于等于 \(0\) 且小于等于 \(9\) 的整数。

对于 \(100\%\) 的数据:\(0 \le n, m \leq {10}^6\)

题目思路

首先,我们发现,直接暴力显然是不行的,那么考虑别的方法。

可以发现,一个 \(n\) 次的多项式是等同于 \(n+1\) 个不同的点的,这个可以从拉格朗日插值法那部分找到证明。所以,我们就考虑将 \(F(x)\) 转化为多个点值计算。但是我们该转化为什么样的点呢?那么先了解一下下面的知识,就会知道如何转化了。

复数

首先,学过高中数学的会知道,什么是复数。那么,现在就讲一下什么是复数:

相信没学过复数的人都会有个疑问,\(x^2=-1\) 的解究竟是多少呢?

那么现在我们定义 \(i^2=-1\) 其中 \(i\) 就是虚数的单位,可以发现,现在此问题就有了 \(i\)\(-i\) 两个解了。

而形如 \(a+bi\) 的数就是复数了,其中 \(a,b\) 是实数,我们称。现在我们讲一下复数的运算。

对于 \(z=a+bi\)\(o=c+di\) 的和,可以发现 \(z+o=a+bi+c+di=(a+c)+(b+d)i\) 也就是其实数部分和虚数部分之和。

而对于 \(z=a+bi\)\(o=c+di\) 的差,那么就是 \(z-o=a+bi-(c+di)=(a-c)+(b-d)i\)

对于 \(z=a+bi\)\(o=c+di\) 的积:

\(z*o=(a+bi)*(c+di)=ac+adi+cbi+bdi^2=ac+adi+cbi-bd=(ac-bd)+(ad+cb)i\)

而对于 \(z=a+bi\)\(o=c+di\) 的商:

\(\frac{z}{o}=\frac{a+bi}{c+di}=\frac{(a+bi)*(c-di)}{(c+di)*(c-di)}=\frac{ac-adi+bci+bd}{c^2-d^2}=\frac{ac+bd}{c^2-d^2}+\frac{bc-ad}{c^2-d^2}i\)

向量

接下来,我们来讲一下向量。

向量,就是一个从原点 \((0,0)\)\((a,b)\) 这一点的有向的量。对于从点 \(A\) 到点 \(B\) 的向量,称其为 \(AB\)

向量也有加减法运算,其为 \(AB+BC=AC\)\(AC-BC=AB\) 这是非常好理解的。可以理解为你从 \(A\) 到了 \(B\),然后从 \(B\) 再到 \(C\) ,就等于你从 \(A\) 到了 \(C\),减法同理。

然后,我们来理解向量和复数的关系,对于一个复数域,其可以画在一个平面上。类似平面直角坐标系,其虚数轴为 \(y\) 轴,实数轴为 \(x\) 轴。那么复数也可以理解为等同于一个向量 \((a,b)\)

弧度制

弧度制即是把 \(360\) 度的角转化为了 \(2\pi\) ,也就是说 \(180^\circ =\pi rad\)

单位根

单位根指 \(z^n=1\) 的解集。

对于 \(z^n=1\) 这个式子,我们可以画一个圆在复平面上,其圆心为 \((0,0)\),半径为 \(1\)。那么,我们可以发现显然 \(z^n=1\) 的解都在圆上,这是因为对于任何一个复数,均可表示为 \(r(cos \theta +isin\theta )\) 其中 \(r\) 称为其模长。这个显然可以画图使用三角函数证明。而我们发现 \(z^n=1\) 显然其解的 \(r\) 定是等于 \(1\) 的,这样就证明了。

那么,我们设 \(w_{n}^{1}\) 为该式的第一个解,其满足正且与 \(x\) 轴的角度最小。然后也依次有 \(n\) 个解,为 \(w_{n}^{k}\),其中 \(w_{n}^{0}=w_{n}^{n}=1\)

我们可以发现其为 \(w_{n}^{k}=cos \frac{2k\pi}{n}+isin \frac{2k\pi}{n}=e^{i\frac{2k\pi }{n}}\),这个可以根据欧拉公式和棣莫弗定理证明,具体的可以上网搜资料,就不在这讲了。

我们接下来默认 \(n\) 为2的正整数次幂。

可以发现 \(w_{2n}^{2k}=cos \frac{4k\pi}{2n}+isin \frac{4k\pi}{2n}=cos \frac{2k\pi}{n}+isin \frac{2k\pi}{n}=w_{n}^{k}\)

还有 \(w_{n}^{k+\frac{n}{2}}=cos(\frac{2k\pi}{n}+\pi)+isin(\frac{2k\pi}{n}+\pi)=-w_{n}^{k}\)

\((w_{n}^{k})^2=cos^2(\frac{2k\pi}{n})-sin^2(\frac{2k\pi}{n})+2cos(\frac{2k\pi}{n})sin(\frac{2k\pi}{n})i=cos(\frac{4k\pi}{n})+isin(\frac{4k\pi}{n})=w_{n}^{2k}\),这里用到了三角函数的一些相关公式。

正题

我们默认 \(n\) 为2的正整数次幂。

重回正题,我们思考如何快速计算对于每个单位根 \(F(x)\) 值。

首先,设 \(F(x)=\sum_{j=0}^{n-1} a_jx^j\),则将其下表按奇偶分类:

\(G(x)=\sum_{j=0}^{\frac{n}{2}-1} a_{2j}x^{j}\)

\(H(x)=\sum_{j=0}^{\frac{n}{2}-1} a_{2j+1}x^{j}\)

容易发现 \(F(x)=G(x^2)+xH(x^2)\),这样的话,带入可得(\(k\le \frac{n}{2}\)):

\(F(w_{n}^{k})=G(w_{n}^{2k})+w_{n}^{k}H(w_{n}^{2k})=G(w_{\frac{n}{2}}^{k})+w_{n}^{k}H(w_{\frac{n}{2}}^{k})\)

\(F(w_{n}^{k+\frac{n}{2}})=G(w_{n}^{2k+n})-w_{n}^{k}H(w_{n}^{2k+n})=G(w_{\frac{n}{2}}^{k})-w_{n}^{k}H(w_{\frac{n}{2}}^{k})\)

这个方法称为 \(DFT\)

这里可以发现,其时间复杂度为 \(T(n)=2T(\frac{n}{2})+n\),可以发现有 \(log n\)层而每层为 \(n\) 个数,所以时间复杂度为 \(T(n)=O(nlog n)\)

然后,我们现在有了两个多项式的每个单位根数值,该如何求呢?

首先,我们可以发现结果多项式 \(H(w_{n}^{k})=F(w_{n}^{k})*G(w_{n}^{k})\)。那么,我们只需要考虑如何还原多项式即可。

考虑一个这样的矩阵:

\[Z= \begin{bmatrix} (w_{n}^{0})^0 & (w_{n}^{0})^1 & (w_{n}^{0})^2 & ... & (w_{n}^{0})^{n-2} & (w_{n}^{0})^{n-1}\\ (w_{n}^{1})^0 & (w_{n}^{1})^1 &...&...&...& (w_{n}^{1})^{n-1}\\ ...&...&...&...&...&...\\ ...&...&...&...&...&...\\ (w_{n}^{n-2})^0 & (w_{n}^{n-2})^1 &...&...&...& (w_{n}^{n-2})^{n-1}\\ (w_{n}^{n-1})^0 & (w_{n}^{n-1})^1 & (w_{n}^{n-1})^2 & ...... & (w_{n}^{n-1})^{n-2} & (w_{n}^{n-1})^{n-1}\\ \end{bmatrix} \]

考虑这个向量,其中 \(h_i\)\(H\) 的系数:

\[X= \begin{bmatrix} h_0\\ h_1\\ h_2\\ ...\\ h_{n-2}\\ h_{n-1}\\ \end{bmatrix} \]

考虑这个向量:

\[Y= \begin{bmatrix} H(w_{n}^{0})\\ H(w_{n}^{1})\\ H(w_{n}^{2})\\ ...\\ H(w_{n}^{n-2})\\ H(w_{n}^{n-1})\ \end{bmatrix} \]

则显然,\(ZX=Y\),我们要解 \(X\) 的值。则 \(X=Z^{-1}Y\)

可以发现,\(Z^{-1}\)\(Z\) 每位取倒数后乘上 \(\frac{1}{n}\) 的结果,证明:

\[Z^{-1}= \begin{bmatrix} \frac{1}{n}(w_{n}^{0})^0 & \frac{1}{n}(w_{n}^{0})^{-1} & \frac{1}{n}(w_{n}^{0})^{-2} & ... & \frac{1}{n}(w_{n}^{0})^{-(n-2)} & \frac{1}{n}(w_{n}^{0})^{-(n-1)}\\ \frac{1}{n}(w_{n}^{1})^0 & \frac{1}{n}(w_{n}^{1})^{-1} &...&...&...& \frac{1}{n}(w_{n}^{1})^{-(n-1)}\\ ...&...&...&...&...&...\\ ...&...&...&...&...&...\\ \frac{1}{n}(w_{n}^{n-2})^0 & \frac{1}{n}(w_{n}^{n-2})^{-1} &...&...&...& \frac{1}{n}(w_{n}^{n-2})^{-(n-1)}\\ \frac{1}{n}(w_{n}^{n-1})^0 & \frac{1}{n}(w_{n}^{n-1})^{-1} & \frac{1}{n}(w_{n}^{n-1})^{-2} & ...... & \frac{1}{n}(w_{n}^{n-1})^{-(n-2)} & \frac{1}{n} (w_{n}^{n-1})^{-(n-1)}\\ \end{bmatrix} \]

那么,\(Z^{-1}Y\) 的第 \(i\)\(h_i\) 为:

\(h_i=\sum_{j=0}^{n-1} \frac{1}{n}(w_{n}^{i})^{-j}H(w_{n}^{j})\)

\(ZX\) 的第 \(k\)\(H(w_{n}^{k})\) 为:

\(H(w_{n}^{k})=\sum_{i=0}^{n-1}(w_{n}^{k})^{i} \sum_{j=0}^{n-1} \frac{1}{n}(w_{n}^{i})^{-j}H(w_{n}^{j})\)

\(H(w_{n}^{k})=\sum_{i=0}^{n-1}(w_{n}^{i})^{k} \sum_{j=0}^{n-1} \frac{1}{n}(w_{n}^{i})^{-j}H(w_{n}^{j})\)

\(H(w_{n}^{k})=\frac{1}{n}\sum_{i=0}^{n-1} \sum_{j=0}^{n-1} (w_{n}^{i})^{k-j}H(w_{n}^{j})\)

\(H(w_{n}^{k})=\frac{1}{n}\sum_{j=0}^{n-1} \sum_{i=0}^{n-1} (w_{n}^{i})^{k-j}H(w_{n}^{j})\)

\(H(w_{n}^{k})=\frac{1}{n}\sum_{j=0}^{n-1} H(w_{n}^{j})\sum_{i=0}^{n-1} (w_{n}^{i})^{k-j}\)

\(H(w_{n}^{k})=\frac{1}{n}\sum_{j\neq k}^{} H(w_{n}^{j})\sum_{i=0}^{n-1} (w_{n}^{i})^{k-j}+\frac{1}{n}H(w_{n}^{k})\sum_{i=0}^{n-1} (w_{n}^{i})^{k-k}\)

\(H(w_{n}^{k})=\frac{1}{n}\sum_{j\neq k}^{} H(w_{n}^{j})\sum_{i=0}^{n-1} (w_{n}^{k-j})^{i}+\frac{1}{n}H(w_{n}^{k})\sum_{i=0}^{n-1} 1\)

\(H(w_{n}^{k})=\frac{1}{n}\sum_{j\neq k}^{} H(w_{n}^{j})\sum_{i=0}^{n-1} (w_{n}^{k-j})^{i}+H(w_{n}^{k})\)

接着,我们要证明:

\(\frac{1}{n}\sum_{j\neq k}^{} H(w_{n}^{j})\sum_{i=0}^{n-1} (w_{n}^{k-j})^{i}=0\)

\(\sum_{i=0}^{n-1} (w_{n}^{k-j})^{i}=0\) 时显然原式值为零,证明:

可以发现,当 \(k-j\) 为奇数时原式 \((w_{n}^{k-j})^{i}+(w_{n}^{k-j})^{i+\frac{n}{2}}=0\),因为 \(\frac{(k-j)\frac{n}{2}}{\frac{n}{2}}\%2=1\),这个是因为 \(w_{n}^{i}=w_{n}^{i+\frac{n}{2}}\) 的原因。

那么,对于 \(\frac{k-j}{2^p}\%2=1\) 的情况,我们只需将 \((w_{n}^{k-j})^{i},(w_{n}^{k-j})^{i+\frac{n}{2^p}}\) 放到一组抵消即可,得证。

然后,我们该如何求呢?

可以发现,我们的操作和 \(DFT\) 没有什么不同,只需把单位根取负结果就是原来的倒数,最后在乘上 \(\frac{1}{n}\) 即可。

对于 \(n\) 不符合 \(2^k\) 条件的,只需要在其前面加部分零使得其符合即可。

时间复杂度 \(O(nlogn)\) .

题目代码

#include<bits/stdc++.h>
using namespace std;
long long n,m,nn=1;
const double pi=acos(-1.0);
double c[4000006];
struct com{
    double x,y;
}a[4000006],b[4000006];
com operator*(com a1,com b1){
    return (com){a1.x*b1.x-a1.y*b1.y,a1.x*b1.y+a1.y*b1.x};
}
com operator/(com a1,com b1){
    return (com){(a1.x*b1.x+a1.y*b1.y)/(b1.x*b1.x+b1.y*b1.y),(a1.y*b1.x-a1.x*b1.y)/(b1.x*b1.x+b1.y*b1.y)};
}
com operator+(com a1,com b1){
    return (com){a1.x+b1.x,a1.y+b1.y};
}
com operator-(com a1,com b1){
    return (com){a1.x-b1.x,a1.y-b1.y};
}
void fft(long long n1,com *a1,long long b1){
    if(n1<=1){
        return ;
    }
    com ax[n1/2],ay[n1/2];
    for(int i=0;i<=n1-1;i+=2){
        ax[i/2]=a1[i];
        ay[i/2]=a1[i+1];
    }
    fft(n1/2,ax,b1);
    fft(n1/2,ay,b1);
    com ww=(com){cos(2.0*pi/n1),b1*sin(2.0*pi/n1)},w=(com){1.0,0.0};
    for(int i=0;i<=n1/2-1;i++,w=w*ww){
        a1[i]=ax[i]+w*ay[i];
        a1[i+n1/2]=ax[i]-w*ay[i];
    }
    return ;
}
int main(){
    cin>>n>>m;
    while(nn<=n+m){
        nn*=2;
    }
    for(int i=0;i<=n;i++){
        cin>>a[i].x;
    }
    for(int i=0;i<=m;i++){
        cin>>b[i].x;
    }
    while(nn<=n+m){
        nn*=2;
    }    
    fft(nn,a,1);
    fft(nn,b,1);
    for(int i=0;i<=nn;i++){
        a[i]=a[i]*b[i];
    }
    fft(nn,a,-1);
    for(int i=0;i<=nn;i++){
        c[i]=a[i].x/nn;
    }
    for(int i=0;i<=n+m;i++){
        cout<<(long long)(c[i]+0.5)<<" ";
    }
    cout<<endl;
	return 0;
}

多项式乘法(NTT)

洛谷题目题意和数据范围

给定一个 \(n\) 次多项式 \(F(x)\),和一个 \(m\) 次多项式 \(G(x)\)

请求出 \(F(x)\)\(G(x)\) 的乘积,系数对 \(998244353\) 取模。

保证输入中的系数均为大于等于 \(0\) 且小于等于 \(9\) 的整数。

对于 \(100\%\) 的数据:\(0 \le n, m \leq {10}^6\)

题目思路

这其实和 \(FFT\) 一样,只是 \(NTT\) 是能更快处理整数和进行取模的。

\(NTT\),即快速数论变换。这是在一个整数模质数的域内进行的,而当 \(n\)\(p-1\) 整除时,是存在本原 \(n\)次方根的,也就是可以使用的。所以有 \(p=qn+1\),然后原根满足 \(g^{qn}\%p=1\),也就是可以将 \(w_{n}^{i}\) 看成 \(g_{n}^{i}=g^{qi}\) 即可。注意:对于递归到某一层,若此层大小为 \(l\) 个数,则 \(g_{n}^{i}=g^{\frac{p-1}{l}i}\)

对于 \(998244353\),其原根为 \(3\)

时间复杂度 \(O(nlogn)\) .

题目代码

#include<bits/stdc++.h>
using namespace std;
long long n,m,nn=1,nn2=1,uv[4000006][2];
const long long mod=998244353;
long long c[4000006];
long long a[4000006],b[4000006];
long long pow2(long long a1,long long b1){
    long long c1=1;
    while(b1!=0){
        if(b1%2==1){
            c1=c1*a1%mod;
        }
        a1=a1*a1%mod;
        b1/=2;
    }
    return c1;
}
void ntt(long long n1,long long *a1,long long b1){
    if(n1<=1){
        return ;
    }
    long long ax[n1/2],ay[n1/2];
    for(int i=0;i<=n1-1;i+=2){
        ax[i/2]=a1[i];
        ay[i/2]=a1[i+1];
    }
    ntt(n1/2,ax,b1);
    ntt(n1/2,ay,b1);
    long long ww,w=1;
    if(b1==1){
        if(uv[n1][0]==0){
            ww=pow2(3ll,(mod-1)*pow2(n1,mod-2)%mod);
            uv[n1][0]=pow2(3ll,(mod-1)*pow2(n1,mod-2)%mod);
        }
        else{
            ww=uv[n1][0];
        }
    }
    else{
        if(uv[n1][1]==0){
            ww=pow2(pow2(3ll,(mod-1)*pow2(n1,mod-2)%mod),mod-2);
            uv[n1][1]=pow2(pow2(3ll,(mod-1)*pow2(n1,mod-2)%mod),mod-2);
        }
        else{
            ww=uv[n1][1];
        }
    }
    for(int i=0;i<=n1/2-1;i++,w=w*ww%mod){
        a1[i]=(ax[i]+w*ay[i])%mod;
        a1[i+n1/2]=(ax[i]-w*ay[i]%mod+mod)%mod;
    }
    return ;
}
int main(){
    cin>>n>>m;
    while(nn<=n+m){
        nn*=2;
    }
    for(int i=0;i<=n;i++){
        cin>>a[i];
    }
    for(int i=0;i<=m;i++){
        cin>>b[i];
    }
    while(nn<=n+m){
        nn*=2;
    }    
    ntt(nn,a,1);
    ntt(nn,b,1);
    for(int i=0;i<=nn;i++){
        a[i]=a[i]*b[i]%mod;
    }
    ntt(nn,a,-1);
    nn2=pow2(nn,mod-2);
    for(int i=0;i<=nn;i++){
        c[i]=a[i]*nn2%mod;
    }
    for(int i=0;i<=n+m;i++){
        cout<<c[i]<<" ";
    }
    cout<<endl;
	return 0;
}

多项式乘法逆

洛谷题目题意和数据范围

给定一个多项式 \(F(x)\) ,请求出一个多项式 \(G(x)\), 满足 \(F(x) \cdot G(x) \equiv 1 \pmod{x^n}\)。系数对 \(998244353\) 取模。

对于 \(100\%\) 的数据,\(1 \leq n \leq 10^5\)\(0 \leq a_i \leq 10^9\)

题目思路

这道题一看就是递归向下。

假设我们已经求出 \(f_0(x) \cdot f(x)\equiv 1 \pmod{x^{\lceil \frac{n}{2}\rceil}}\),那么:

\(f_0(x) \cdot f(x)-1\equiv 0 \pmod{x^{\lceil \frac{n}{2}\rceil}}\)

\((f_0(x)f(x)-1)^2\equiv 0 \pmod{x^n}\)

\(f_0(x)^2f(x)^2-2f_0(x)f(x)+1\equiv 0 \pmod{x^n}\)

\(f_0(x)^2f(x)^2-2f_0(x)f(x)\equiv -1 \pmod{x^n}\)

\(2f_0(x)f(x)-f_0(x)^2f(x)^2\equiv 1 \pmod{x^n}\)

\(f(x)(2f_0(x)-f_0(x)^2f(x))\equiv 1 \pmod{x^n}\)

\(f(x)(f_0(x)(2-f_0(x)f(x)))\equiv 1 \pmod{x^n}\)

容易发现,此时 \(g(x)=f_0(x)(2-f_0(x)f(x))\),所以我们可以发现,只需不断递归即可。具体的,从最低项直接暴力计算逆元,然后往上计算即可。时间复杂度:\(T(n)=T(n/2)+O(nlogn)=O(nlogn)\)

时间复杂度 \(O(nlogn)\)

题目代码

#include<bits/stdc++.h>
using namespace std;
long long n,m,nn=1,nn2=1,uv[4000006][2];
const long long mod=998244353;
long long c[4000006];
long long a[4000006],b[4000006],d[4000006];
long long pow2(long long a1,long long b1){
    long long c1=1;
    while(b1!=0){
        if(b1%2==1){
            c1=c1*a1%mod;
        }
        a1=a1*a1%mod;
        b1/=2;
    }
    return c1;
}
void ntt(long long n1,long long *a1,long long b1){
    if(n1<=1){
        return ;
    }
    long long ax[n1/2],ay[n1/2];
    for(int i=0;i<=n1-1;i+=2){
        ax[i/2]=a1[i];
        ay[i/2]=a1[i+1];
    }
    ntt(n1/2,ax,b1);
    ntt(n1/2,ay,b1);
    long long ww,w=1;
    if(b1==1){
        if(uv[n1][0]==0){
            ww=pow2(3ll,(mod-1)*pow2(n1,mod-2)%mod);
            uv[n1][0]=pow2(3ll,(mod-1)*pow2(n1,mod-2)%mod);
        }
        else{
            ww=uv[n1][0];
        }
    }
    else{
        if(uv[n1][1]==0){
            ww=pow2(pow2(3ll,(mod-1)*pow2(n1,mod-2)%mod),mod-2);
            uv[n1][1]=pow2(pow2(3ll,(mod-1)*pow2(n1,mod-2)%mod),mod-2);
        }
        else{
            ww=uv[n1][1];
        }
    }
    for(int i=0;i<=n1/2-1;i++,w=w*ww%mod){
        a1[i]=(ax[i]+w*ay[i])%mod;
        a1[i+n1/2]=(ax[i]-w*ay[i]%mod+mod)%mod;
    }
    return ;
}
int main(){
    cin>>n;
    for(int i=0;i<=n-1;i++){
        cin>>d[i];
    }
    a[0]=pow2(d[0],mod-2);
	while(nn<=n+n){
		nn*=2;
		nn2=pow2(nn*2ll%mod,mod-2);
		for(int i=0;i<=nn-1;i++){
			c[i]=b[i]=0;
		}
//		for(int i=0;i<=nn/2-1;i++){
//			c[i]=a[i];
//		}
		ntt(nn*2,a,1);
		for(int i=0;i<=nn/2-1;i++){
			c[i]=a[i];
		}
		for(int i=0;i<=nn-1;i++){
			b[i]=d[i];
		}
		ntt(nn*2,b,1);
		for(int i=0;i<=nn*2-1;i++){
			a[i]=a[i]*(2ll-b[i]*a[i]%mod+mod)%mod;
		}
		ntt(nn*2,a,-1);
		for(int i=0;i<=nn*2-1;i++){
			a[i]=a[i]*nn2%mod;
			if(i>=nn){
				a[i]=0;
			}
		}           
	} 
    for(int i=0;i<=n-1;i++){
        cout<<a[i]<<" ";
    }
    cout<<endl;
	return 0;
}

多项式除法&多项式取模

洛谷题目题意和数据范围

给定一个 \(n\) 次多项式 \(F(x)\) 和一个 \(m\) 次多项式 \(G(x)\) ,请求出多项式 \(Q(x)\), \(R(x)\),满足以下条件:

  • \(Q(x)\) 次数为 \(n-m\)\(R(x)\) 次数小于 \(m\)
  • \(F(x) = Q(x) * G(x) + R(x)\)

所有的运算在模 \(998244353\) 意义下进行。

对于所有数据,\(1 \le m < n \le 10^5\),给出的系数均属于 \([0, 998244353) \cap \mathbb{Z}\)

题目思路

我们先设:

\(f(x)=F(\frac{1}{x})\)

\(g(x)=G(\frac{1}{x})\)

\(q(x)=Q(\frac{1}{x})\)

\(r(x)=R(\frac{1}{x})\)

\(f_0(x)=x^nF(\frac{1}{x})\)

\(g_0(x)=x^mG(\frac{1}{x})\)

\(q_0(x)=x^{n-m}Q(\frac{1}{x})\)

\(r_0(x)=x^{m-1}R(\frac{1}{x})\)

显然,\(F(\frac{1}{x})=Q(\frac{1}{x})*G(\frac{1}{x})+R(\frac{1}{x})\),则:

\(f(x)=q(x)*g(x)+r(x)\)

\(x^nf(x)=x^mq(x)*x^{n-m}g(x)+x^{n-m+1}x^{m-1}r(x)\)

\(f_0(x)=q_0(x)*g_0(x)+x^{n-m+1}r_0(x)\)

则如果取模 \(x^{n-m+1}\),将彻底消除余数影响,而 \(q_k(x)\) 次数为 \(x^{n-m}\),所以不受影响。

接着,我们可以通过求逆直接求出 \(q_k(x)\),而反其系数则得到了 \(Q{x}\),然后计算取模即可得 \(R(x)\)

时间复杂度 \(O(nlogn)\)

题目代码

#include<bits/stdc++.h>
using namespace std;
long long n,m,nn=1,nn2=1,uv[4000006][2];
const long long mod=998244353;
long long a[4000006],b[4000006],c[4000006],d[4000006],f[4000006],g[2000006],g2[2000006];
long long pow2(long long a1,long long b1){
    long long c1=1;
    while(b1!=0){
        if(b1%2==1){
            c1=c1*a1%mod;
        }
        a1=a1*a1%mod;
        b1/=2;
    }
    return c1;
}
void ntt(long long n1,long long *a1,long long b1){
    if(n1<=1){
        return ;
    }
    long long ax[n1/2],ay[n1/2];
    for(int i=0;i<=n1-1;i+=2){
        ax[i/2]=a1[i];
        ay[i/2]=a1[i+1];
    }
    ntt(n1/2,ax,b1);
    ntt(n1/2,ay,b1);
    long long ww,w=1;
    if(b1==1){
        if(uv[n1][0]==0){
            ww=pow2(3ll,(mod-1)*pow2(n1,mod-2)%mod);
            uv[n1][0]=pow2(3ll,(mod-1)*pow2(n1,mod-2)%mod);
        }
        else{
            ww=uv[n1][0];
        }
    }
    else{
        if(uv[n1][1]==0){
            ww=pow2(pow2(3ll,(mod-1)*pow2(n1,mod-2)%mod),mod-2);
            uv[n1][1]=pow2(pow2(3ll,(mod-1)*pow2(n1,mod-2)%mod),mod-2);
        }
        else{
            ww=uv[n1][1];
        }
    }
    for(int i=0;i<=n1/2-1;i++,w=w*ww%mod){
        a1[i]=(ax[i]+w*ay[i])%mod;
        a1[i+n1/2]=(ax[i]-w*ay[i]%mod+mod)%mod;
    }
    return ;
}
int main(){
    cin>>n>>m;
    for(int i=0;i<=n;i++){
        cin>>f[i];
    }
    for(int i=0;i<=m;i++){
        cin>>g[i];
        g2[i]=g[i];
    }
    for(int i=0;i<=n/2;i++){
        swap(f[i],f[n-i]);
    }
    for(int i=0;i<=m/2;i++){
        swap(g[i],g[m-i]);
    }
    for(int i=n-m+1;i<=m;i++){
        g[i]=0;
    }
    a[0]=pow2(g[0],mod-2);
    while(nn<=n-m+1+n-m+1){
        nn*=2;
        nn2=pow2(nn*2ll%mod,mod-2);
        for(int i=0;i<=nn-1;i++){
            c[i]=b[i]=0;
        }
        for(int i=0;i<=nn/2-1;i++){
            c[i]=a[i];
        }
        ntt(nn*2,a,1);
        for(int i=0;i<=nn-1;i++){
            b[i]=g[i];
        }
        ntt(nn*2,b,1);
        for(int i=0;i<=nn*2-1;i++){
            a[i]=a[i]*b[i]%mod;
        }
        ntt(nn*2,a,-1);
        for(int i=0;i<=nn*2-1;i++){
            a[i]=a[i]*nn2%mod;
            a[i]=mod-a[i];
        }   
        a[0]=(a[0]+2ll)%mod;
        ntt(nn*2,a,1);
        ntt(nn*2,c,1);
        for(int i=0;i<=nn*2-1;i++){
            a[i]=a[i]*c[i]%mod;
        }
        ntt(nn*2,a,-1); 
        for(int i=0;i<=nn*2-1;i++){
            a[i]=a[i]*nn2%mod;
            if(i>=nn){
                a[i]=0;
            }
        }           
    } 
    nn=1;
    while(nn<=n+n){
        nn*=2;
    }
    nn2=pow2(nn,mod-2);
    for(int i=n-m+1;i<=nn-1;i++){
        a[i]=0;
    }
    ntt(nn,a,1);
    for(int i=0;i<=nn-1;i++){
        b[i]=f[i];
    }
    ntt(nn,b,1);
    for(int i=0;i<=nn;i++){
        a[i]=a[i]*b[i]%mod;
    }    
    ntt(nn,a,-1);
    for(int i=0;i<=nn-1;i++){
        a[i]=a[i]*nn2%mod;
        if(i>=n-m+1){
            a[i]=0;
        }
    }      
    for(int i=0;i<=(n-m)/2;i++){
        swap(a[i],a[n-m-i]);
        //cout<<a[i]<<" ";
    }
    for(int i=0;i<=n-m;i++){
        //swap(a[i],a[n-m-i]);
        cout<<a[i]<<" ";
    }
    cout<<endl;
    for(int i=0;i<=m;i++){
        g[i]=g2[i];
    }   
    for(int i=0;i<=n/2;i++){
        swap(f[i],f[n-i]);
    }   
    nn=1;
    while(nn<=n+n){
        nn*=2;
    }
    nn2=pow2(nn,mod-2);    
    for(int i=0;i<=nn-1;i++){
        b[i]=c[i]=0;
    }
    for(int i=0;i<=n-m;i++){
        c[i]=a[i];
    }    
    ntt(nn,c,1);
    for(int i=0;i<=m;i++){
        b[i]=g[i];
    }    
    ntt(nn,b,1);
    for(int i=0;i<=nn;i++){
        c[i]=c[i]*b[i]%mod;
    }    
    ntt(nn,c,-1);    
    for(int i=0;i<=nn-1;i++){
        c[i]=c[i]*nn2%mod;
    }   
    for(int i=0;i<=n;i++){
        c[i]=(f[i]-c[i]+mod)%mod;
    }    
    for(int i=0;i<=m-1;i++){
        cout<<c[i]<<" ";
    }
    cout<<endl;
	return 0;
}

多项式多点求值

洛谷题目题意和数据范围

给定一个 \(n\) 次多项式 \(f(x)\) ,现在请你对于 \(i \in [1,m]\) ,求出 \(f(a_i)\)

\(n,m \in [1,64000]\)\(a_i,[x^i]f(x) \in [0,998244352]\)

\([x^i]f(x)\) 表示 \(f(x)\)\(i\) 次项系数。

题目思路

首先同样考虑递归解决,那么,我们先解决前面的一半。

设:\(X_0=\{x_1,x_2,x_3,x_4,x_5,......,x_{\lfloor \frac{n}{2}\rfloor}\}\)

\(X_1=\{x_{\lfloor \frac{n}{2}\rfloor+1},......,x_n\}\)

则我们先构造函数:

\(g_0(x)=\prod_{i=1}^{\lfloor \frac{n}{2}\rfloor} (x-x_i)\)

则这个函数对于所有满足 \(x\isin X_0\) 的数来说都是 \(g_0(x)=0\),这样的话可以发现我们对于前半部分,只需要知道其对于 \(h_0(x)=f(x)\% g_0(x)\) 的值即可。

而对于 \(X_1\) 中的元素,我们同样定义函数:

\(g_1(x)=\prod_{i=\lfloor \frac{n}{2}\rfloor+1}^{n} (x-x_i)\)

则这个函数对于所有满足 \(x\isin X_1\) 的数来说都是 \(g_1(x)=0\),这样的话可以发现我们对于前半部分,只需要知道其对于 \(h_1(x)=f(x)\% g_1(x)\) 的值即可。

这样的话每次都将递归的规模给减半,则时间复杂度为 \(T(n)=2T(n/2)+O(nlogn)=O(nlog^2n)\)

时间复杂度为 \(O(nlog^2n)\)

题目代码

#include<bits/stdc++.h>
using namespace std;
long long n,m,nn=1,nn2=1,uv[400005][2];
const int mod=998244353;
long long a[400005],b[400005],c[400005],d[400005],f[400005],g[2000006],g2[2000006],ff[400005],qwerty[400005],yy[400005],xx[400005];
vector<long long> v[1600006];
inline long long pow2(long long a1,long long b1){
	long long c1=1;
	while(b1!=0){
		if(b1%2==1){
			c1=c1*a1%mod;
		}
		a1=a1*a1%mod;
		b1/=2;
	}
	return c1;
}
inline void ntt(long long n1,long long *a1,long long b1){
	if(n1<=1){
		return ;
	}
	long long ax[n1/2],ay[n1/2];
	for(int i=0;i<=n1-1;i+=2){
		ax[i/2]=a1[i];
		ay[i/2]=a1[i+1];
	}
	ntt(n1/2,ax,b1);
	ntt(n1/2,ay,b1);
	long long ww,w=1;
	if(b1==1){
		if(uv[n1][0]==0){
			ww=pow2(3ll,(mod-1)*pow2(n1,mod-2)%mod);
			uv[n1][0]=pow2(3ll,(mod-1)*pow2(n1,mod-2)%mod);
		}
		else{
			ww=uv[n1][0];
		}
	}
	else{
		if(uv[n1][1]==0){
			ww=pow2(pow2(3ll,(mod-1)*pow2(n1,mod-2)%mod),mod-2);
			uv[n1][1]=pow2(pow2(3ll,(mod-1)*pow2(n1,mod-2)%mod),mod-2);
		}
		else{
			ww=uv[n1][1];
		}
	}
	for(int i=0;i<=n1/2-1;i++,w=w*ww%mod){
		a1[i]=(ax[i]+w*ay[i])%mod;
		a1[i+n1/2]=(ax[i]-w*ay[i]%mod+mod)%mod;
	}
	return ;
}
inline void ntt0(long long n1,long long *a1,long long b1,long long l,long long r){
	//cout<<l<<" "<<r<<endl;
	if(l==r){
		//cout<<l<<" "<<r<<endl;
		a1[0]=(-xx[l]+mod)%mod;
		a1[1]=1;
		v[b1].push_back((-xx[l]+mod)%mod);
		v[b1].push_back(1);
		//cout<<n1<<":";
		for(int i=0;i<=n1*2-1;i++){
			v[b1].push_back(a1[i]);
			//cout<<a1[i]<<" ";
		}
		//cout<<endl;        
		return ;
	}
	long long ax[n1*2+1],ay[n1*2+1],n12=pow2(n1*2ll%mod,mod-2);
	for(int i=0;i<=n1*2-1;i++){
		ax[i]=ay[i]=0;
	}
	long long mid=(l+r)/2;
	ntt0(n1/2,ax,b1*2,l,mid);
	ntt0(n1/2,ay,b1*2+1,mid+1,r);
	//cout<<l<<" "<<r<<endl;
	ntt(n1*2,ax,1);
	ntt(n1*2,ay,1);
	for(int i=0;i<=n1*2-1;i++){
		ax[i]=ax[i]*ay[i]%mod;
	}
	ntt(n1*2,ax,-1);
	//cout<<n1<<":";
	for(int i=0;i<=n1*2-1;i++){
		ax[i]=ax[i]*n12%mod;
		v[b1].push_back(ax[i]);
		//cout<<ax[i]<<" ";
		a1[i]=ax[i];
		ax[i]=ay[i]=0;
	}
	//cout<<endl;
	return ;
}
long long ui[400005];
inline void ntt2(long long n1,long long *a1,long long b1,long long l,long long r){
	//cout<<l<<" "<<r<<endl;
	if(l+10000>=r){
        for(int i=l;i<=r;i++){
            ui[i-l+1]=1;
        }
        for(int i=0;i<=n1/2;i++){
    		for(int u=l;u<=r;u++){
				yy[u]=(yy[u]+ui[u-l+1]*a1[i]);
                ui[u-l+1]=ui[u-l+1]*xx[u]%mod;
			}
            if(i%6==0){
                for(int u=l;u<=r;u++){
                    yy[u]%=mod;
                }
            }
		}
        for(int u=l;u<=r;u++){
            yy[u]%=mod;
        }
		//cout<<endl;
		return ;
	}
	long long n,m,nn=1,cff[n1*2+1],cff2[n1*2+1];
	long long mid=(l+r)/2;
	for(int i=0;i<=n1*2;i++){
		f[i]=a[i]=b[i]=c[i]=g[i]=g2[i]=cff[i]=cff2[i]=0;
	}
	n=m=0;
	for(int i=0;i<v[b1*2].size();i++){
		g[i]=v[b1*2][i];
		//cout<<g[i]<<" ";
	}
	for(int i=v[b1*2].size()-1;i>=0;i--){
		if(g[i]!=0){
			m=i;
			break;
		}
	}
	//cout<<endl;
	for(int i=0;i<=n1-1;i++){
		f[i]=a1[i];
		//cout<<a1[i]<<" ";
	}
	for(int i=n1-1;i>=0;i--){
		if(f[i]!=0){
			n=i;
			break;
		}
	}
	//cout<<endl;
	// n=n1-1;
	// m=v[b1*2+1].size()-1;
	for(int i=0;i<=m;i++){
		g2[i]=g[i];
	}
	for(int i=0;i<=n/2;i++){
		swap(f[i],f[n-i]);
	}
	for(int i=0;i<=m/2;i++){
		swap(g[i],g[m-i]);
	}
	for(int i=n-m+1;i<=m;i++){
		g[i]=0;
	}
	a[0]=pow2(g[0],mod-2);
	while(nn<=n-m+1+n-m+1){
		nn*=2;
		nn2=pow2(nn*2ll%mod,mod-2);
		for(int i=0;i<=nn-1;i++){
			c[i]=b[i]=0;
		}
		//		for(int i=0;i<=nn/2-1;i++){
		//			c[i]=a[i];
		//		}
		ntt(nn*2,a,1);
		for(int i=0;i<=nn/2-1;i++){
			c[i]=a[i];
		}
		for(int i=0;i<=nn-1;i++){
			b[i]=g[i];
		}
		ntt(nn*2,b,1);
		for(int i=0;i<=nn*2-1;i++){
			a[i]=a[i]*(2ll-b[i]*a[i]%mod+mod)%mod;
		}
		ntt(nn*2,a,-1);
		for(int i=0;i<=nn*2-1;i++){
			a[i]=a[i]*nn2%mod;
			if(i>=nn){
				a[i]=0;
			}
		}           
	} 
	nn=1;
	while(nn<=n+n){
		nn*=2;
	}
	nn2=pow2(nn,mod-2);
	for(int i=n-m+1;i<=nn-1;i++){
		a[i]=0;
	}
	ntt(nn,a,1);
	for(int i=0;i<=nn-1;i++){
		b[i]=f[i];
	}
	ntt(nn,b,1);
	for(int i=0;i<=nn;i++){
		a[i]=a[i]*b[i]%mod;
	}    
	ntt(nn,a,-1);
	for(int i=0;i<=nn-1;i++){
		a[i]=a[i]*nn2%mod;
		if(i>=n-m+1){
			a[i]=0;
		}
	}      
	for(int i=0;i<=(n-m)/2;i++){
		swap(a[i],a[n-m-i]);
		//cout<<a[i]<<" ";
	}
	// for(int i=0;i<=n-m;i++){
	//     //swap(a[i],a[n-m-i]);
	//     cout<<a[i]<<" ";
	// }
	//cout<<endl;
	for(int i=0;i<=m;i++){
		g[i]=g2[i];
	}   
	for(int i=0;i<=n/2;i++){
		swap(f[i],f[n-i]);
	}   
	nn=1;
	while(nn<=n+n){
		nn*=2;
	}
	nn2=pow2(nn,mod-2);    
	for(int i=0;i<=nn-1;i++){
		b[i]=c[i]=0;
	}
	for(int i=0;i<=n-m;i++){
		c[i]=a[i];
	}    
	ntt(nn,c,1);
	for(int i=0;i<=m;i++){
		b[i]=g[i];
	}    
	ntt(nn,b,1);
	for(int i=0;i<=nn;i++){
		c[i]=c[i]*b[i]%mod;
	}    
	ntt(nn,c,-1);    
	for(int i=0;i<=nn-1;i++){
		c[i]=c[i]*nn2%mod;
	}   
	//cout<<n<<" "<<m<<endl;
	for(int i=0;i<=n1*2;i++){
		c[i]=(f[i]-c[i]+mod)%mod;
		cff[i]=c[i];
		//cff2[i]=c[i];
		//cout<<c[i]<<" ";
	} 
	//cout<<endl;
	ntt2(n1/2,cff,b1*2,l,mid);
	for(int i=0;i<=n1*2;i++){
		f[i]=a[i]=b[i]=c[i]=g[i]=g2[i]=cff[i]=cff2[i]=0;
	}
	for(int i=0;i<v[b1*2+1].size();i++){
		g[i]=v[b1*2+1][i];
		//cout<<g[i]<<" ";
	}
	//cout<<endl;
	m=0;
	n=0;
	for(int i=v[b1*2+1].size()-1;i>=0;i--){
		if(g[i]!=0){
			m=i;
			break;
		}
	}
	for(int i=0;i<=n1-1;i++){
		f[i]=a1[i];
		//cout<<a1[i]<<" ";
	}
	for(int i=n1-1;i>=0;i--){
		if(f[i]!=0){
			n=i;
			break;
		}
	}
	//cout<<endl;
	nn=1;
	// n=n1-1;
	// m=v[b1*2].size()-1;
	for(int i=0;i<=m;i++){
		g2[i]=g[i];
	}
	for(int i=0;i<=n/2;i++){
		swap(f[i],f[n-i]);
	}
	for(int i=0;i<=m/2;i++){
		swap(g[i],g[m-i]);
	}
	for(int i=n-m+1;i<=m;i++){
		g[i]=0;
	}
	a[0]=pow2(g[0],mod-2);
	while(nn<=n-m+1+n-m+1){
		nn*=2;
		nn2=pow2(nn*2ll%mod,mod-2);
		for(int i=0;i<=nn-1;i++){
			c[i]=b[i]=0;
		}
//		for(int i=0;i<=nn/2-1;i++){
//			c[i]=a[i];
//		}
		ntt(nn*2,a,1);
		for(int i=0;i<=nn/2-1;i++){
			c[i]=a[i];
		}
		for(int i=0;i<=nn-1;i++){
			b[i]=g[i];
		}
		ntt(nn*2,b,1);
		for(int i=0;i<=nn*2-1;i++){
			a[i]=a[i]*(2ll-b[i]*a[i]%mod+mod)%mod;
		}
		ntt(nn*2,a,-1);
		for(int i=0;i<=nn*2-1;i++){
			a[i]=a[i]*nn2%mod;
			if(i>=nn){
				a[i]=0;
			}
		}           
	} 
	nn=1;
	while(nn<=n+n){
		nn*=2;
	}
	nn2=pow2(nn,mod-2);
	for(int i=n-m+1;i<=nn-1;i++){
		a[i]=0;
	}
	ntt(nn,a,1);
	for(int i=0;i<=nn-1;i++){
		b[i]=f[i];
	}
	ntt(nn,b,1);
	for(int i=0;i<=nn;i++){
		a[i]=a[i]*b[i]%mod;
	}    
	ntt(nn,a,-1);
	for(int i=0;i<=nn-1;i++){
		a[i]=a[i]*nn2%mod;
		if(i>=n-m+1){
			a[i]=0;
		}
	}      
	for(int i=0;i<=(n-m)/2;i++){
		swap(a[i],a[n-m-i]);
		//cout<<a[i]<<" ";
	}
	// for(int i=0;i<=n-m;i++){
	//     //swap(a[i],a[n-m-i]);
	//     cout<<a[i]<<" ";
	// }
	// cout<<endl;
	for(int i=0;i<=m;i++){
		g[i]=g2[i];
	}   
	for(int i=0;i<=n/2;i++){
		swap(f[i],f[n-i]);
	}   
	nn=1;
	while(nn<=n+n){
		nn*=2;
	}
	nn2=pow2(nn,mod-2);    
	for(int i=0;i<=nn-1;i++){
		b[i]=c[i]=0;
	}
	for(int i=0;i<=n-m;i++){
		c[i]=a[i];
	}    
	ntt(nn,c,1);
	for(int i=0;i<=m;i++){
		b[i]=g[i];
	}    
	ntt(nn,b,1);
	for(int i=0;i<=nn;i++){
		c[i]=c[i]*b[i]%mod;
	}    
	ntt(nn,c,-1);    
	for(int i=0;i<=nn-1;i++){
		c[i]=c[i]*nn2%mod;
	}   
	for(int i=0;i<=n1*2;i++){
		c[i]=(f[i]-c[i]+mod)%mod;
		//cff[i]=c[i];
		cff2[i]=c[i];
		//cout<<c[i]<<" ";
	} 
	//cout<<endl;
	//cout<<n<<endl;
	ntt2(n1/2,cff2,b1*2+1,mid+1,r);
	return ;
}
int main(){
	//freopen("dddddd01.in","r",stdin);
	cin>>n>>m;
	for(int i=0;i<=n;i++){
		scanf("%lld",&ff[i]);
	}
	for(int i=1;i<=m;i++){
		scanf("%lld",&xx[i]);
	}
	//n/=6;
	//m/=2;
	while(nn<=n+m){
		nn*=2;
	}
	//nn/=2;
	//cout<<n<<" "<<m<<endl;
	ntt0(nn,qwerty,1,1,m);
	ntt2(nn,ff,1,1,m);
	for(int i=1;i<=m;i++){
		cout<<yy[i]<<'\n';
	}
	return 0;
}

多项式快速插值

洛谷题目题意和数据范围

给出 \(n\) 个点 \((x_i, y_i)\)

求一个 \(n-1\) 次的多项式 \(f(x)\),使得 \(f(x_i)\equiv y_i\pmod{998244353}\)

\(1 \leqslant n \leqslant 100000\)

\(0 \leqslant x_i, y_i \lt 998244353\)

保证 \(x_i\) 互不相同

对于 \(30\%\) 的数据,\(n \leqslant 5000\)

注意,你输出的数必须是 \([0, 998244353)\) 范围内的整数。

题目思路

终于到了最后的一道题了,这道题同样也是十分有趣的好题。

考虑拉格朗日插值法:

\(f(x)=\sum_{i=1}^{n} y_i\prod_{j\neq i}^{} \frac{x-x_j}{x_i-x_j}\)

\(f_0(x)=\prod_{j=1}^{n} (x-x_j)\)

\(\prod_{j\neq i}^{} (x_i-x_j)=lim_{x\to x_i} \frac{\prod_{j=1}^{n} (x-x_j)}{x-x_i}=f_0'(x_i)\)

\(f(x)=\sum_{i=1}^{n} \frac{y_i}{f_0'(x_i)}\prod_{j\neq i}^{} (x-x_j)\)

那么首先我们就要先计算 \(f_0(x)\) 的系数,我们考虑分治,设:

\(f_1(x)=\prod_{j=1}^{\lfloor \frac{n}{2}\rfloor} (x-x_j)\)

\(f_2(x)=\prod_{j=\lfloor \frac{n}{2}\rfloor+1}^{n} (x-x_j)\)

\(f_0(x)=f_1(x)\cdot f_2(x)\),我们可以发现其时间复杂度也是 \(O(nlog^2n)\)

接着求导后用多点求值即可,这样就求出了所有的 \(f_0'(x_i)\)

接下来,我们设 \(u_i=\frac{y_i}{f_0'(x_i)}\),考虑计算 \(f(x)\),对于 \(n=1\) 的时候,有 \(f(x)=u_1\)\(f_0(x)=x-x_1\),那么设:

\(g_0(x)=\sum_{i=1}^{\lfloor \frac{n}{2}\rfloor} u_i\prod_{j\neq i\land j\le \lfloor \frac{n}{2}\rfloor}^{} (x-x_j)\)

\(g_1(x)=\sum_{i=\lfloor \frac{n}{2}\rfloor+1}^{n} u_i\prod_{j\neq i\land \lfloor \frac{n}{2}\rfloor+1 \le j\le n}^{} (x-x_j)\)

\(f_1(x)=\prod_{j=1}^{\lfloor \frac{n}{2}\rfloor} (x-x_j)\)

\(f_2(x)=\prod_{j=\lfloor \frac{n}{2}\rfloor+1}^{n} (x-x_j)\)

则同样 \(f_0(x)=f_1(x)\cdot f_2(x)\),而 \(f(x)=g_0(x)f_2(x)+g_1(x)f_1(x)\)

时间复杂度为 \(O(nlog^2n)\)

题目代码

#include<bits/stdc++.h>
using namespace std;
long long n,m,nn=1,nn2=1,uv[600005][2];
const int mod=998244353;
const long long rr=(__int128)((__int128)1<<64)/mod;
const long long moc=1ll*mod*mod;
long long a[600005],b[600005],c[600005],d[600005],f[600005],g[600005],g2[600005],ff[600005],qwerty[600005],yy[600005],xx[600005],lr;
long long xq[600005],yq[600005],uq[600005],fe[600005],me[600005];
vector<long long> v[3200006];
inline int mo(long long a1){
    a1-=((__int128)a1*rr>>64)*mod;
    return a1>=mod?a1-mod:a1;
}
inline long long pow2(long long a1,long long b1){
	long long c1=1;
	while(b1!=0){
		if(b1%2==1){
			c1=c1*a1%mod;
		}
		a1=a1*a1%mod;
		b1/=2;
	}
	return c1;
}
inline void ntt(long long n1,long long *a1,long long b1){
	if(n1<=1){
		return ;
	}
	long long ax[n1/2],ay[n1/2];
	for(int i=0;i<=n1-1;i+=2){
		ax[i/2]=a1[i];
		ay[i/2]=a1[i+1];
	}
	ntt(n1/2,ax,b1);
	ntt(n1/2,ay,b1);
	long long ww,w=1;
	if(b1==1){
		if(uv[n1][0]==0){
			ww=pow2(3ll,(mod-1)*pow2(n1,mod-2)%mod);
			uv[n1][0]=pow2(3ll,(mod-1)*pow2(n1,mod-2)%mod);
		}
		else{
			ww=uv[n1][0];
		}
	}
	else{
		if(uv[n1][1]==0){
			ww=pow2(pow2(3ll,(mod-1)*pow2(n1,mod-2)%mod),mod-2);
			uv[n1][1]=pow2(pow2(3ll,(mod-1)*pow2(n1,mod-2)%mod),mod-2);
		}
		else{
			ww=uv[n1][1];
		}
	}
	for(int i=0;i<=n1/2-1;i++,w=w*ww%mod){
		a1[i]=(ax[i]+w*ay[i])%mod;
		a1[i+n1/2]=(ax[i]-w*ay[i]+moc)%mod;
	}
	return ;
}
inline void ntt0(long long n1,long long *a1,long long b1,long long l,long long r){
	//cout<<l<<" "<<r<<endl;
	if(l==r){
		//cout<<l<<" "<<r<<endl;
		a1[0]=(-xx[l]+mod)%mod;
		a1[1]=1;
		v[b1].push_back((-xx[l]+mod)%mod);
		v[b1].push_back(1);
		//cout<<n1<<":";
		for(int i=0;i<=n1*2-1;i++){
			v[b1].push_back(a1[i]);
			//cout<<a1[i]<<" ";
		}
		//cout<<endl;        
		return ;
	}
	long long ax[n1*2+1],ay[n1*2+1],n12=pow2(n1*2ll%mod,mod-2);
	for(int i=0;i<=n1*2-1;i++){
		ax[i]=ay[i]=0;
	}
	long long mid=(l+r)/2;
	ntt0(n1/2,ax,b1*2,l,mid);
	ntt0(n1/2,ay,b1*2+1,mid+1,r);
	//cout<<l<<" "<<r<<endl;
	ntt(n1*2,ax,1);
	ntt(n1*2,ay,1);
	for(int i=0;i<=n1*2-1;i++){
		ax[i]=ax[i]*ay[i]%mod;
	}
	ntt(n1*2,ax,-1);
	//cout<<n1<<":";
	for(int i=0;i<=n1*2-1;i++){
		ax[i]=ax[i]*n12%mod;
		v[b1].push_back(ax[i]);
		//cout<<ax[i]<<" ";
		a1[i]=ax[i];
		ax[i]=ay[i]=0;
	}
	//cout<<endl;
	return ;
}
long long ui[400005];
inline void ntt2(long long n1,long long *a1,long long b1,long long l,long long r){
	//cout<<l<<" "<<r<<endl;
	if(l+2000>=r){
        for(int i=l;i<=r;i++){
            ui[i-l+1]=1;
        }
        for(int i=0;i<=n1/2;i++){
    		for(int u=l;u<=r;u++){
				yy[u]=(yy[u]+ui[u-l+1]*a1[i]);
                ui[u-l+1]=ui[u-l+1]*xx[u]%mod;
			}
            if(i%6==0){
                for(int u=l;u<=r;u++){
                    yy[u]%=mod;
                }
            }
		}
        for(int u=l;u<=r;u++){
            yy[u]%=mod;
        }
		//cout<<endl;
		return ;
	}
	long long n,m,nn=1,cff[n1*2+1],cff2[n1*2+1];
	long long mid=(l+r)/2;
	for(int i=0;i<=n1*2;i++){
		f[i]=a[i]=b[i]=c[i]=g[i]=g2[i]=cff[i]=cff2[i]=0;
	}
	n=m=0;
	for(int i=0;i<v[b1*2].size();i++){
		g[i]=v[b1*2][i];
		//cout<<g[i]<<" ";
	}
	for(int i=v[b1*2].size()-1;i>=0;i--){
		if(g[i]!=0){
			m=i;
			break;
		}
	}
	//cout<<endl;
	for(int i=0;i<=n1-1;i++){
		f[i]=a1[i];
		//cout<<a1[i]<<" ";
	}
	for(int i=n1-1;i>=0;i--){
		if(f[i]!=0){
			n=i;
			break;
		}
	}
	//cout<<endl;
	// n=n1-1;
	// m=v[b1*2+1].size()-1;
	for(int i=0;i<=m;i++){
		g2[i]=g[i];
	}
	for(int i=0;i<=n/2;i++){
		swap(f[i],f[n-i]);
	}
	for(int i=0;i<=m/2;i++){
		swap(g[i],g[m-i]);
	}
	for(int i=n-m+1;i<=m;i++){
		g[i]=0;
	}
	a[0]=pow2(g[0],mod-2);
	while(nn<=n-m+1){
		nn*=2;
		nn2=pow2(nn*2ll%mod,mod-2);
		for(int i=0;i<=nn-1;i++){
			c[i]=b[i]=0;
		}
		//		for(int i=0;i<=nn/2-1;i++){
		//			c[i]=a[i];
		//		}
		ntt(nn*2,a,1);
		for(int i=0;i<=nn/2-1;i++){
			c[i]=a[i];
		}
		for(int i=0;i<=nn-1;i++){
			b[i]=g[i];
		}
		ntt(nn*2,b,1);
		for(int i=0;i<=nn*2-1;i++){
			a[i]=a[i]*(2ll-b[i]*a[i]%mod+mod)%mod;
		}
		ntt(nn*2,a,-1);
		for(int i=0;i<=nn*2-1;i++){
			a[i]=a[i]*nn2%mod;
			if(i>=nn){
				a[i]=0;
			}
		}           
	} 
	nn=1;
	while(nn<=n+n){
		nn*=2;
	}
	nn2=pow2(nn,mod-2);
	for(int i=n-m+1;i<=nn-1;i++){
		a[i]=0;
	}
	ntt(nn,a,1);
	for(int i=0;i<=nn-1;i++){
		b[i]=f[i];
	}
	ntt(nn,b,1);
	for(int i=0;i<=nn;i++){
		a[i]=a[i]*b[i]%mod;
	}    
	ntt(nn,a,-1);
	for(int i=0;i<=nn-1;i++){
		a[i]=a[i]*nn2%mod;
		if(i>=n-m+1){
			a[i]=0;
		}
	}      
	for(int i=0;i<=(n-m)/2;i++){
		swap(a[i],a[n-m-i]);
		//cout<<a[i]<<" ";
	}
	// for(int i=0;i<=n-m;i++){
	//     //swap(a[i],a[n-m-i]);
	//     cout<<a[i]<<" ";
	// }
	//cout<<endl;
	for(int i=0;i<=m;i++){
		g[i]=g2[i];
	}   
	for(int i=0;i<=n/2;i++){
		swap(f[i],f[n-i]);
	}   
	nn=1;
	while(nn<=n+n){
		nn*=2;
	}
	nn2=pow2(nn,mod-2);    
	for(int i=0;i<=nn-1;i++){
		b[i]=c[i]=0;
	}
	for(int i=0;i<=n-m;i++){
		c[i]=a[i];
	}    
	ntt(nn,c,1);
	for(int i=0;i<=m;i++){
		b[i]=g[i];
	}    
	ntt(nn,b,1);
	for(int i=0;i<=nn;i++){
		c[i]=c[i]*b[i]%mod;
	}    
	ntt(nn,c,-1);    
	for(int i=0;i<=nn-1;i++){
		c[i]=c[i]*nn2%mod;
	}   
	//cout<<n<<" "<<m<<endl;
	for(int i=0;i<=n1*2;i++){
		c[i]=(f[i]-c[i]+mod)%mod;
		cff[i]=c[i];
		//cff2[i]=c[i];
		//cout<<c[i]<<" ";
	} 
	//cout<<endl;
	ntt2(n1/2,cff,b1*2,l,mid);
	for(int i=0;i<=n1*2;i++){
		f[i]=a[i]=b[i]=c[i]=g[i]=g2[i]=cff[i]=cff2[i]=0;
	}
	for(int i=0;i<v[b1*2+1].size();i++){
		g[i]=v[b1*2+1][i];
		//cout<<g[i]<<" ";
	}
	//cout<<endl;
	m=0;
	n=0;
	for(int i=v[b1*2+1].size()-1;i>=0;i--){
		if(g[i]!=0){
			m=i;
			break;
		}
	}
	for(int i=0;i<=n1-1;i++){
		f[i]=a1[i];
		//cout<<a1[i]<<" ";
	}
	for(int i=n1-1;i>=0;i--){
		if(f[i]!=0){
			n=i;
			break;
		}
	}
	//cout<<endl;
	nn=1;
	// n=n1-1;
	// m=v[b1*2].size()-1;
	for(int i=0;i<=m;i++){
		g2[i]=g[i];
	}
	for(int i=0;i<=n/2;i++){
		swap(f[i],f[n-i]);
	}
	for(int i=0;i<=m/2;i++){
		swap(g[i],g[m-i]);
	}
	for(int i=n-m+1;i<=m;i++){
		g[i]=0;
	}
	a[0]=pow2(g[0],mod-2);
	while(nn<=n-m+1){
		nn*=2;
		nn2=pow2(nn*2ll%mod,mod-2);
		for(int i=0;i<=nn-1;i++){
			c[i]=b[i]=0;
		}
//		for(int i=0;i<=nn/2-1;i++){
//			c[i]=a[i];
//		}
		ntt(nn*2,a,1);
		for(int i=0;i<=nn/2-1;i++){
			c[i]=a[i];
		}
		for(int i=0;i<=nn-1;i++){
			b[i]=g[i];
		}
		ntt(nn*2,b,1);
		for(int i=0;i<=nn*2-1;i++){
			a[i]=a[i]*(2ll-b[i]*a[i]%mod+mod)%mod;
		}
		ntt(nn*2,a,-1);
		for(int i=0;i<=nn*2-1;i++){
			a[i]=a[i]*nn2%mod;
			if(i>=nn){
				a[i]=0;
			}
		}           
	} 
	nn=1;
	while(nn<=n+n){
		nn*=2;
	}
	nn2=pow2(nn,mod-2);
	for(int i=n-m+1;i<=nn-1;i++){
		a[i]=0;
	}
	ntt(nn,a,1);
	for(int i=0;i<=nn-1;i++){
		b[i]=f[i];
	}
	ntt(nn,b,1);
	for(int i=0;i<=nn;i++){
		a[i]=a[i]*b[i]%mod;
	}    
	ntt(nn,a,-1);
	for(int i=0;i<=nn-1;i++){
		a[i]=a[i]*nn2%mod;
		if(i>=n-m+1){
			a[i]=0;
		}
	}      
	for(int i=0;i<=(n-m)/2;i++){
		swap(a[i],a[n-m-i]);
		//cout<<a[i]<<" ";
	}
	// for(int i=0;i<=n-m;i++){
	//     //swap(a[i],a[n-m-i]);
	//     cout<<a[i]<<" ";
	// }
	// cout<<endl;
	for(int i=0;i<=m;i++){
		g[i]=g2[i];
	}   
	for(int i=0;i<=n/2;i++){
		swap(f[i],f[n-i]);
	}   
	nn=1;
	while(nn<=n){
		nn*=2;
	}
	nn2=pow2(nn,mod-2);    
	for(int i=0;i<=nn-1;i++){
		b[i]=c[i]=0;
	}
	for(int i=0;i<=n-m;i++){
		c[i]=a[i];
	}    
	ntt(nn,c,1);
	for(int i=0;i<=m;i++){
		b[i]=g[i];
	}    
	ntt(nn,b,1);
	for(int i=0;i<=nn;i++){
		c[i]=c[i]*b[i]%mod;
	}    
	ntt(nn,c,-1);    
	for(int i=0;i<=nn-1;i++){
		c[i]=c[i]*nn2%mod;
	}   
	for(int i=0;i<=n1*2;i++){
		c[i]=(f[i]-c[i]+mod)%mod;
		//cff[i]=c[i];
		cff2[i]=c[i];
		//cout<<c[i]<<" ";
	} 
	//cout<<endl;
	//cout<<n<<endl;
	ntt2(n1/2,cff2,b1*2+1,mid+1,r);
	return ;
}
inline void ntt3(long long n1,long long *f1,long long *m1,long long l,long long r){
    if(l==r){
        f1[0]=uq[l];
        m1[0]=(mod-xq[l])%mod;
        m1[1]=1;
        return ;
    }
    long long fx[n1*2+1],fy[n1*2+1],mx[n1*2+1],my[n1*2+1],fg[n1*2+1],nn,nn2;
    long long mid=(l+r)/2;
    for(int i=0;i<=n1*2;i++){
        fx[i]=fy[i]=mx[i]=my[i]=0;
    }
    ntt3(n1/2,fx,mx,l,mid);
    ntt3(n1/2,fy,my,mid+1,r);
    nn=n1*2;
    ntt(nn,fx,1);
    ntt(nn,fy,1);
    ntt(nn,mx,1);
    ntt(nn,my,1);
    nn2=pow2(nn,mod-2);    
	for(int i=0;i<=nn;i++){
		f1[i]=(fx[i]*my[i])%mod;
	}    
	ntt(nn,f1,-1);    
	for(int i=0;i<=nn;i++){
		f1[i]=f1[i]*nn2%mod;
	}  
	for(int i=0;i<=nn;i++){
		fg[i]=(fy[i]*mx[i])%mod;
	}    
	ntt(nn,fg,-1);    
	for(int i=0;i<=nn;i++){
		fg[i]=fg[i]*nn2%mod;
	}  
	for(int i=0;i<=nn;i++){
		f1[i]=(f1[i]+fg[i])%mod;
	}  
	for(int i=0;i<=nn;i++){
		m1[i]=mx[i]*my[i]%mod;
	}    
	ntt(nn,m1,-1);    
	for(int i=0;i<=nn;i++){
		m1[i]=m1[i]*nn2%mod;
	}  
    return ;
}
int main(){
	//freopen("dddddd01.in","r",stdin);
	cin>>n;
	for(int i=1;i<=n;i++){
		scanf("%lld%lld",&xq[i],&yq[i]);
        xx[i]=xq[i];
	}
    //n/=2;
    m=n;
	//n/=6;
	//m/=2;
	while(nn<=n+m){
		nn*=2;
	}
	//nn/=2;
	//cout<<n<<" "<<m<<endl;
	ntt0(nn/2,qwerty,1,1,m);
    for(int i=1;i<=nn;i++){
        ff[i-1]=v[1][i]*i%mod;
        //cout<<ff[i]<<" ";
    }
    //cout<<endl;
	ntt2(nn,ff,1,1,m);
	for(int i=1;i<=m;i++){
		uq[i]=yq[i]*pow2(yy[i],mod-2)%mod;
        //cout<<uq[i]<<" ";
	}
    //cout<<endl;
    ntt3(nn/2,fe,me,1,m);
    for(int i=0;i<=n-1;i++){
        cout<<fe[i]<<" ";
    }
    cout<<endl;
	return 0;
}

学习感受

十分神奇的方法,非常的有意思!

参考文献

oiwiki-快速傅里叶变换

oiwiki-快速数论变换

oiwiki-多项式初等函数

oiwiki-多项式多点求值|快速插值

luogu-题解 P3803 【【模板】多项式乘法(FFT)】

luogu-题解 P3803 【【模板】多项式乘法(NTT)】2

还有一些也忘记了,毕竟最初是3个月前学的,现在只是回顾。。。。。。

posted @ 2026-06-17 20:40  bz02_2023f2  阅读(23)  评论(1)    收藏  举报