5 条题解
-
2
更好的阅读体验:https://blog.csdn.net/tenkuo/article/details/150500190
诱导公式: https://blog.csdn.net/tenkuo/article/details/153632690

// a 数组一开始存的是系数,后面通过 a1 和 a2 计算得出每一个 k 对应的 y 值 // limit:数组的大小,*a 当前数组 void FFT(int limit, complex *a) { if(limit == 1) return; // 递归终止条件:只有一个常数项 //最后 n = 1的情况只有一个常数 a0,本来 a数组里面就存着,啥也不干直接返回就好 complex a1[limit>>1], a2[limit>>1]; // 按下标奇偶性分类 for(int i = 0; i < limit; i += 2) { a1[i>>1] = a[i]; // 偶数下标 a2[i>>1] = a[i+1]; // 奇数下标 } // 递归处理两个子问题 FFT(limit >> 1, a1); FFT(limit >> 1, a2); // 计算单位根 complex Wn = complex( cos(2.0 * Pi / limit), sin(2.0 * Pi / limit) ); //pi 就是圆周率 π,在完整代码中会讲怎么定义,Wn 是一开始先计算的 w_n^1 complex w = complex( 1, 0 ); // 当前的 w_n^k // 此时 w 还只是一个乘积基底,也就是 1,所以等于 1 + 0*i // 合并结果,计算每一个 w_n^k的 y for(int i = 0; i < (limit>>1); i++, w = w * Wn) { //蝴蝶操作(代码后面有解释为啥叫这个) complex t = w * a2[i]; //因为 w * a2[i]计算了两次,设个变量能省点时间 a[i] = a1[i] + t; a[i + (limit>>1)] = a1[i] - t; } }
https://www.luogu.com.cn/problem/P3803#include<bits/stdc++.h> using namespace std; const int N = 3e6 + 10; //这里得开大点,最好是 2 * (maxn + maxm) const double Pi = acos(-1.0); // acos(-1.0)是 π 的精确值(自行百度) struct Complex { // 我这里自定义复数结构体啦,c++库里也有 double x, y; // x为实部,y为虚部 } a[N], b[N]; // 两个多项式的数组 Complex operator+(Complex a, Complex b) { //重载复数加法 return {a.x + b.x, a.y + b.y}; } Complex operator-(Complex a, Complex b) { //重载复数减法 return {a.x - b.x, a.y - b.y}; } Complex operator*(Complex a, Complex b) { return {a.x * b.x - a.y * b.y , a.x * b.y + a.y * b.x}; // 重载复数乘法 (a+bi) * (c+di) = (ac-bd) + (ad+bc)i } int n, m, l, r[N]; // r是位逆序 int limit; // FFT变换长度 void FFT(Complex *A, int type) { //type:1表示正变换,-1表示逆变换,后面有解释 for (int i = 0; i < limit; i++) if(i < r[i]){ swap(A[i], A[r[i]]); //相互是位逆序的交换,只交换一次,避免重复交换 } for (int mid = 1; mid < limit; mid <<= 1) { // mid是当前子问题的半长 //Wn = cos( 2 * Pi / limit ) + i * type * sin( 2 * Pi / limit ) //关于逆变换的 w_n^{-k},就等于 w_n^{n - k},sin值要变成负的,直接乘是 -1的 type就好 Complex Wn = { cos(Pi / mid), type * sin(Pi / mid) }; // R是当前子问题的完整长度,j表示当前处理到哪个位置 for(int R = mid << 1, j = 0; j < limit; j += R) { Complex w = {1, 0}; // 枚举子问题的左半部分(0到mid-1) for(int k = 0; k < mid; k++, w = w * Wn) { // 蝴蝶操作 // 这里可以把 j + k 当作 i,x 当作 a1[i],y 当作 a2[i] Complex x = A[j + k]; Complex y = w * A[j + mid + k]; A[j + k] = x + y; A[j + mid + k] = x - y; } } } } int main() { ios::sync_with_stdio(false); cin.tie(0); int n, m; cin >> n; cin >> m; for (int i = 0; i <= n; i++) { cin >> a[i].x; // 实部为系数,虚部默认为 0 a[i].y = 0; } for(int i = 0; i <= m; i++) { cin >> b[i].x; b[i].y = 0; } limit = 1; l = 0; // 是这个意思 limit < n + m + 1 while (limit <= n + m) { // 实际运用中并不严格要求 n 和 m 都是形如 2^k 的数 // 我们会计算出 >= n + m + 1 最小的二进制数 limit,用 limit 来 FFT // 相当于将多出来的位补 0 limit <<= 1; l++; } for (int i = 0; i < limit; i++) { // 计算二进制位逆序 r[i] = (r[i >> 1] >> 1) | ((i & 1) << (l - 1)); //i>>1:去掉 i 的最低位,r[i>>1]:得到 i>>1 的位逆序 //r[i>>1]>>1:将位逆序右移一位(为新位腾出空间) //i&1:获取 i 的最低位,(i&1)<<(l-1):将最低位移到最高位 //最后合并两部分 //自己试一下验证正确性 // 因为 i >> 1 后,前面肯定会空出一个该有的位 // 所以 r[i >> 1] 的最后一个位是 0 // 这个 0 本来是属于 i 的最低位的,现在移走让最低位到最前面 } // FFT正变换 FFT(a, 1); FFT(b, 1); for (int i = 0; i < limit; i++) { a[i] = a[i] * b[i]; //两两点值相乘 } // FFT逆变换 FFT(a, -1); for(int i = 0; i <= n + m; i++) { cout << (int)(a[i].x / limit + 0.5) << " "; //逆变换后需要除以 limit // 为什么要 + 0.5 四舍五入,理论上应该得到整数 42 // 但由于 fft 浮点误差,实际可能是: // 42.0000000001 或 41.9999999999 } cout << "\n"; return 0; }
#include <bits/stdc++.h> using namespace std; typedef long long LL; const int M = 3e6 + 10; //这里得开大点,最好是 2 * (maxn + maxm) const LL P = 998244353; LL qpow(LL a, LL b) { LL res = 1; a %= P; while (b) { if (b & 1) { res = res * a %P; } a = a * a %P; b /= 2; } return res; } LL a[M], b[M], r[M]; int limit, l; void NTT(LL *A, LL type) { //type:原根的特定幂次(正变换用原根,逆变换用原根的逆元) for (int i = 0; i < limit; i++) if(i < r[i]){ swap(A[i], A[r[i]]); } for (int mid = 1; mid < limit; mid <<= 1) { //mid 是当前半长 // 计算当前长度 2 * mid 对应的单位根:x ^ {limit / (2 * mid)} // 为什么是 limit / (2 * mid)?就相当于原来的 g^{(P - 1) / limit} 上面的幂次 *当前长度 /limit //就等于 g^{(P - 1) / 当前长度} LL Wn = qpow( type, limit / (2 * mid) ); // R是当前子问题的完整长度,j表示当前处理到哪个位置 for(int R = mid << 1, j = 0; j < limit; j += R) { LL w = 1; // 初始化当前单位根为 1(即 w_n^0) for(int k = 0; k < mid; k++, w = w * Wn %P) { // 蝴蝶操作 LL x = A[j + k]; LL y = w * A[j + mid + k] %P; A[j + k] = (x + y)%P; A[j + mid + k] = (x - y + P)%P; //这里一定要 + P!!不然会输出负数!! } } } } int main() { ios::sync_with_stdio(false); cin.tie(0); int n, m; cin >> n; cin >> m; for (int i = 0; i <= n; i++) { cin >> a[i]; } for(int i = 0; i <= m; i++) { cin >> b[i]; } limit = 1; l = 0; while (limit <= n + m) { limit <<= 1; l++; } for (int i = 0; i < limit; i++) { // 计算二进制位逆序 r[i] = (r[i >> 1] >> 1) | ((i & 1) << (l - 1)); } //(P-1)/N 是N次单位根 LL t = qpow(3ll, (P - 1) / limit); // 执行 NTT正变换 NTT(a, t); NTT(b, t); for (int i = 0; i < limit; i++) { a[i] = a[i] * b[i] % P; } // 计算原根的逆元(用于逆变换) LL inv_t = qpow(t, P - 2); // 执行 NTT逆变换 NTT(a, inv_t); LL invl = qpow(limit, P - 2); //计算 N在模 P下的逆元 for (int i = 0; i < n + m + 1; i++) { cout << a[i] * invl % P << " "; // 逆变换后需要除以 limit(乘以 limit的逆元) } cout << endl; return 0; }update 2025.12.7:修正了原文表达模糊的部分,添加了更细致的公式推导。
update 2025.12.14:重写了 6.FFT 逆变换 和 7.FFT 的正确性证明。
update 2026.8.2:修正了定义域错误,添加代码注释。
-
1
//递归版(日常推荐使用,105ms): #include<bits/stdc++.h>//快速数论变换(递归版)(NTT) #define LL long long const int M=4e5+10;//要开4倍 const LL P=(7ll<<26)+1;//(119ll<<23)+1;//(7ll<<26)+1;//(17ll<<27+1); LL qpow(LL a,LL b){LL res=1;for(;b;b>>=1,a=a*a%P)if(b&1)res=res*a%P;return res;} LL A[M],B[M],C[M]; void NTT(LL A[],LL n,LL x) { if(n==1)return;//省略A[0]=A[0],实际是Y[0]=A[0] LL A1[n/2],A2[n/2];//A1,A2必须在函数内定义 for(LL i=0;i<n/2;++i)A1[i]=A[2*i],A2[i]=A[2*i+1]; NTT(A1,n/2,x*x%P); NTT(A2,n/2,x*x%P); for(LL i=0,xi=1;i<n/2;++i,xi=xi*x%P) { A[i] = (A1[i] + A2[i]*xi) % P; A[i+n/2] = ((A1[i] - A2[i]*xi) % P + P) % P; } } int main() { LL n, m;scanf("%lld%lld", &n, &m);n++;m++; for(LL i=0;i<n;i++)scanf("%lld", &A[i]); for(LL i=0;i<m;i++)scanf("%lld", &B[i]); LL N=1;while(N<n+m-1)N<<=1;LL inv_N=qpow(N,P-2); LL x=qpow(3ll,(P-1)/N); NTT(A,N,x); NTT(B,N,x); for(LL i=0;i<N;++i)C[i]=A[i]*B[i]%P; LL inv_x=qpow(x,P-2); NTT(C,N,inv_x); for(LL i=0;i<n+m-1;++i)printf("%lld ",C[i]*inv_N% P); return 0; }//非递归版(日常推荐使用,90ms): #include <bits/stdc++.h>//快速数论变换(非递归版:常用)(NTT) #define LL long long using namespace std; const int M=4e5+10; const LL P=(7ll<<26)+1; LL qpow(LL a,LL b){LL res=1;for(;b;b>>=1,a=a*a%P)if(b&1)res=res*a%P;return res;} LL A[M],B[M],C[M],r[M]; void NTT(LL A[],LL n,LL x) { for(int i=0;i<n;++i)if(i<r[i])swap(A[i],A[r[i]]); for(LL m=2;m<=n;m<<=1) { LL xm=qpow(x,n/m); for(LL i=0;i<n;i+=m) { for(LL j=0,xj=1;j<m/2;++j,xj=xj*xm%P) { LL t1=A[i+j],t2=A[i+j+m/2]*xj%P; A[i+j] =(t1+t2)%P; A[i+j+m/2]=(t1-t2+P)%P; } } } } int main() { int n,m;scanf("%d%d",&n,&m);n++;m++; for(int i=0;i<n;i++)scanf("%lld",&A[i]); for(int i=0;i<m;i++)scanf("%lld",&B[i]); LL N=1;while(N<n+m-1)N<<=1;LL inv_N=qpow(N, P-2); for(LL i=0; i<N;++i)r[i]=r[i/2]/2+(i&1)*N/2; LL x=qpow(3ll,(P-1)/N); NTT(A,N,x); NTT(B,N,x); for(LL i=0;i<N;++i)C[i]=A[i]*B[i]%P; LL inv_x=qpow(x,P-2); NTT(C,N,inv_x); for(LL i=0;i<n+m-1;++i)printf("%lld ",C[i]*inv_N%P); return 0; }//数论变换(教学版,超时30分)(number-theoretic transform, NTT): #include <bits/stdc++.h>//数论变换(教学版)(number-theoretic transform, NTT) #define LL long long using namespace std; const int M = 4e5 + 10;//要开4倍 const LL P = (7ll << 26) + 1;//(119ll<<23)+1;//(7ll<<26)+1;//(17ll<<27+1); LL A[M], YA[M], B[M], YB[M], C[M], YC[M], YCn[M], g, gi, inv; LL qpow(LL a, int b){LL res = 1;for (; b; b >>= 1, a = a * a % P)if (b & 1) res = res * a % P;return res;} LL X[M]; void NTT(LL Y[], LL A[], int n, int op) { LL w = qpow(op == 1 ? g : gi, (P - 1) / n), x = 1; for (int i = 0; i < n; ++i, x = x * w % P) { X[0] = 1; for (int j = 0; j < n; ++j, X[j] = X[j - 1] * x % P) Y[i] = (Y[i] + A[j] * X[j] % P) % P; } } int main() { int n, m;scanf("%d%d", &n, &m);n++;m++; for (int i = 0; i < n; i++)scanf("%lld", &A[i]); for (int i = 0; i < m; i++)scanf("%lld", &B[i]); n = n + m - 1; int N = 1;while (N < n)N <<= 1; g = 3;gi = qpow(g, P - 2);inv = qpow(N, P - 2); NTT(YA, A, N, 1); NTT(YB, B, N, 1); for (int i = 0; i < N; ++i)YC[i] = YA[i] * YB[i] % P; NTT(YCn, YC, N, -1); for (int i = 0; i < n; ++i)printf("%lld ", YCn[i] * inv % P); return 0; }//快速数论变换(递归版:常用)(Fast number-theoretic transform, FNTT) #include <bits/stdc++.h>//快速数论变换(递归版)(FNTT) #define LL long long const int M = 4e5 + 10;//要开4倍 const LL P = (7ll << 26) + 1;//(119ll<<23)+1;//(7ll<<26)+1;//(17ll<<27+1); LL A[M], YA[M], B[M], YB[M], C[M], YC[M], YCn[M], g, gi, inv; LL qpow(LL a, int b){LL res = 1;for (; b; b >>= 1, a = a * a % P)if (b & 1)res = res * a % P;return res;} void FNTT(LL Y[], LL A[], int n, int op) { if (n == 1){Y[0] = A[0];return;} LL A1[n / 2], Y1[n / 2], A2[n / 2], Y2[n / 2]; for (int i = 0; i < n / 2; ++i)A1[i] = A[2 * i], A2[i] = A[2 * i + 1]; FNTT(Y1, A1, n / 2, op); FNTT(Y2, A2, n / 2, op); LL w = qpow(op == 1 ? g : gi, (P - 1) / n), x = 1; for (int i = 0; i < n / 2; ++i, x = x * w % P) { Y[i] = (Y1[i] + Y2[i] * x) % P; Y[i + n / 2] = ((Y1[i] - Y2[i] * x) % P + P) % P; } } int main() { int n, m;scanf("%d%d", &n, &m);n++;m++; for (int i = 0; i < n; i++)scanf("%lld", &A[i]); for (int i = 0; i < m; i++)scanf("%lld", &B[i]); n = n + m - 1; int N = 1;while (N < n)N <<= 1; g = 3;gi = qpow(g, P - 2);inv = qpow(N, P - 2); FNTT(YA, A, N, 1); FNTT(YB, B, N, 1); for (int i = 0; i < N; ++i)YC[i] = YA[i] * YB[i] % P; FNTT(YCn, YC, N, -1); for (int i = 0; i < n; ++i)printf("%lld ", YCn[i] * inv % P); return 0; }//快速数论变换(非递归版)(Fast number-theoretic transform, FNTT) #include <bits/stdc++.h>//快速数论变换(非递归版:常用)(FNTT) #define LL long long using namespace std; const int M = 4e5 + 10;//要开4倍 const LL P = (7ll << 26) + 1;//(119ll<<23)+1;//(7ll<<26)+1;//(17ll<<27+1); LL A[M], YA[M], B[M], YB[M], C[M], YC[M], YCn[M], g, gi, inv; int r[M]; LL qpow(LL a, int b){LL res = 1;for (; b; b >>= 1, a = a * a % P)if (b & 1)res = res * a % P;return res;} void FNTT(LL Y[], LL A[], int n, int op) { for (int i = 0; i < n; ++i)if (i < r[i])swap(A[i], A[r[i]]); for (int i = 0; i < n; ++i)Y[i] = A[i]; for (int m = 2; m <= n; m <<= 1){ LL w = qpow(op == 1 ? g : gi, (P - 1) / m); for (int i = 0; i < n; i += m) { LL x = 1; for (int j = 0; j < m / 2; ++j, x = x * w % P){ LL t1 = Y[i + j], t2 = Y[i + j + m / 2] * x % P; Y[i + j] = (t1 + t2) % P; Y[i + j + m / 2] = (t1 - t2 + P) % P; } } } } int main() { int n, m;scanf("%d%d", &n, &m);n++;m++; for (int i = 0; i < n; i++)scanf("%lld", &A[i]); for (int i = 0; i < m; i++)scanf("%lld", &B[i]); n = n + m - 1; int N = 1;while (N < n)N <<= 1; for (int i = 0; i < N; ++i)r[i] = r[i / 2] / 2 + (i & 1) * N / 2; g = 3;gi = qpow(g, P - 2);inv = qpow(N, P - 2); FNTT(YA, A, N, 1); FNTT(YB, B, N, 1); for (int i = 0; i < N; ++i)YC[i] = YA[i] * YB[i] % P; FNTT(YCn, YC, N, -1); for (int i = 0; i < n; ++i)printf("%lld ", YCn[i] * inv % P); return 0; }//傅里叶变换(危险的教学版,不用!)(Fourier transform,FT) #include<bits/stdc++.h>//傅里叶变换(危险的教学版)(Fourier transform,FT) #define complex complex<double> using namespace std; const double PI=acos(-1.0); const int N=2e5+10;//要开2倍 complex A[N],YA[N],B[N],YB[N],C[N],YC[N],YCn[N]; void FT(complex Y[],complex A[],int n,int op) { complex w1({cos(2*PI/n),sin(2*PI/n)*op}), w({1,0}); for(int i=0;i<n;++i,w*=w1) { complex wk({1,0}); for(int j=0;j<n;j++,wk*=w)Y[i]=Y[i]+A[j]*wk; } } int main() { int n,m;scanf("%d%d",&n,&m);n++;m++; for(int i=0;i<n;++i)scanf("%lf",&A[i]); for(int i=0;i<m;++i)scanf("%lf",&B[i]); n=n+m-1; FT(YA,A,n,1); FT(YB,B,n,1); for(int i=0;i<n;++i)YC[i]=YA[i]*YB[i]; FT(YCn,YC,n,-1); for(int i=0;i<n;++i) printf("%d ",int(YCn[i].real()/n+0.5) ); return 0; }//傅里叶变换(教学版,不用!)(Fourier transform,FT) #include<bits/stdc++.h>//傅里叶变换(教学版)(Fourier transform,FT) #define complex complex<double> using namespace std; const double PI=acos(-1.0); const int N=4e5+10;//要开4倍 complex A[N],YA[N],B[N],YB[N],C[N],YC[N],YCn[N]; void FT(complex Y[],complex A[],int n,int op) { complex w1({cos(2*PI/n),sin(2*PI/n)*op}), w({1,0}); for(int i=0;i<n;++i,w*=w1) { complex wk({1,0}); for(int j=0;j<n;j++,wk*=w)Y[i]=Y[i]+A[j]*wk; } } int main() { int n,m;scanf("%d%d",&n,&m);n++;m++; for(int i=0;i<n;++i)scanf("%lf",&A[i]); for(int i=0;i<m;++i)scanf("%lf",&B[i]); n=n+m-1;int lim=1;while(lim<n)lim<<=1; FT(YA,A,lim,1); FT(YB,B,lim,1); for(int i=0;i<lim;++i)YC[i]=YA[i]*YB[i]; FT(YCn,YC,lim,-1); for(int i=0;i<n;++i) printf("%d ",int(YCn[i].real()/lim+0.5) ); return 0; }//快速傅里叶变换(递归版,不用!)(Fast Fourier Transform,FFT) #include<bits/stdc++.h>//快速傅里叶变换(递归版)(Fast Fourier Transform,FFT) #define complex complex<double> using namespace std; const double PI=acos(-1.0); const int N=4e5+10;//要开4倍 complex A[N],YA[N],B[N],YB[N],C[N],YC[N],YCn[N]; void FFT(complex Y[],complex A[],int n,int op) { if(n==1){Y[0]=A[0];return ;} complex A1[n/2],Y1[n/2],A2[n/2],Y2[n/2]; for(int i=0;i<n/2;++i) A1[i]=A[2*i],A2[i]=A[2*i+1]; FFT(Y1,A1,n/2,op);FFT(Y2,A2,n/2,op); complex w1({cos(2*PI/n),sin(2*PI/n)*op}), wk({1,0}); for(int i=0;i<n/2;++i,wk*=w1) { Y[i]=Y1[i]+Y2[i]*wk; Y[i+n/2]=Y1[i]-Y2[i]*wk; } } int main() { int n,m;scanf("%d%d",&n,&m);n++;m++; for(int i=0;i<n;++i)scanf("%lf",&A[i]); for(int i=0;i<m;++i)scanf("%lf",&B[i]); n=n+m-1;int lim=1;while(lim<n)lim<<=1; FFT(YA,A,lim,1); FFT(YB,B,lim,1); for(int i=0;i<lim;++i)YC[i]=YA[i]*YB[i]; FFT(YCn,YC,lim,-1); for(int i=0;i<n;++i) printf("%d ",int(YCn[i].real()/lim+0.5) ); return 0; }//快速傅里叶变换(非递归版,不用!)(Fast Fourier Transform,FFT) #include<bits/stdc++.h>//快速傅里叶变换(非递归版)(Fast Fourier Transform,FFT) #define complex complex<double> using namespace std; const double PI=acos(-1.0); const int N=4e5+10;//要开4倍 complex A[N],YA[N],B[N],YB[N],C[N],YC[N],YCn[N]; int r[N]; void FFT(complex Y[],complex A[],int n,int op) { for(int i=0;i<n;++i)if(i<r[i])swap(A[i],A[r[i]]); for(int i=0;i<n;++i)Y[i]=A[i]; for(int m=2;m<=n;m<<=1) { complex w1({cos(2*PI/m),sin(2*PI/m)*op}); for(int i=0;i<n;i+=m) { complex wk({1,0}); for(int j=0;j<m/2;++j,wk*=w1) { complex x=Y[i+j],y=Y[i+j+m/2]*wk; Y[i+j]=x+y; Y[i+j+m/2]=x-y; } } } } int main() { int n,m;scanf("%d%d",&n,&m);n++;m++; for(int i=0;i<n;++i)scanf("%lf",&A[i]); for(int i=0;i<m;++i)scanf("%lf",&B[i]); n=n+m-1;int lim=1;while(lim<n)lim<<=1; for(int i=0;i<lim;++i)r[i]=r[i/2]/2+(i&1)*lim/2; FFT(YA,A,lim,1); FFT(YB,B,lim,1); for(int i=0;i<lim;++i)YC[i]=YA[i]*YB[i]; FFT(YCn,YC,lim,-1); for(int i=0;i<n;++i) printf("%d ",int(YCn[i].real()/lim+0.5) ); return 0; } -
-1
我的 FFT:
#include<bits/stdc++.h> using namespace std; #define int long long #define complex complex<double> const int N=4e5+10; const double pi=acos(-1.0); complex a[N],b[N],c[N]; void fft(complex a[],int n,int op) { if(n==1)return ; complex a1[n/2],a2[n/2]; for(int i=0;i<n/2;i++)a1[i]=a[i*2],a2[i]=a[i*2+1]; fft(a1,n/2,op);fft(a2,n/2,op); complex w1({cos(2*pi/n),sin(2*pi/n)*op}),wk({1,0}); for(int i=0;i<n/2;i++,wk*=w1) { a[i]=a1[i]+a2[i]*wk; a[i+n/2]=a1[i]-a2[i]*wk; } } void calc(complex a[],int al,complex b[],int bl,complex c[]) { int lim=1;while(lim<al+bl-1)lim<<=1; fft(a,lim,1);fft(b,lim,1); for(int i=0;i<lim;i++)c[i]=a[i]*b[i]; fft(c,lim,-1); } signed main() { int n,m;cin>>n>>m;n++,m++; for(int i=0;i<n;i++)cin>>a[i]; for(int i=0;i<m;i++)cin>>b[i]; calc(a,n,b,m,c); int lim=1;while(lim<n+m-1)lim<<=1; for(int i=0;i<n+m-1;i++)cout<<(int)(c[i].real()/lim+0.5)<<' '; return 0; }我的 NTT:
#include<bits/stdc++.h> using namespace std; #define int long long const int N=4e5+10,P=(7ll<<26)+1; int qpow(int a,int b){int ans=1;for(;b;b>>=1,a=a*a%P)if(b&1)ans=ans*a%P;return ans;} int a[N],b[N],c[N]; void ntt(int s[],int n,int x) { if(n==1)return; int s1[n/2],s2[n/2]; for(int i=0;i<n/2;i++)s1[i]=s[2*i],s2[i]=s[2*i+1]; ntt(s1,n/2,x*x%P);ntt(s2,n/2,x*x%P); for(int i=0,xi=1;i<n/2;i++,xi=xi*x%P) { s[i]=(s1[i]+s2[i]*xi)%P; s[i+n/2]=((s1[i]-s2[i]*xi)%P+P)%P; } } signed main() { int n,m;cin>>n>>m;n++,m++; for(int i=0;i<n;i++)cin>>a[i]; for(int i=0;i<m;i++)cin>>b[i]; int D=1;while(D<n+m-1)D<<=1;int inv=qpow(D,P-2); int x=qpow(3,(P-1)/D); ntt(a,D,x);ntt(b,D,x); for(int i=0;i<D;i++)c[i]=a[i]*b[i]%P; int inv1=qpow(x,P-2); ntt(c,D,inv1); for(int i=0;i<n+m-1;i++)cout<<c[i]*inv%P<<' '; return 0; } -
-1
#include<bits/stdc++.h> #define complex complex<double> using namespace std; const double PI = acos(-1.0);/*获得pi的值,为赋值w1*/ const int N = 4e5 + 10; complex A[N], B[N];/*存储多项式的系数*/ complex YA[N], YB[N];/*存储转换后的点值*/ complex YC[N]/*点值乘积*/, C[N]/*逆变换成系数*/; void FFT(complex Y[]/*点值*/, complex A[]/*系数*/, int n/*采样点个数*/, int op/*正逆变换*/) { if (n == 1/*只剩下常数项*/) {Y[0] = A[0]/*点值即为系数*/; return ;} complex A1[n / 2]/*偶数项系数*/, Y1[n / 2]/*偶数项系数*/; complex A2[n / 2]/*奇数项系数*/, Y2[n / 2]/*奇数项点值*/; for (int i = 0; i < n / 2; i++)/*奇偶项分离*/ A1[i] = A[2 * i], A2[i] = A[2 * i + 1]; FFT(Y1, A1, n / 2, op); FFT(Y2, A2, n / 2, op);/*递归处理*/ complex w1 = {cos(2 * PI / n), sin(2 * PI / n) * op}; complex wk = {1, 0}; //正变换: 从w0到w(k-1) 逆变换:从wk到w1 //本来要计算n个采样点, 现在只需要计算n/2次 for (int i = 0; i < n / 2; i ++, wk *= w1) { Y[i] = Y1[i] + Y2[i] * wk; Y[i + n / 2] = Y1[i] - Y2[i] * wk; } } int main() { int n, m; cin >> n >> m; n++; m++; for (int i = 0; i < n; ++i) cin >> A[i]; for (int i = 0; i < m; ++i) cin >> B[i]; n = n + m - 1; int lim = 1; while (lim < n) lim <<= 1;/*计算出采样点的数量, 必须是2的幂*/ FFT(YA, A, lim, 1);/*将A(x)用点值表示*/ FFT(YB, B, lim, 1);/*将B(x)用点值表示*/ for (int i = 0; i < lim; ++i) YC[i] = YA[i] * YB[i];/*点值乘法*/ FFT(C, YC, lim, -1);/*逆变换会系数表示法*/ for (int i = 0; i < n; ++i) cout << int(C[i].real() / lim + 0.5) << ' '; return 0; } -
-1
快速数论变换NTT算法
[TOC]
一. 问题
给出两个多项式:一个 项 次 多项式 和一个 项 次多项式 ,,求的各项系数()。
形式如下:
$A(x)=a_0+a_1 * x+a_2 * x^2+ \dots +a_{n-1} * x^{n-1}$
$B(x)=b_0+b_1 * x+b_2 * x^2+ \dots +b_{m-1} * x^{m-1}$
$C(x)=c_0+c_1 * x+c_2 * x^2+ \dots +c_{n+m-2} * x^{n+m-2}$
快速数论变换算法(number-theoretic transform, NTT)是一种计算带模数卷积(convolution)的快速算法。二. 多项式的表示
1. 系数表示:
2. 点值表示:$A(x)=\{(x_0,y_0),(x_1,y_1), \dots ,(x_{n-1},y_{n-1})\}$,要求 的值各不相同
已知 系数表示 容易得到 点值表示。
已知 点值表示 也能得到 系数表示。
即:给定n个不同的点可以确定n-1次函数曲线方程的系数。
NTT算法的重点:理解如何快速算点值。三. 多项式的点值计算
1. 取代入和,求点值
因 ,
可得:2. 如何取?
取
关于 ,当 时,有以下特殊性质:- 互异性 的值两两不同。
- 周期性 因 ,故 (此处及以下省略 ) ,则有: 。
- 对称性
扩大 的选取个数 ,使得 ,且 。
设 ,由于 且 , 所以 必是整数。
设 ,固有:。
取: ,此时有:。
说明:,又因 两两不同,故,而不可能等于 。(模意义下的 等于 )
3. 计算
假设 ,
取:$3^0,3^{58720256},3^{58720256 \times 2},3^{58720256 \times 3},3^{58720256 \times 4},3^{58720256 \times 5},3^{5872025 \times 6},3^{58720256 \times 7}$
即 取:
具体计算 的值如下:
$Y_0=A(x_0)=A(g^0)= a_0 + a_1 + a_2 + a_3 + a_4 + a_5 + a_6 + a_7$
$Y_1=A(x_1)=A(g^1)= a_0 + a_1g + a_2g^2 + a_3g^3 + a_4g^4 + a_5g^5 + a_6g^6 + a_7g^7$
$Y_2=A(x_2)=A(g^2)= a_0 + a_1g^2 + a_2g^4 + a_3g^6 + a_4g^8 + a_5g^{10} + a_6g^{12} + a_7g^{14}$
$Y_3=A(x_3)=A(g^3)= a_0 + a_1g^3 + a_2g^6 + a_3g^9 + a_4g^{12} + a_5g^{15} + a_6g^{18} + a_7g^{21}$
$Y_4=A(x_4)=A(g^4)= a_0 + a_1g^4 + a_2g^8 + a_3g^{12} + a_4g^{16} + a_5g^{20} + a_6g^{24} + a_7g^{28}$
$Y_5=A(x_5)=A(g^5)= a_0 + a_1g^5 + a_2g^{10} + a_3g^{15} + a_4g^{20} + a_5g^{25} + a_6g^{30} + a_7g^{35}$
$Y_6=A(x_6)=A(g^6)= a_0 + a_1g^6 + a_2g^{12} + a_3g^{18} + a_4g^{24} + a_5g^{30} + a_6g^{36} + a_7g^{42}$
$Y_7=A(x_7)=A(g^7)= a_0 + a_1g^7 + a_2g^{14} + a_3g^{21} + a_4g^{28} + a_5g^{35} + a_6g^{42} + a_7g^{49}$
因:$A(x)= a_0 + a_1x + a_2x^2 + a_3x^3 + a_4x^4 + a_5x^5 + a_6x^6 + a_7x^7$
把偶数项和奇数项分开如下:
$A(x)= (a_0 + a_2x^2 + a_4x^4 + a_6x^6) + (a_1 + a_3x^2 + a_5x^4 + a_7x^6) \times x$
设:,
则有: ,具体:
$Y_0=A(x_0)=A(g^0)= A_1((g^0)^2)+A_2((g^0)^2) \times g^0$
$Y_1=A(x_1)=A(g^1)= A_1((g^1)^2)+A_2((g^1)^2) \times g^1$
$Y_2=A(x_2)=A(g^2)= A_1((g^2)^2)+A_2((g^2)^2) \times g^2$
$Y_3=A(x_3)=A(g^3)= A_1((g^3)^2)+A_2((g^3)^2) \times g^3$
$Y_4=A(x_4)=A(g^4)= A_1((g^4)^2)+A_2((g^4)^2) \times g^4$
$Y_5=A(x_5)=A(g^5)= A_1((g^5)^2)+A_2((g^5)^2) \times g^5$
$Y_6=A(x_6)=A(g^6)= A_1((g^6)^2)+A_2((g^6)^2) \times g^6$
$Y_7=A(x_7)=A(g^7)= A_1((g^7)^2)+A_2((g^7)^2) \times g^7$
代码如下:for(LL i=0,xi=1;i<N;++i,xi=xi*x%P) { A(i)=(A1(i)+A2(i)*xi)%P; }注:以上程序块中, 表示 ,而 表示 , 表示
考虑周期性:
考虑对称性:
则有:
对比项目: 和 ,
表达式的左边的变量有:
表达式的右边的变量只有:对比 和
对比 和
对比 和
对比 和
发现只要枚举 ,就可以获得: 则代码中循环次数可以减半,如下:
for(LL i=0,xi=1;i<N/2;++i,xi=xi*x%P) { A(i) =(A1(i)+A2(i)*xi)%P; A(i+N/2)=(A1(i)-A2(i)*xi+P)%P; }故 ,可在 内解决。
同样可得:
以及:
到此,NTT 算法已经学会一半:在 内得到的N个不同的点值。另外一半:已知N个点值反推 的各项系数。四. 已知多项式的点值求系数
已知点值表示:$C(x)=\{(x_0,y_0),(x_1,y_1), \dots ,(x_{N-1},y_{N-1})\}$,求的各项系数。
1. 把 的 值作为新的多项式的系数
$C'(x)=y_0+y_1 * x+y_2 * x^2+ \dots +y_{N-1} * x^{N-1}=\sum\limits_{i=0}^{N-1}y_i * x^i$
2. 取:代入 得到 个新点值:
3. 分析 有惊喜:
$y'_k=C'(g^{-k})=\sum\limits_{i=0}^{N-1}y_i*(g^{-k})^i$--1式
又因:--2式
把2式的代入1式,得到:
$y'_k=C'(g^{-k})=\sum\limits_{i=0}^{N-1}\sum\limits_{j=0}^{N-1}c_j*(g^i)^j*(g^{-k})^i$
$=\sum\limits_{i=0}^{N-1}\sum\limits_{j=0}^{N-1}c_j * g^{i*j}*g^{-k*i}$
$=\sum\limits_{i=0}^{N-1}\sum\limits_{j=0}^{N-1}c_j * g^{i*(j-k)}$
$=\sum\limits_{i=0}^{N-1}\sum\limits_{j=0}^{N-1}c_j * (g^{(j-k)})^i$
*
先分析 :
当时, $=\sum\limits_{i=0}^{N-1}(g^0)^i=\sum\limits_{i=0}^{N-1}1=N$
当时, $=\frac{(g^{(j-k)})^N - 1}{g^{(j-k)} - 1}=\frac{(g^N)^{(j-k)} - 1}{g^{(j-k)} - 1}=\frac{1^{(j-k)} - 1}{g^{(j-k)} - 1}=0$
故: *
$=c_0*0+c_1*0+c_2*0+ \dots +c_k*N+ \dots +c_{N-2}*0+c_{N-1}*0$
因: ,可得: 。
也即是:只要求出 ,根据,可以轻松得到 的各项系数
实际上 的系数只有 项, ,多出的系数为0()。五. 代码与改进
1. 题目:LOJ108多项式乘法
2. 代码
代码1:数论变换(递归版,105ms)
#include<bits/stdc++.h>//快速数论变换(递归版)(NTT) #define LL long long const int M=4e5+10;//要开4倍 const LL P=(7ll<<26)+1;//(119ll<<23)+1;//(7ll<<26)+1;//(17ll<<27+1); LL qpow(LL a,LL b){LL res=1;for(;b;b>>=1,a=a*a%P)if(b&1)res=res*a%P;return res;} LL A[M],B[M],C[M]; void NTT(LL A[],LL n,LL x) { if(n==1)return;//省略A[0]=A[0],实际是Y[0]=A[0] LL A1[n/2],A2[n/2];//A1,A2必须在函数内定义 for(LL i=0;i<n/2;++i)A1[i]=A[2*i],A2[i]=A[2*i+1]; NTT(A1,n/2,x*x%P); NTT(A2,n/2,x*x%P); for(LL i=0,xi=1;i<n/2;++i,xi=xi*x%P) { A[i] = (A1[i] + A2[i]*xi) % P; A[i+n/2] = ((A1[i] - A2[i]*xi) % P + P) % P; } } int main() { LL n, m;scanf("%lld%lld", &n, &m);n++;m++; for(LL i=0;i<n;i++)scanf("%lld", &A[i]); for(LL i=0;i<m;i++)scanf("%lld", &B[i]); LL N=1;while(N<n+m-1)N<<=1;LL inv_N=qpow(N,P-2); LL x=qpow(3ll,(P-1)/N); NTT(A,N,x); NTT(B,N,x); for(LL i=0;i<N;++i)C[i]=A[i]*B[i]%P; LL inv_x=qpow(x,P-2); NTT(C,N,inv_x); for(LL i=0;i<n+m-1;++i)printf("%lld ",C[i]*inv_N% P); return 0; }代码2:快速数论变换(非递归版,90ms)
#include <bits/stdc++.h>//快速数论变换(非递归版:常用)(NTT) #define LL long long using namespace std; const int M=4e5+10; const LL P=(7ll<<26)+1; LL qpow(LL a,LL b){LL res=1;for(;b;b>>=1,a=a*a%P)if(b&1)res=res*a%P;return res;} LL A[M],B[M],C[M],r[M]; void NTT(LL A[],LL n,LL x) { for(int i=0;i<n;++i)if(i<r[i])swap(A[i],A[r[i]]); for(LL m=2;m<=n;m<<=1) { LL xm=qpow(x,n/m); for(LL i=0;i<n;i+=m) { for(LL j=0,xj=1;j<m/2;++j,xj=xj*xm%P) { LL t1=A[i+j],t2=A[i+j+m/2]*xj%P; A[i+j] =(t1+t2)%P; A[i+j+m/2]=(t1-t2+P)%P; } } } } int main() { int n,m;scanf("%d%d",&n,&m);n++;m++; for(int i=0;i<n;i++)scanf("%lld",&A[i]); for(int i=0;i<m;i++)scanf("%lld",&B[i]); LL N=1;while(N<n+m-1)N<<=1;LL inv_N=qpow(N, P-2); for(LL i=0; i<N;++i)r[i]=r[i/2]/2+(i&1)*N/2; LL x=qpow(3ll,(P-1)/N); NTT(A,N,x); NTT(B,N,x); for(LL i=0;i<N;++i)C[i]=A[i]*B[i]%P; LL inv_x=qpow(x,P-2); NTT(C,N,inv_x); for(LL i=0;i<n+m-1;++i)printf("%lld ",C[i]*inv_N%P); return 0; }
- 1
信息
- ID
- 564
- 时间
- 1000ms
- 内存
- 256MiB
- 难度
- 8
- 标签
- 递交数
- 329
- 已通过
- 38
- 上传者