快速傅里叶变换(FFT)/快速数论变换(NTT)

https://www.luogu.com.cn/problem/P3803

给定 \(n\) 次多项式 \(A(x)\)\(m\) 次多项式 \(B(x)\),求乘积 \(C(x)\).

单位根

\(z^n = 1\)\(n\) 个不同的复数根,记为 \(\omega_n\),相当于在复平面内将圆 \(n\) 等分,不同的根可以由一个根的若干次表示.

\[\omega_{n}^{k} = e^{\frac{2\pi ik}{n}} = \cos(\frac{2\pi k}{n})+i\sin(\frac{2\pi k}{n}) \]

单位根具有三个重要的性质

  • \(\omega_{n}^{k+n}=\omega_{n}^{k}\)

  • \(\omega_{n}^{k+\frac{n}{2}}=-\omega_{n}^{k}\)

  • \(\omega_{2n}^{2k} = \omega_{n}^{k}\)

多项式的点值表示

根据代数基本定理,一个 \(n-1\) 次的多项式可以被 \(n\) 个不同的点值唯一表示,也就是

\[(x_0,A(x_0)),(x_1,A(x_1)),\cdots,(x_{n-1},A(x_{n-1})) \]

\(n\) 个不同的点,我们选用 \(n\) 个单位根.

两多项式相乘,直接将点值相乘即可,因此 \(FFT\) 的核心思路就是将系数表示转化成点值表示,相乘,转化回系数表示.

FFT:将系数转化成点值

对于多项式 \(A(x)\),令 \(A_0(x)\) 为偶次项组成的多项式,\(A_1(x)\) 为奇数次组成的多项式.

将奇偶次拆分

\[A(x) = A_0(x^2)+x\cdot A_1(x^2) \]

代入 \(x=\omega_{n}^{k}\),利用单位根的性质化简

\[A(\omega_{n}^{k}) = A_0(\omega_{\frac{n}{2}}^{k})+\omega_{n}^{k}\cdot A_1(\omega_{\frac{n}{2}}^{k}) \]

