快速傅里叶变换(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+n}=\omega_{n}^{k}\)
-
\(\omega_{n}^{k+\frac{n}{2}}=-\omega_{n}^{k}\)
-
\(\omega_{2n}^{2k} = \omega_{n}^{k}\)
多项式的点值表示
根据代数基本定理,一个 \(n-1\) 次的多项式可以被 \(n\) 个不同的点值唯一表示,也就是
这 \(n\) 个不同的点,我们选用 \(n\) 个单位根.
两多项式相乘,直接将点值相乘即可,因此 \(FFT\) 的核心思路就是将系数表示转化成点值表示,相乘,转化回系数表示.
FFT:将系数转化成点值
对于多项式 \(A(x)\),令 \(A_0(x)\) 为偶次项组成的多项式,\(A_1(x)\) 为奇数次组成的多项式.
将奇偶次拆分
代入 \(x=\omega_{n}^{k}\),利用单位根的性质化简
同时,对于 \(k' = k+\frac{n}{2}\)(\(k\lt \frac{n}{2}\))
因此可以分治计算.
方便起见,将两多项式的位数都补到 \(2\) 的幂次.
IFFT:将点值转化成系数
和 \(FFT\) 基本一样,区别仅在于将单位根的指数取反,最后系数全部除以 \(n\).
迭代实现 FFT
递归实现的常数较大,考虑用迭代实现.
对于 \((a_0,a_1,a_2,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;
}

浙公网安备 33010602011771号