jb 都能看懂的FFT(?)解析
浅谈 FFT
单位根
定义:单位根 \(w_n^k\) 表示将单位圆 \(n\) 等分,在第 \(k\) 个复数的值。
显然的 \(w_{n}^0 = w_n^n=1\)
而 \(w_n^{k}\) 也可以理解为 \(x^n=1\) 的解集。
对于 \(w_n^k\) 。因为 \(e^{2k\pi i} =1\),两边同时开 \(n\) 次方, 可以表示为 \(\sqrt[n]{1} = e^{2k\pi i/n} = \cos(2k\pi/n) + i\sin(2k\pi/n)\)
故 \(w_n^k = \cos(2k\pi/n) + i\sin(2k\pi/n)\)。
还有 \(w_{2n}^{2k} = w_{n}^k , w_n^{k+n/2} = -w_n^k\)
这些带入一下就可以得到。
DFT
我们快速傅里叶变换(FFT)的作用是快速对多项式进行乘积的操作。
对于一个多项式 \(A(x) = \sum_i^{n-1} a_ix^i\) ,用 \(\{a_i\}\) 来表示的叫做系数表示法,这种方法只能做到 \(O(n^2)\)
对于一个多项式 \(A(x) = \sum_i^{n-1} a_ix^i\) ,可以用 \(n\) 个数对 \(\{x_i,y_i\}\) 来表示,表示 \(A(x_i) = y_i\),这种表示方法叫做点值表示法,这很形象。这个方法可以选择 \(n\) 个点算出 \(y3_{i} = y1_i\times y2_i\),然后就可以得到了。但是还是 \(O(n^2)\)
我们显然需要优化下面的算法。
\(A(x) = \sum_i a_i x^i\) 我们对 \(i\) 奇偶分类。
\(A_0(x) = a_0 + a_2x^2 + a_4x^4... + a_{n-1}x^{n-1}\)
\(A_1(x) = a_1x + a_3x^3 + a_5x^5... + a_{n-2}x^{n-2}\)
然后转换成一般的多项式的形式。
\(A'_0(y) = a_0 + a_2y + a_4y^2... + a_{n-1}y^{(n-1)/2}\)
\(A_1'(y) = a_1 + a_3y+a_5y^2 +a_{n-2}y^{(n-3)/2}\)
\(A(x) = A'_0(x^2) + xA'_1(x^2)\)
然后分别带入 \(x=w_n^k,x=w_n^{k+n/2},k\in[1,n/2)\)
可得:
\(A(w_n^k)=A'_0(w_{n/2}^k) + w_n^kA'_1(w_{n/2}^k)\)
\(A(w_n^{k+n/2})=A'_0(w_{n/2}^k) - w_n^kA'_1(w_{n/2}^k)\)
两者极为相似,问题转化成形式相同规模更小的子问题,通过递归可以解决。
时间复杂度可以达到 \(O(n\log n)\)
IDFT
但是我们发现,我们一般使用的多项式都是系数表示法的,这个点值表示法根本不用,所以导致我们这个算法还不能用。
那我们得搞一个逆运算,让点值表示法转化成系数表示法。
当然实际上结论就是将原来的 \(w_n^k\) 改成 \(w_n^{-k}\) 最后除以 \(n\) 即可。
优化
- 用迭代来替代递归,这样常数大降。
- 使用蝶形优化,这样可以内存调用更快吧。
NTT
好吧,我们发现这个 FFT 有个严重的问题就是精度。我们肯定会找一个数论里的东西来替代,我们找到了原根!
原根的可以了解一下,就是在 \(w^n\equiv 1(\bmod p)\) 一般 \(p=998244353\) 此时 \(w=3\) ,带入即可。
#include <bits/stdc++.h>
#define cp complex<double>
using namespace std;
const int N=3E6+5;
const double pi = acos(-1.0);
int n,m,len=1,lim;
cp omg[N],inv[N],a[N],b[N];
int ans[N],nxt[N];
void init(){
while(len<=n+m)len*=2,lim++;
for(int i=0;i<len;i++){
omg[i] = cp(cos(2*pi*i/len),sin(2*pi*i/len));
inv[i] = conj(omg[i]);
nxt[i]= (nxt[i>>1]>>1)|( (i&1)<<(lim-1) ) ;
}
}
void FFT(cp *a,cp *omg){
for(int i=0;i<len;i++)
if(i<nxt[i]) swap(a[i],a[nxt[i]]);
for(int mid=1;mid<len;mid<<=1){
for(int R=mid<<1,j=0;j<len;j+=R){
for(int k=0;k<mid;k++){
cp x=a[j+k],y=omg[len / (mid << 1) * k]*a[j+mid+k];
a[j+k]=x+y;
a[j+mid+k]=x-y;
}
}
}
return;
}
int main() {
scanf("%d%d",&n,&m);
for(int i=0,x;i<=n;i++)
scanf("%d",&x),a[i].real(x);
for(int i=0,x;i<=m;i++)
scanf("%d",&x),b[i].real(x);
init();
FFT(a,omg);
FFT(b,omg);
for(int i=0;i<len;i++)
a[i] *=b[i];
FFT(a,inv);
for(int i=0;i<=n+m;i++)
printf("%d ",(int)floor(a[i].real()/len+0.5));
return 0;
}
分治FFT
用 CDQ 写即可。
#include <bits/stdc++.h>
#define int long long
using namespace std;
const int N = (1 << 18) + 5, mod = 998244353, G = 3, Gi = 332748118;
int n, f[N], g[N], a[N], b[N], nxt[N];
int qpow(int a, int b) {
int res = 1;
while (b) {
if (b & 1) res = res * a % mod;
a = a * a % mod;
b >>= 1;
}
return res;
}
// 这里的 len 和 lim 需要在每次卷积时动态计算
void NTT(int *f, int len, int type) {
int lim = 0; while ((1 << lim) < len) lim++;
for (int i = 0; i < len; i++) {
nxt[i] = (nxt[i >> 1] >> 1) | ((i & 1) << (lim - 1));
if (i < nxt[i]) swap(f[i], f[nxt[i]]);
}
for (int mid = 1; mid < len; mid <<= 1) {
int Wn = qpow(type == 1 ? G : Gi, (mod - 1) / (mid << 1));
for (int i = 0; i < len; i += (mid << 1)) {
int w = 1;
for (int j = 0; j < mid; j++, w = w * Wn % mod) {
int x = f[i + j], y = w * f[i + j + mid] % mod;
f[i + j] = (x + y) % mod;
f[i + j + mid] = (x - y + mod) % mod;
}
}
}
if (type == -1) {
int inv = qpow(len, mod - 2);
for (int i = 0; i < len; i++) f[i] = f[i] * inv % mod;
}
}
void solve(int l, int r) {
if (l == r) return;
int mid = (l + r) >> 1;
solve(l, mid);
int L = 1; while (L <= (r - l + 1)) L <<= 1;
for (int i = 0; i < L; i++) a[i] = b[i] = 0;
for (int i = l; i <= mid; i++) a[i - l] = f[i];
for (int i = 1; i <= r - l; i++) b[i - 1] = g[i];
NTT(a, L, 1); NTT(b, L, 1);
for (int i = 0; i < L; i++) a[i] = a[i] * b[i] % mod;
NTT(a, L, -1);
for (int i = mid + 1; i <= r; i++)
f[i] = (f[i] + a[i - l - 1]) % mod;
solve(mid + 1, r);
}
signed main() {
if (scanf("%lld", &n) == EOF) return 0;
for (int i = 1; i < n; i++) scanf("%lld", &g[i]);
f[0] = 1; // 边界条件
solve(0, n - 1);
for (int i = 0; i < n; i++) printf("%lld%c", f[i], i == n - 1 ? '\n' : ' ');
return 0;
}
End

浙公网安备 33010602011771号