多项式乘法与快速傅里叶变换

本文介绍了一种高效计算多项式乘法的方法——快速傅里叶变换(FFT),通过将多项式从系数表示转换为点值表示再进行乘法运算,最后将结果转换回系数表示。这种方法能显著降低计算复杂度。

多项式乘法与快速傅里叶变换

问题介绍

试想这样一个问题,求两个多项式

f(x)=i=0n1aixif(x)=∑i=0n−1aixi

g(x)=i=0m1bixig(x)=∑i=0m−1bixi

的乘积
f(x)g(x)=i=0n+m2j+k=i(aj+bk)xif(x)g(x)=∑i=0n+m−2∑j+k=i(aj+bk)xi

使用传统的方法至少需要 O(n2)O(n2) 的复杂度,下面介绍快速傅里叶变换,将这个过程加速到 O(nlogn)O(nlog⁡n).

问题的快速解法

首先考虑如何用其他方式表示多项式 f(x)=n1i=0aixif(x)=∑i=0n−1aixi .

任取n个不同的数(可以是整数、实数,甚至是复数)

x0,x1,,xn1x0,x1,⋯,xn−1

将其代入 f(x)f(x) 中,就得到一个线性方程组
 f(x0)=y0 f(x1)=y1  f(xn1)=yn1{ f(x0)=y0 f(x1)=y1 ⋯ f(xn−1)=yn−1

只要 nn 足够大,就能够唯一地确定一个多项式,换言之,上述方程组可以表示一个多项式,将这两种多项式的表示方法分别称为系数表示和点值表示.

利用快速傅里叶变换来求多项式乘积的总体思路是

  1. 选取合适的 n 个不同的数 x0,x1,,xn1x0,x1,⋯,xn−1
    • 将多项式 f(x)f(x)g(x)g(x) 转化为点值表示(称为离散傅里叶变换,简称DFTDFT)
    • 计算 f(x)g(x)f(x)g(x) 的点值表示
    • f(x)g(x)f(x)g(x) 转化为系数表示(称为逆离散傅里叶变换,简称 DFT1DFT−1)
    • 下面本文将分步讲解上述过程.

      1. 选取合适的 nn 个不同的数 x0,x1,,xn1

      我们选取复数域上 1n1nnn 个不同的值(或称 nnn 次单位复根)作为 x0,x1,,xn1 的值,即

      xk=ωkn=e2kπin,k=0,2,,n1xk=ωnk=e2kπin,k=0,2,⋯,n−1

      至于指数形式的复数 e2kπine2kπin ,用大家所熟知的欧拉公式即可求得其代数形式
      eiθ=cosθ+isinθeiθ=cos⁡θ+isin⁡θ

      经过简单计算可知
      ωk+mnn=cos(2kπn+2πm)+isin(2kπn+2πm)=cos2kπn+isin2kπn=ωkn,mZωnk+mn=cos⁡(2kπn+2πm)+isin⁡(2kπn+2πm)=cos⁡2kπn+isin⁡2kπn=ωnk,m∈Z

      nn 为偶数时
      (ωnk)2=(e2kπin)2=e2kπin/2=ωn/2k=ωn/2k mod n/2

      其中 a mod ba mod baa 除以 b 的余数,上述两等式将在后文中使用.

      我们为什么要费尽周折选取如此复杂的点呢?是为了使用快速傅里叶变换.

      2. 将多项式 f(x)f(x)g(x)g(x) 转化为点值表示

      考虑多项式 f(x)=n1i=0aixif(x)=∑i=0n−1aixi ,当 n=2m,mZ+n=2m,m∈Z+ 时(当不满足该条件时,向 f(x)f(x) 补充系数为0的高次项来扩大 nn 使其满足该条件),将其化为两个多项式

      f[0](x)=a0+a2x++an2xn22

      f[1](x)=a1+a3x++an1xn22f[1](x)=a1+a3x+⋯+an−1xn−22

      则有
      f(x)=f[0](x2)+xf[1](x2)f(x)=f[0](x2)+xf[1](x2)

      进而
      f(ωkn)=f[0](ωk mod n/2n/2)+ωknf[1](ωk mod n/2n/2)f(ωnk)=f[0](ωn/2k mod n/2)+ωnkf[1](ωn/2k mod n/2)

      也就是说,要求 f(x)f(x)nn 个不同点处的值,只需要求 f[0](x)f[1](x)f[1](x)n2n2 个不同点处的值,由于 n=2m,mZ+n=2m,m∈Z+ ,可对 f[0](x)f[0](x)f[1](x)f[1](x) 重复进行上述过程,最终经过 mm 步后得到 n 个函数
      f[0](x)=a0,f[1](x)=a1,,f[n1](x)=an1f[0](x)=a0,f[1](x)=a1,⋯,f[n−1](x)=an−1

      之后回推得到 f(x)f(x) 的点值表示,上述过程就是快速傅里叶变换的过程,复杂度为 O(nm)O(nm)O(nlogn)O(nlog⁡n).

      当然,还需要对 g(x)g(x) 进行同样的变换.

      3. 计算 f(x)g(x)f(x)g(x) 的点值表示

      点值表示的优点是可以快速地求出两个选取了相同点值的多项式的乘积,例如多项式

       f(x0)=y0 f(x1)=y1  f(xn1)=yn1{ f(x0)=y0 f(x1)=y1 ⋯ f(xn−1)=yn−1

      与多项式
       g(x0)=z0 g(x1)=z1  g(xn1)=zn1{ g(x0)=z0 g(x1)=z1 ⋯ g(xn−1)=zn−1

      的乘积
       f(x0)g(x0)=y0z0 f(x1)g(x1)=y1z1  f(xn1)g(xn1)=yn1zn1{ f(x0)g(x0)=y0z0 f(x1)g(x1)=y1z1 ⋯ f(xn−1)g(xn−1)=yn−1zn−1

      只需要 O(n)O(n) 的复杂度即可求得.

      4. 将 f(x)g(x)f(x)g(x) 转化为系数表示

      下面以 f(x)f(x) 为例,讲解如何将多项式从点值表示转化为系数表示,此过程又称多项式的插值.

      f(x)f(x) 的点值表示写成矩阵形式 Y=VnAY=VnA

       y0 y1 y2 y3  yn1= 1 1 1 1  11ω1nω2nω3nωn1n1ω2nω4nω6nω2(n1)n1ω3nω6nω9nω3(n1)n1ωn1nω2(n1)nω3(n1)nω(n1)(n1)n a0 a1 a2 a3  an1( y0 y1 y2 y3 ⋮ yn−1)=( 1111⋯1 1ωn1ωn2ωn3⋯ωnn−1 1ωn2ωn4ωn6⋯ωn2(n−1) 1ωn3ωn6ωn9⋯ωn3(n−1) ⋮⋮⋮⋮⋱⋮ 1ωnn−1ωn2(n−1)ωn3(n−1)⋯ωn(n−1)(n−1))( a0 a1 a2 a3 ⋮ an−1)

      此处矩阵 VnVn 中的1可视为 ω0nωn0 .

      现在我们已知的是 YYVn ,要求的是 AAVn 是一范德蒙德矩阵,可求得其逆矩阵

      V1n=1n 1 1 1  11ω1nω2nω(n1)n1ω2nω4nω2(n1)n1ω3nω6nω3(n1)n1ω(n1)nω2(n1)nω(n1)(n1)nVn−1=1n( 1111⋯1 1ωn−1ωn−2ωn−3⋯ωn−(n−1) 1ωn−2ωn−4ωn−6⋯ωn−2(n−1) ⋮⋮⋮⋮⋱⋮ 1ωn−(n−1)ωn−2(n−1)ωn−3(n−1)⋯ωn−(n−1)(n−1))

      因此 A=V1nYA=Vn−1Y

       a0 a1 a2 a3  an1=1n 1 1 1 1  11ω1nω2nω3nω(n1)n1ω2nω4nω6nω2(n1)n1ω3nω6nω9nω3(n1)n1ω(n1)nω2(n1)nω3(n1)nω(n1)(n1)n y0 y1 y2 y3  yn1( a0 a1 a2 a3 ⋮ an−1)=1n( 1111⋯1 1ωn−1ωn−2ωn−3⋯ωn−(n−1) 1ωn−2ωn−4ωn−6⋯ωn−2(n−1) 1ωn−3ωn−6ωn−9⋯ωn−3(n−1) ⋮⋮⋮⋮⋱⋮ 1ωn−(n−1)ωn−2(n−1)ωn−3(n−1)⋯ωn−(n−1)(n−1))( y0 y1 y2 y3 ⋮ yn−1)

      也就是说,只需将 YYA 对换,将 ωknωnk 换成 ωknωn−k ,再乘上系数 1n1n ,进行类似步骤2的变换,即可进行逆快速傅里叶变换,算法复杂度同样为 O(nlogn)O(nlog⁡n) .

      按照上述方法将 f(x)g(x)f(x)g(x) 转化为系数表示,本题得解.

      代码实现

      下面给出计算整系数多项式乘积的C++代码

      #include <bits/stdc++.h>
      using namespace std;
      const double pi=acos(-1.0);
      struct cpx
      {
          double x,y;
          cpx(double x=0.0,double y=0.0){this->x=x;this->y=y;}
          cpx operator + (const cpx &b)const{return cpx(x+b.x,y+b.y);}
          cpx operator - (const cpx &b)const{return cpx(x-b.x,y-b.y);}
          cpx operator * (const cpx &b)const{return cpx(x*b.x-y*b.y,b.x*y+x*b.y);}
      };
      inline void Rader(cpx F[],int len)
      {
          int j=len>>1;
          for(int i=1;i<len-1;i++)
          {
              if(i<j)swap(F[i],F[j]);
              int k=len>>1;
              while(j>=k)
              {
                  j-=k;
                  k>>=1;
              }
              if(j<k)j+=k;
          }
      }
      inline cpx w(int n,int k)
      {
          return cpx(cos(2*k*pi/n),sin(2*k*pi/n));
      }
      cpx temp[10005];
      inline void FFT(cpx f[],int len,int flag)
      {
          Rader(f,len);
          int wei=-1,tt=len;
          while(tt)
          {
              wei++;
              tt>>=1;
          }
          for(int it=1;it<=wei;it++)
          {
              for(int i=0;i<len;i++)
              {
                  int x=-1>>it<<it;
                  temp[i]=f[(i&x)+(i&~x>>1)]+w(1<<it,-flag*(i&~x))*f[((i>>it<<1|1)<<it-1)+(i&~x>>1)];
              }
              for(int i=0;i<len;i++)f[i]=temp[i];
          }
          if(flag==-1)for(int i=0;i<len;i++)f[i].x/=len;
      }
      inline void Convolution(cpx f[],int n,cpx g[],int m)
      {
          int len=1;
          while(len<2*max(n,m))len<<=1;
          for(int i=n;i<len;i++)f[i]=cpx(0.0,0.0);
          FFT(f,len,1);
          for(int i=m;i<len;i++)g[i]=cpx(0.0,0.0);
          FFT(g,len,1);
          for(int i=0;i<len;i++)f[i]=f[i]*g[i];
          FFT(f,len,-1);
      }
      cpx f[1005],g[1005];
      int n,m;
      int main()
      {
          while(~scanf("%d",&n))
          {
              for(int i=0;i<n;i++)
              {
                  f[i]=cpx(0.0,0.0);
                  scanf("%lf",&f[i].x);
              }
              scanf("%d",&m);
              for(int i=0;i<m;i++)
              {
                  g[i]=cpx(0.0,0.0);
                  scanf("%lf",&g[i].x);
              }
              Convolution(f,n,g,m);
              for(int i=0;i<=n+m;i++)printf(" %.f",f[i].x);
          }
          return 0;
      }

      例:

      (2x3+x+1)(6x2+2x+3)=12x5+4x4+12x3+8x2+5x+3(2x3+x+1)(6x2+2x+3)=12x5+4x4+12x3+8x2+5x+3

      Input:
      4
      1 1 0 2
      3
      3 2 6
      Output:
      3 5 8 12 4 12 0 0
评论
添加红包

请填写红包祝福语或标题

红包个数最小为10个

红包金额最低5元

当前余额3.43前往充值 >
需支付:10.00
成就一亿技术人!
领取后你会自动成为博主和红包主的粉丝 规则
hope_wisdom
发出的红包
实付
使用余额支付
点击重新获取
扫码支付
钱包余额 0

抵扣说明:

1.余额是钱包充值的虚拟货币,按照1:1的比例进行支付金额的抵扣。
2.余额无法直接购买下载,可以购买VIP、付费专栏及课程。

余额充值