快速数论变换

定义

对于只有整数参与的多项式运算,使用快速数论变换(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;
}
作者

xqmmcqs

发布于

2018-03-03

更新于

2026-09-19

许可协议

评论

Your browser is out-of-date!

Update your browser to view this website correctly.&npsb;Update my browser now

×