同时,对于 \(k' = k+\frac{n}{2}\)\(k\lt \frac{n}{2}\)

\[A(\omega_{n}^{k+\frac{n}{2}}) = A_0(\omega_{\frac{n}{2}}^{k})-\omega_{n}^{k}\cdot A_1(\omega_{\frac{n}{2}}^{k}) \]

因此可以分治计算.

方便起见,将两多项式的位数都补到 \(2\) 的幂次.

IFFT:将点值转化成系数

\(FFT\) 基本一样,区别仅在于将单位根的指数取反,最后系数全部除以 \(n\).

迭代实现 FFT

递归实现的常数较大,考虑用迭代实现.

对于 \((a_0,a_1,a_2,a_3)\),分治时数组变化为

\[(a_0,a_1,a_2,a_3)\rightarrow (a_0,a_2),(a_1,a_3)\rightarrow (a_0),(a_2),(a_1),(a_3) \]

最终位置的下标实际上是原下标二进制表达反转,因此先做位逆序置换,自底向上合并.

实现多项式乘法

  • \(A\)\(B\) 次数补到 \(2\) 的幂次,转化成复数.

  • 使用 \(FFT\),将 \(A\)\(B\) 的系数表示转化成点值表示.

  • 点值表示相乘,得到结果多项式 \(C\) 的点值表示.

  • 使用 \(IFFT\),将 \(C\) 的点值表示转化成系数表示,取实部并四舍五入.

时间复杂度 \(\mathcal{O}(L\log L)\)\(L\) 是参与 \(FFT\) 的长度.

递归实现代码

//author:kzssCCC

#include <bits/stdc++.h>
using namespace std;
using ll = long long;
using cd = complex<double>;
const double pi = acos(-1);

void fft(vector<cd>& vec,int n,int op){
	if (n==1) return;
	vector<cd> vec0(n/2),vec1(n/2);
	for (int i=0;i<n/2;i++){
		vec0[i] = vec[i<<1];
		vec1[i] = vec[i<<1|1];
	}

	cd w(cos(2*pi/n),sin(2*pi/n*(op==0?1:-1)));
	fft(vec0,n/2,op);
	fft(vec1,n/2,op);

	cd cur(1,0);
	for (int i=0;i<n/2;i++){
		vec[i] = vec0[i]+cur*vec1[i];
		vec[i+n/2] = vec0[i]-cur*vec1[i];
		cur*=w;
	}
}

void solve(){
	int n,m;
	cin >> n >> m;
	n++,m++;
	vector<int> a(n),b(m);
	for (int i=0;i<n;i++){
		cin >> a[i];
	}
	for (int i=0;i<m;i++){
		cin >> b[i];
	}

	vector<cd> A(a.begin(),a.end());
	vector<cd> B(b.begin(),b.end());
	int N = 1;
	while (N<n+m-1){
		N<<=1;
	}
	A.resize(N);
	B.resize(N);

	fft(A,N,0);
	fft(B,N,0);
	vector<cd> C(N);
	for (int i=0;i<N;i++){
		C[i] = A[i]*B[i];
	}
	fft(C,N,1);
	for (int i=0;i<N;i++){
		C[i] /= N;
	}

	for (int i=0;i<n+m-1;i++){
		cout << llround(C[i].real()) << ' ';
	}
	cout << '\n';
}

int main(){
	ios::sync_with_stdio(false);
	cin.tie(0);
	
	int t = 1;
	// cin >> t;
	while (t--) solve();

	return 0;
}

迭代实现代码

//author:kzssCCC

#include <bits/stdc++.h>
using namespace std;
using ll = long long;
using cd = complex<double>;
const double pi = acos(-1);

void fft(vector<cd>& A,int n,int op){
	int B = 31-__builtin_clz(n);
	for (int i=0;i<n;i++){
		int rev = 0;
		for (int j=0;j<B;j++){
			rev <<= 1;
			rev |= i>>j&1;
		}
		if (i<rev){
			swap(A[i],A[rev]);
		}
	}

	for (int len=2;len<=n;len<<=1){
		cd w(cos(2*pi/len),sin(2*pi/len*(op==0?1:-1)));
		for (int i=0;i<n;i+=len){
			cd cur(1,0);
			for (int j=0;j<len/2;j++){
				cd u = A[i+j];
				cd v = A[i+j+len/2];
				A[i+j] = u+cur*v;
				A[i+j+len/2] = u-cur*v;
				cur*=w;
			}
		}
	}

	if (op==1){
		for (int i=0;i<n;i++){
			A[i] /= n;
		}
	}
}

void solve(){
	int n,m;
	cin >> n >> m;
	n++,m++;
	vector<int> a(n),b(m);
	for (int i=0;i<n;i++){
		cin >> a[i];
	}
	for (int i=0;i<m;i++){
		cin >> b[i];
	}

	vector<cd> A(a.begin(),a.end());
	vector<cd> B(b.begin(),b.end());
	int N = 1;
	while (N<n+m-1){
		N<<=1;
	}
	A.resize(N);
	B.resize(N);

	fft(A,N,0);
	fft(B,N,0);
	vector<cd> C(N);
	for (int i=0;i<N;i++){
		C[i] = A[i]*B[i];
	}
	fft(C,N,1);

	for (int i=0;i<n+m-1;i++){
		cout << llround(C[i].real()) << ' ';
	}
	cout << '\n';
}

int main(){
	ios::sync_with_stdio(false);
	cin.tie(0);
	
	int t = 1;
	// cin >> t;
	while (t--) solve();

	return 0;
}

快速数论变换(NTT)

\(FFT\) 基本一样,运算在模 \(P\) 下进行,用 \(g^{\frac{P-1}{n}}\) 代替单位根,逆变换时把 \(g\) 替换成 \(g\) 的逆元.

\(P\) 一般取 \(998244353\),原根 \(g\) 一般取 \(3\).

代码

//author:kzssCCC

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

const int MOD = 998244353;

ll qpow(ll a,ll b){
	ll res = 1;
	while (b){
		if (b&1){
			res = res*a%MOD;
		}
		a = a*a%MOD;
		b >>= 1;
	}
	return res;
}

const int g = 3;
const int invg = qpow(g,MOD-2);

void ntt(vector<ll>& vec,int n,int op){
	int B = 31-__builtin_clz(n);
	for (int i=0;i<n;i++){
		int rev = 0;
		for (int j=0;j<B;j++){
			rev <<= 1;
			rev |= i>>j&1;
		}
		if (i<rev){
			swap(vec[i],vec[rev]);
		}
	}

	for (int len=2;len<=n;len<<=1){
		ll w = qpow(op==0?g:invg,(MOD-1)/len);
		for (int i=0;i<n;i+=len){
			ll cur = 1;
			for (int j=0;j<len/2;j++){
				ll u = vec[i+j];
				ll v = vec[i+j+len/2];
				vec[i+j] = (u+cur*v%MOD)%MOD;
				vec[i+j+len/2] = (u-cur*v%MOD+MOD)%MOD;
				cur = cur*w%MOD;
			}	
		}
	}

	if (op==1){
		for (int i=0;i<n;i++){
			vec[i] = vec[i]*qpow(n,MOD-2)%MOD;
		}
	}
}

void solve(){
	int n,m;
	cin >> n >> m;
	n++,m++;
	vector<ll> a(n),b(m);
	for (int i=0;i<n;i++){
		cin >> a[i];
	}
	for (int i=0;i<m;i++){
		cin >> b[i];
	}

	int N = 1;
	while (N<n+m-1){
		N<<=1;
	}
	a.resize(N);
	b.resize(N);

	ntt(a,N,0);
	ntt(b,N,0);
	vector<ll> c(N);
	for (int i=0;i<N;i++){
		c[i] = a[i]*b[i]%MOD;
	}
	ntt(c,N,1);
	
	for (int i=0;i<n+m-1;i++){
		cout << c[i] << ' ';
	}
	cout << '\n';
}

int main(){
	ios::sync_with_stdio(false);
	cin.tie(0);
	
	int t = 1;
	// cin >> t;
	while (t--) solve();

	return 0;
}
posted @ 2026-07-06 15:53  kzssCCC  阅读(2)  评论(0)    收藏  举报