FFT 学习笔记
Fast Furina Transform(确信)
我恨数学
注意:笔者尚且没有理解 IDFT,因此在 IDFT 只能背板。
前置知识
注意:如果你已经学过多项式,虚数,单位根,那么可以直接跳过这一部分。
多项式
定义:形如 \(\sum\limits_{i=0}^{n-1} a_i \times x^i\) 的式子叫做多项式。
多项式有两种表示形式:一种是如同上面那样的系数表示,另一种则是点值表示。具体的,则是给出 \(n\) 个 \((x,y)\) 的坐标,满足 \((x,y)\) 在图像上。而这 \(n\) 个点可以唯一得出多项式。
而 FFT,解决的就是多项式乘法。具体的,就是把原先为 \(O(n^2)\) 的乘法变为 \(O(n \log n)\)。
优化方向
这里顺便给出FFT大概是如何优化的。我们需要将系数式和点值式以 \(O(n \log n)\) 的复杂度完成互化,然后乘法部分将在点值式上进行。
虚数
形如 \(a+bi\) 的数字。加减乘除都是显然的,这里需要学习的虚数的向量表示。
具体的,虚数可以写成 \(a \cos \theta + bi \sin \theta\),其中,\(\theta\) 就是这个向量的幅角。然后,虚数相乘就变成了模长相乘,角相加。
然后再讲下单位圆,具体的,就是半径为 1 的,坐标上的圆。因此,此时我们就不需要考虑模长的事情了。
单位根
\(\omega_n\) 表示把单位圆 \(n\) 等分,形成的向量幅角最小的就是单位根。相应的,还有 \(\omega_n^2\) 至 \(\omega_n^n\)。当然,\(\omega_n^0 = \omega_n^n = 1\)。
从下面开始,我们规定 \(n\) 是 2 的幂次,\(m\) = \(n/2\)。
然后是一些重要的性质。\(w_{2n}^{2k} = w_n^k\),\(w_n^{k+n/2}=-w_n^k\),证明显然。
FFT
我们先尝试将一个系数式换成点值式,如果直接做,显然是 \(O(n^2)\) 的,我们不能接受。
因此我们要考虑一些优良的东西。令 \(b_k=\sum\limits_{i=0}^{n-1} a_i \times \omega_n^{ki}\)。那么我们定义 \(A(x) = \sum\limits_{i=0}^{n-1} = a_i \times x^i = \sum\limits_{i=0}^{m-1} a_{2i}x^{2i} + x\sum\limits_{i=0}^{m-1}a_{2i+1}x^{2i}\)。
所以,我们定义 \(A0(x) = \sum\limits_{i=0}^{m-1} a_{2i} \times x^i\),\(A1(x) = \sum\limits_{i=0}^{m-1} a_{2i+1} \times x^i\)。则 \(A(x) = A0(x^2) + x \cdot A1(x^2)\)
然后带入 \(x = \omega_n^k\),我们有 \(A(\omega_n^k)=A0(\omega_m^k)+\omega_n^k\cdot A1(\omega_m^k)\)。然后,我们发现 \(A(\omega_n^{k+m})=A0(\omega_m^k)-\omega_k^n\cdot A1(\omega_m^k)\)。因此,我们可以通过算出一半的 \(A\) 得到全部的 \(A\)。而里面的结构显然也是类似于递推的。
IDFT
这里我不会。只能给出公式 \(a_k = \frac{1}{n} \sum\limits_{i=0}^{n-1} b_i \cdot w_n^{-ki}\)。然后化简一下,同样可以做 FFT。
Code(递归)
至此,递归的写法已出。
#include<bits/stdc++.h>
#include <complex>
using namespace std;
#define IOS ios::sync_with_stdio(false);cin.tie(0),cout.tie(0)
#define File(s) freopen(s".in","r",stdin);freopen(s".out","w",stdout)
#define LL long long
#define fi first
#define se second
const int N = (1 << 20) | 5;
const double PI = acos(-1);
struct Complex{
double x,y;
Complex(double x=0,double y=0): x(x),y(y) {}
};
Complex operator + (Complex a,Complex b){return Complex(a.x+b.x,a.y+b.y);}
Complex operator - (Complex a,Complex b){return Complex(a.x-b.x,a.y-b.y);}
Complex operator * (Complex a,Complex b){return Complex(a.x*b.x-a.y*b.y,a.x*b.y+a.y*b.x);}
Complex a[N],b[N];
int n,m;
void FFT(int limit,Complex *a,int type){
if(limit == 1) return ;
Complex a0[limit>>1],a1[limit>>1];
for(int i=0;i<limit;i+=2)
a0[i>>1] = a[i],a1[i>>1] = a[i+1];
FFT(limit>>1,a0,type);
FFT(limit>>1,a1,type);
Complex wn(cos(2.0*PI/limit),type*sin(2.0*PI/limit)),w = Complex(1,0);
for(int i=0;i<(limit>>1);i++,w=w*wn){
a[i] = a0[i] + w * a1[i];
a[i+(limit>>1)] = a0[i] - w * a1[i];
}
return ;
}
int main(){
IOS;
cin >> n >> m;
n ++ ;
m ++ ;
for(int i=0;i<n;i++)
cin >> a[i].x;
for(int i=0;i<m;i++)
cin >> b[i].x;
int limit = 1;
while(limit < (n+m)) limit <<= 1;
FFT(limit,a,1);
FFT(limit,b,1);
for(int i=0;i<limit;i++)
a[i] = a[i] * b[i];
FFT(limit,a,-1);
for(int i=0;i<n+m-1;i++){
double x = a[i].x / (1.0 * limit);
cout << int(x+0.5) << " ";
}
return 0;
}
优化
这个代码无法通过,不是复杂度不够优,而是频繁的内存复制导致的。
因此,我们需要优化。假设 \(n=8\),观察到递归的最下层的顺序对应的上层编号为 \(0,4,2,6,1,3,5,7\),发现实际上就是编号二进制反过来,因此排序一下就行了。
Code
#include<bits/stdc++.h>
using namespace std;
#define IOS ios::sync_with_stdio(false);cin.tie(0),cout.tie(0)
#define File(s) freopen(s".in","r",stdin);freopen(s".out","w",stdout)
#define LL long long
#define fi first
#define se second
const double PI = acos(-1.0);
const int N = (1 << 21) | 5;
struct Complex{
double x,y;
Complex(double x=0,double y=0) : x(x),y(y) {}
};
Complex operator + (Complex a,Complex b){return Complex(a.x+b.x,a.y+b.y);}
Complex operator - (Complex a,Complex b){return Complex(a.x-b.x,a.y-b.y);}
Complex operator * (Complex a,Complex b){return Complex(a.x*b.x-a.y*b.y,a.x*b.y+a.y*b.x);}
int n,m,r[N];
Complex a[N],b[N];
int limit = 1,l=0;
void FFT(Complex *a,int 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){
Complex wn(cos(PI/mid),type*sin(PI/mid));
int R = (mid << 1);
for(int j=0;j<limit;j+=R){
Complex w(1,0);
for(int k=0;k<mid;k++,w=w*wn){
Complex x = a[j+k],y = w * a[j+k+mid];
a[j+k] = x + y;
a[j+k+mid] = x - y;
}
}
}
return ;
}
int main(){
IOS;
cin >> n >> m;
for(int i=0;i<=n;i++) cin >> a[i].x;
for(int i=0;i<=m;i++) cin >> b[i].x;
while(limit <= (n+m)){
limit *= 2;
l ++ ;
}
r[0] = 0;
for(int i=1;i<=limit;i++)
r[i] = (r[i>>1]>>1) | ((i&1) << l-1);
FFT(a,1);
FFT(b,1);
for(int i=0;i<limit;i++)
a[i] = a[i] * b[i];
FFT(a,-1);
for(int i=0;i<=n+m;i++){
double x = a[i].x / limit;
cout << int(x+0.5) << " ";
}
return 0;
}

浙公网安备 33010602011771号