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\) 即可。

优化

  1. 用迭代来替代递归,这样常数大降。
  2. 使用蝶形优化,这样可以内存调用更快吧。

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

posted @ 2026-04-13 15:26  hnczy  阅读(20)  评论(2)    收藏  举报