定义
对于只有整数参与的多项式运算,使用快速数论变换(Number-Theoretic
Transform,NTT)可以避免浮点运算带来的精度误差。
思路与简单证明
快速数论变换(NTT)在有限域中使用单位根完成多项式的采样与插值,从而避免
FFT 的浮点误差。设质数 \(p\) 满足 \(p=nq+1\),且 \(g\) 是模 \(p\) 的原根,则
\[
\omega_n^k = g^{qk}
\]
是 \(n\) 次单位根。因为 \(g\) 的阶为 \(p-1\),所以 \(\omega_n^0, \omega_n^1, \ldots,
\omega_n^{n-1}\) 两两不同。
当 \(d \mid n\)
时,单位根满足消去性质
\[
\omega_{dn}^{dk} = \omega_n^k.
\]
当 \(n\)
为偶数时,还满足折半性质
\[
\omega_n^{n/2} = -1.
\]
常用模数为 \(1004535809 =
2^{21}\times479+1\) 和 \(998244353 =
2^{23}\times7\times17+1\),两者的原根均可取 \(3\)。逆变换时需要乘以变换长度的模逆元。
实现
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77
| #include<cstdio> #include<cstring> #include<algorithm> #include<cmath> #define INF 0x3f3f3f3f using namespace std; typedef long long LL; typedef double db; inline int read() { int x=0,f=1; char ch=getchar(); while(ch<'0'||ch>'9') { if(ch=='-')f=-1; ch=getchar(); } while(ch>='0'&&ch<='9') { x=(x<<1)+(x<<3)+ch-'0'; ch=getchar(); } return x*f; } const int MAXN=(1 << 20),MOD=998244353,G=3; int n,m,x[MAXN],y[MAXN]; int qpow(int a,int b) { int ret=1; for(;b;b>>=1,a=(LL)a*a%MOD) if(b&1)ret=(LL)ret*a%MOD; return ret; } void rader(int n,int *x) { for(int i=0,j=0;i<n;++i) { if(i<j)swap(x[i],x[j]); int k=n>>1; while((j^=k)<k)k>>=1; } return; } void NTT(int n,int *x,int flag) { rader(n,x); for(int len=2;len<=n;len<<=1) { int omega_n=qpow(G,(MOD-1)/len); if(flag==-1)omega_n=qpow(omega_n,MOD-2); for(int i=0;i<n;i+=len) { int omega=1; for(int j=i;j<i+(len>>1);++j) { int t1=x[j],t2=(LL)omega*x[j+(len>>1)]%MOD; x[j]=((LL)t1+t2)%MOD; x[j+(len>>1)]=((LL)t1-t2+MOD)%MOD; omega=(LL)omega*omega_n%MOD; } } } if(flag==-1) { int inv_n=qpow(n,MOD-2); for(int i=0;i<n;++i) x[i]=(LL)x[i]*inv_n%MOD; } return; } int main() { n=read();m=read(); for(int i=0;i<=n;++i)x[i]=read(); for(int i=0;i<=m;++i)y[i]=read(); m+=n;n=1; while(n<=m)n<<=1; NTT(n,x,1); NTT(n,y,1); for(int i=0;i<n;++i)x[i]=(LL)x[i]*y[i]%MOD; NTT(n,x,-1); for(int i=0;i<=m;++i)printf("%d ",x[i]); puts(""); return 0; }
|