5 条题解

  • 2
    @ 2025-12-14 16:13:42

    更好的阅读体验:https://blog.csdn.net/tenkuo/article/details/150500190

    诱导公式: https://blog.csdn.net/tenkuo/article/details/153632690

    
    // a 数组一开始存的是系数,后面通过 a1 和 a2 计算得出每一个 k 对应的 y 值 
    // limit:数组的大小,*a 当前数组
    void FFT(int limit, complex *a) { 
        if(limit == 1) return;  // 递归终止条件:只有一个常数项
        //最后 n = 1的情况只有一个常数 a0,本来 a数组里面就存着,啥也不干直接返回就好 
        
        complex a1[limit>>1], a2[limit>>1];
        // 按下标奇偶性分类
        for(int i = 0; i < limit; i += 2) {
            a1[i>>1] = a[i];    // 偶数下标
            a2[i>>1] = a[i+1];  // 奇数下标
        }
        
        // 递归处理两个子问题
        FFT(limit >> 1, a1);
        FFT(limit >> 1, a2);
        
        // 计算单位根
        complex Wn = complex( cos(2.0 * Pi / limit), sin(2.0 * Pi / limit) );
        //pi 就是圆周率 π,在完整代码中会讲怎么定义,Wn 是一开始先计算的 w_n^1 
        complex w = complex( 1, 0 );  // 当前的 w_n^k 
        // 此时 w 还只是一个乘积基底,也就是 1,所以等于 1 + 0*i
        
        // 合并结果,计算每一个 w_n^k的 y 
        for(int i = 0; i < (limit>>1); i++, w = w * Wn) {  //蝴蝶操作(代码后面有解释为啥叫这个)
            complex t = w * a2[i]; //因为 w * a2[i]计算了两次,设个变量能省点时间 
            a[i] = a1[i] + t;
            a[i + (limit>>1)] = a1[i] - t;
        }
    }
    

    https://www.luogu.com.cn/problem/P3803

    #include<bits/stdc++.h> 
    using namespace std;
     
    const int N = 3e6 + 10;  //这里得开大点,最好是 2 * (maxn + maxm)
    const double Pi = acos(-1.0); // acos(-1.0)是 π 的精确值(自行百度) 
     
    struct Complex {   // 我这里自定义复数结构体啦,c++库里也有 
        double x, y;  // x为实部,y为虚部
    } a[N], b[N];  // 两个多项式的数组
     
    Complex operator+(Complex a, Complex b) {  //重载复数加法 
        return {a.x + b.x, a.y + b.y}; 
    }
     
    Complex operator-(Complex a, Complex b) {   //重载复数减法 
        return {a.x - b.x, a.y - b.y}; 
    }
     
    Complex operator*(Complex a, Complex b) { 
        return {a.x * b.x - a.y * b.y , a.x * b.y + a.y * b.x}; 
    		// 重载复数乘法 (a+bi) * (c+di) = (ac-bd) + (ad+bc)i
    }
     
    int n, m, l, r[N];  // r是位逆序
    int limit;  // FFT变换长度
     
    void FFT(Complex *A, int type) {   //type:1表示正变换,-1表示逆变换,后面有解释 
        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) {   // mid是当前子问题的半长
            //Wn = cos( 2 * Pi / limit ) + i * type * sin( 2 * Pi / limit ) 
            //关于逆变换的 w_n^{-k},就等于 w_n^{n - k},sin值要变成负的,直接乘是 -1的 type就好 
            Complex Wn = { cos(Pi / mid), type * sin(Pi / mid) };
            
            // R是当前子问题的完整长度,j表示当前处理到哪个位置
            for(int R = mid << 1, j = 0; j < limit; j += R) {
                Complex w = {1, 0};
                
                // 枚举子问题的左半部分(0到mid-1)
                for(int k = 0; k < mid; k++, w = w * Wn) {  // 蝴蝶操作
                // 这里可以把 j + k 当作 i,x 当作 a1[i],y 当作 a2[i] 
                    Complex x = A[j + k];
                    Complex y = w * A[j + mid + k]; 
                    
                    A[j + k] = x + y;   
                    A[j + mid + k] = x - y; 
                }
            }
        }
    }
     
    int main() {
    	ios::sync_with_stdio(false);
    	cin.tie(0); 
    	
        int n, m;
        cin >> n;
        cin >> m;
        for (int i = 0; i <= n; i++) {
            cin >> a[i].x;  // 实部为系数,虚部默认为 0
            a[i].y = 0;
        }
        
        for(int i = 0; i <= m; i++) {
            cin >> b[i].x;
            b[i].y = 0;
        }
        
        limit = 1;
        l = 0;
        // 是这个意思 limit < n + m + 1
        while (limit <= n + m) {   // 实际运用中并不严格要求 n 和 m 都是形如 2^k 的数
            // 我们会计算出 >= n + m + 1 最小的二进制数 limit,用 limit 来 FFT
            // 相当于将多出来的位补 0
            limit <<= 1; 
            l++; 
        }
        
        for (int i = 0; i < limit; i++) {  // 计算二进制位逆序
            r[i] = (r[i >> 1] >> 1) | ((i & 1) << (l - 1));
            //i>>1:去掉 i 的最低位,r[i>>1]:得到 i>>1 的位逆序
            //r[i>>1]>>1:将位逆序右移一位(为新位腾出空间)
            //i&1:获取 i 的最低位,(i&1)<<(l-1):将最低位移到最高位
            //最后合并两部分
            //自己试一下验证正确性 
            // 因为 i >> 1 后,前面肯定会空出一个该有的位
            // 所以 r[i >> 1] 的最后一个位是 0
            // 这个 0 本来是属于 i 的最低位的,现在移走让最低位到最前面
        }
        
        // FFT正变换
        FFT(a, 1);
        FFT(b, 1);
        
        for (int i = 0; i < limit; i++) {
            a[i] = a[i] * b[i];  //两两点值相乘
        }
        
        // FFT逆变换
        FFT(a, -1);
        
        for(int i = 0; i <= n + m; i++) {
            cout << (int)(a[i].x / limit + 0.5) << " ";   //逆变换后需要除以 limit
            // 为什么要 + 0.5 四舍五入,理论上应该得到整数 42
            // 但由于 fft 浮点误差,实际可能是:
            // 42.0000000001  或  41.9999999999
    	}
    	cout << "\n";
        return 0;
    }
    
    

    #include <bits/stdc++.h> 
    using namespace std;
     
    typedef long long LL;
    const int M = 3e6 + 10;  //这里得开大点,最好是 2 * (maxn + maxm)
    const LL P = 998244353;  
     
    LL qpow(LL a, LL b) {
        LL res = 1;
        a %= P;
        while (b) {
        	if (b & 1) {
        		res = res * a %P;
    		}
    		a = a * a %P;
    		b /= 2;
    	}
        return res;
    }
     
    LL a[M], b[M], r[M];
    int limit, l;
     
    void NTT(LL *A, LL type) {  //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) {   //mid 是当前半长 
            // 计算当前长度 2 * mid 对应的单位根:x ^ {limit / (2 * mid)}
            // 为什么是 limit / (2 * mid)?就相当于原来的 g^{(P - 1) / limit} 上面的幂次 *当前长度 /limit
    		//就等于 g^{(P - 1) / 当前长度} 
            LL Wn = qpow( type, limit / (2 * mid) );
            
            // R是当前子问题的完整长度,j表示当前处理到哪个位置
            for(int R = mid << 1, j = 0; j < limit; j += R) {
                LL w = 1;  // 初始化当前单位根为 1(即 w_n^0)
                
                for(int k = 0; k < mid; k++, w = w * Wn %P) {  // 蝴蝶操作
                    LL x = A[j + k];
                    LL y = w * A[j + mid + k] %P; 
                    
                    A[j + k] = (x + y)%P;   
                    A[j + mid + k] = (x - y + P)%P;   //这里一定要 + P!!不然会输出负数!! 
                }
            }
        }
    }
     
    int main() {
    	ios::sync_with_stdio(false);
    	cin.tie(0); 
    	
        int n, m;
        cin >> n;
        cin >> m;
        
        for (int i = 0; i <= n; i++) {
            cin >> a[i]; 
        }
        
        for(int i = 0; i <= m; i++) {
            cin >> b[i];
        }
        
        limit = 1;
        l = 0;
        while (limit <= n + m) {
            limit <<= 1; 
            l++; 
        }
        
        for (int i = 0; i < limit; i++) {  // 计算二进制位逆序
            r[i] = (r[i >> 1] >> 1) | ((i & 1) << (l - 1));
        }
        
        //(P-1)/N 是N次单位根
        LL t = qpow(3ll, (P - 1) / limit);
        
        // 执行 NTT正变换
        NTT(a, t);
        NTT(b, t);
        
        for (int i = 0; i < limit; i++) {
            a[i] = a[i] * b[i] % P;  
        }
        
        // 计算原根的逆元(用于逆变换)
        LL inv_t = qpow(t, P - 2);
        
        // 执行 NTT逆变换
        NTT(a, inv_t);
        
        LL invl = qpow(limit, P - 2); //计算 N在模 P下的逆元
        
        for (int i = 0; i < n + m + 1; i++) {
            cout << a[i] * invl % P << " ";   // 逆变换后需要除以 limit(乘以 limit的逆元)
    	}
    	cout << endl;
        return 0;
    }
    
    

    update 2025.12.7:修正了原文表达模糊的部分,添加了更细致的公式推导。

    update 2025.12.14:重写了 6.FFT 逆变换 和 7.FFT 的正确性证明。

    update 2026.8.2:修正了定义域错误,添加代码注释。

    • 1
      @ 2026-2-8 0:19:19

      G41 快速傅里叶变换 FFT算法 多项式乘法

      G43 快速数论变换 NTT算法

      //递归版(日常推荐使用,105ms):
      #include<bits/stdc++.h>//快速数论变换(递归版)(NTT)
      #define LL long long
      const int M=4e5+10;//要开4倍
      const LL P=(7ll<<26)+1;//(119ll<<23)+1;//(7ll<<26)+1;//(17ll<<27+1);
      LL qpow(LL a,LL b){LL res=1;for(;b;b>>=1,a=a*a%P)if(b&1)res=res*a%P;return res;}
      LL A[M],B[M],C[M];
      void NTT(LL A[],LL n,LL x)
      {
          if(n==1)return;//省略A[0]=A[0],实际是Y[0]=A[0] 
          LL A1[n/2],A2[n/2];//A1,A2必须在函数内定义 
          for(LL i=0;i<n/2;++i)A1[i]=A[2*i],A2[i]=A[2*i+1];
          NTT(A1,n/2,x*x%P);
          NTT(A2,n/2,x*x%P);
          for(LL i=0,xi=1;i<n/2;++i,xi=xi*x%P)
          {
              A[i]     =  (A1[i] + A2[i]*xi) % P;
              A[i+n/2] = ((A1[i] - A2[i]*xi) % P + P) % P;
          }
      }
      int main()
      {
          LL n, m;scanf("%lld%lld", &n, &m);n++;m++;
          for(LL i=0;i<n;i++)scanf("%lld", &A[i]);
          for(LL i=0;i<m;i++)scanf("%lld", &B[i]);
          LL N=1;while(N<n+m-1)N<<=1;LL inv_N=qpow(N,P-2);
          LL x=qpow(3ll,(P-1)/N);
          NTT(A,N,x);
          NTT(B,N,x);
          for(LL i=0;i<N;++i)C[i]=A[i]*B[i]%P;
          LL inv_x=qpow(x,P-2);
          NTT(C,N,inv_x);
          for(LL i=0;i<n+m-1;++i)printf("%lld ",C[i]*inv_N% P);
          return 0;
      }
      
      //非递归版(日常推荐使用,90ms):
      #include <bits/stdc++.h>//快速数论变换(非递归版:常用)(NTT)
      #define LL long long
      using namespace std;
      const int M=4e5+10;
      const LL P=(7ll<<26)+1;
      LL qpow(LL a,LL b){LL res=1;for(;b;b>>=1,a=a*a%P)if(b&1)res=res*a%P;return res;}
      LL A[M],B[M],C[M],r[M];
      void NTT(LL A[],LL n,LL x)
      {
          for(int i=0;i<n;++i)if(i<r[i])swap(A[i],A[r[i]]);
          for(LL m=2;m<=n;m<<=1)
          {
              LL xm=qpow(x,n/m);
              for(LL i=0;i<n;i+=m)
              {
                  for(LL j=0,xj=1;j<m/2;++j,xj=xj*xm%P)
                  {
                      LL t1=A[i+j],t2=A[i+j+m/2]*xj%P;
                      A[i+j]    =(t1+t2)%P;
                      A[i+j+m/2]=(t1-t2+P)%P;
                  }
              }
          }
      }
      int main()
      {
          int n,m;scanf("%d%d",&n,&m);n++;m++;
          for(int i=0;i<n;i++)scanf("%lld",&A[i]);
          for(int i=0;i<m;i++)scanf("%lld",&B[i]);
          LL N=1;while(N<n+m-1)N<<=1;LL inv_N=qpow(N, P-2);
          for(LL i=0; i<N;++i)r[i]=r[i/2]/2+(i&1)*N/2;
          LL x=qpow(3ll,(P-1)/N);
          NTT(A,N,x);
          NTT(B,N,x);
          for(LL i=0;i<N;++i)C[i]=A[i]*B[i]%P;
          LL inv_x=qpow(x,P-2);
          NTT(C,N,inv_x);
          for(LL i=0;i<n+m-1;++i)printf("%lld ",C[i]*inv_N%P);
          return 0;
      }
      
      //数论变换(教学版,超时30分)(number-theoretic transform, NTT):
      #include <bits/stdc++.h>//数论变换(教学版)(number-theoretic transform, NTT)
      #define LL long long
      using namespace std;
      const int M = 4e5 + 10;//要开4倍
      const LL P = (7ll << 26) + 1;//(119ll<<23)+1;//(7ll<<26)+1;//(17ll<<27+1);
      LL A[M], YA[M], B[M], YB[M], C[M], YC[M], YCn[M], g, gi, inv;
      LL qpow(LL a, int b){LL res = 1;for (; b; b >>= 1, a = a * a % P)if (b & 1) res = res * a % P;return res;}
      LL X[M];
      void NTT(LL Y[], LL A[], int n, int op)
      {
          LL w = qpow(op == 1 ? g : gi, (P - 1) / n), x = 1;
          for (int i = 0; i < n; ++i, x = x * w % P)
          {
              X[0] = 1;
              for (int j = 0; j < n; ++j, X[j] = X[j - 1] * x % P)
                  Y[i] = (Y[i] + A[j] * X[j] % P) % P;
          }
      }
      int main()
      {
          int n, m;scanf("%d%d", &n, &m);n++;m++;
          for (int i = 0; i < n; i++)scanf("%lld", &A[i]);
          for (int i = 0; i < m; i++)scanf("%lld", &B[i]);
          n = n + m - 1;
          int N = 1;while (N < n)N <<= 1;
          g = 3;gi = qpow(g, P - 2);inv = qpow(N, P - 2);
          NTT(YA, A, N, 1);
          NTT(YB, B, N, 1);
          for (int i = 0; i < N; ++i)YC[i] = YA[i] * YB[i] % P;
          NTT(YCn, YC, N, -1);
          for (int i = 0; i < n; ++i)printf("%lld ", YCn[i] * inv % P);
          return 0;
      }
      
      //快速数论变换(递归版:常用)(Fast number-theoretic transform, FNTT)
      #include <bits/stdc++.h>//快速数论变换(递归版)(FNTT)
      #define LL long long
      const int M = 4e5 + 10;//要开4倍
      const LL P = (7ll << 26) + 1;//(119ll<<23)+1;//(7ll<<26)+1;//(17ll<<27+1);
      LL A[M], YA[M], B[M], YB[M], C[M], YC[M], YCn[M], g, gi, inv;
      LL qpow(LL a, int b){LL res = 1;for (; b; b >>= 1, a = a * a % P)if (b & 1)res = res * a % P;return res;}
      void FNTT(LL Y[], LL A[], int n, int op)
      {
          if (n == 1){Y[0] = A[0];return;}
          LL A1[n / 2], Y1[n / 2], A2[n / 2], Y2[n / 2];
          for (int i = 0; i < n / 2; ++i)A1[i] = A[2 * i], A2[i] = A[2 * i + 1];
          FNTT(Y1, A1, n / 2, op);
          FNTT(Y2, A2, n / 2, op);
          LL w = qpow(op == 1 ? g : gi, (P - 1) / n), x = 1;
          for (int i = 0; i < n / 2; ++i, x = x * w % P)
          {
              Y[i] = (Y1[i] + Y2[i] * x) % P;
              Y[i + n / 2] = ((Y1[i] - Y2[i] * x) % P + P) % P;
          }
      }
      int main()
      {
          int n, m;scanf("%d%d", &n, &m);n++;m++;
          for (int i = 0; i < n; i++)scanf("%lld", &A[i]);
          for (int i = 0; i < m; i++)scanf("%lld", &B[i]);
          n = n + m - 1;
          int N = 1;while (N < n)N <<= 1;
          g = 3;gi = qpow(g, P - 2);inv = qpow(N, P - 2);
          FNTT(YA, A, N, 1);
          FNTT(YB, B, N, 1);
          for (int i = 0; i < N; ++i)YC[i] = YA[i] * YB[i] % P;
          FNTT(YCn, YC, N, -1);
          for (int i = 0; i < n; ++i)printf("%lld ", YCn[i] * inv % P);
          return 0;
      }
      
      //快速数论变换(非递归版)(Fast number-theoretic transform, FNTT)
      #include <bits/stdc++.h>//快速数论变换(非递归版:常用)(FNTT)
      #define LL long long
      using namespace std;
      const int M = 4e5 + 10;//要开4倍
      const LL P = (7ll << 26) + 1;//(119ll<<23)+1;//(7ll<<26)+1;//(17ll<<27+1);
      LL A[M], YA[M], B[M], YB[M], C[M], YC[M], YCn[M], g, gi, inv;
      int r[M];
      LL qpow(LL a, int b){LL res = 1;for (; b; b >>= 1, a = a * a % P)if (b & 1)res = res * a % P;return res;}
      void FNTT(LL Y[], LL A[], int n, int op)
      {
          for (int i = 0; i < n; ++i)if (i < r[i])swap(A[i], A[r[i]]);
          for (int i = 0; i < n; ++i)Y[i] = A[i];
          for (int m = 2; m <= n; m <<= 1){
              LL w = qpow(op == 1 ? g : gi, (P - 1) / m);
              for (int i = 0; i < n; i += m)
              {
                  LL x = 1;
                  for (int j = 0; j < m / 2; ++j, x = x * w % P){
                      LL t1 = Y[i + j], t2 = Y[i + j + m / 2] * x % P;
                      Y[i + j] = (t1 + t2) % P;
                      Y[i + j + m / 2] = (t1 - t2 + P) % P;
                  }
              }
          }
      }
      int main()
      {
          int n, m;scanf("%d%d", &n, &m);n++;m++;
          for (int i = 0; i < n; i++)scanf("%lld", &A[i]);
          for (int i = 0; i < m; i++)scanf("%lld", &B[i]);
          n = n + m - 1;
          int N = 1;while (N < n)N <<= 1;
          for (int i = 0; i < N; ++i)r[i] = r[i / 2] / 2 + (i & 1) * N / 2;
          g = 3;gi = qpow(g, P - 2);inv = qpow(N, P - 2);
          FNTT(YA, A, N, 1);
          FNTT(YB, B, N, 1);
          for (int i = 0; i < N; ++i)YC[i] = YA[i] * YB[i] % P;
          FNTT(YCn, YC, N, -1);
          for (int i = 0; i < n; ++i)printf("%lld ", YCn[i] * inv % P);
          return 0;
      }
      
      //傅里叶变换(危险的教学版,不用!)(Fourier transform,FT)
      #include<bits/stdc++.h>//傅里叶变换(危险的教学版)(Fourier transform,FT)
      #define complex complex<double>
      using namespace std;
      const double PI=acos(-1.0);
      const int N=2e5+10;//要开2倍 
      complex A[N],YA[N],B[N],YB[N],C[N],YC[N],YCn[N];
      void FT(complex Y[],complex A[],int n,int op)
      {
      	complex w1({cos(2*PI/n),sin(2*PI/n)*op}), w({1,0});
      	for(int i=0;i<n;++i,w*=w1)
      	{
      		complex wk({1,0});
      		for(int j=0;j<n;j++,wk*=w)Y[i]=Y[i]+A[j]*wk;
      	}
      }
      int main()
      {
          int n,m;scanf("%d%d",&n,&m);n++;m++;
          for(int i=0;i<n;++i)scanf("%lf",&A[i]);
          for(int i=0;i<m;++i)scanf("%lf",&B[i]);
          n=n+m-1;
          FT(YA,A,n,1);
      	FT(YB,B,n,1);
          for(int i=0;i<n;++i)YC[i]=YA[i]*YB[i];
          FT(YCn,YC,n,-1);
          for(int i=0;i<n;++i) printf("%d ",int(YCn[i].real()/n+0.5) );
          return 0;
      }
      
      //傅里叶变换(教学版,不用!)(Fourier transform,FT)
      #include<bits/stdc++.h>//傅里叶变换(教学版)(Fourier transform,FT)
      #define complex complex<double>
      using namespace std;
      const double PI=acos(-1.0);
      const int N=4e5+10;//要开4倍 
      complex A[N],YA[N],B[N],YB[N],C[N],YC[N],YCn[N];
      void FT(complex Y[],complex A[],int n,int op)
      {
      	complex w1({cos(2*PI/n),sin(2*PI/n)*op}), w({1,0});
      	for(int i=0;i<n;++i,w*=w1)
      	{
      		complex wk({1,0});
      		for(int j=0;j<n;j++,wk*=w)Y[i]=Y[i]+A[j]*wk;
      	}
      }
      int main()
      {
          int n,m;scanf("%d%d",&n,&m);n++;m++;
          for(int i=0;i<n;++i)scanf("%lf",&A[i]);
          for(int i=0;i<m;++i)scanf("%lf",&B[i]);
          n=n+m-1;int lim=1;while(lim<n)lim<<=1;
          FT(YA,A,lim,1);
      	FT(YB,B,lim,1);
          for(int i=0;i<lim;++i)YC[i]=YA[i]*YB[i];
          FT(YCn,YC,lim,-1);
          for(int i=0;i<n;++i) printf("%d ",int(YCn[i].real()/lim+0.5) );
          return 0;
      }
      
      //快速傅里叶变换(递归版,不用!)(Fast Fourier Transform,FFT)
      #include<bits/stdc++.h>//快速傅里叶变换(递归版)(Fast Fourier Transform,FFT)
      #define complex complex<double>
      using namespace std;
      const double PI=acos(-1.0);
      const int N=4e5+10;//要开4倍
      complex A[N],YA[N],B[N],YB[N],C[N],YC[N],YCn[N];
      void FFT(complex Y[],complex A[],int n,int op)
      {
      	if(n==1){Y[0]=A[0];return ;}
      	complex A1[n/2],Y1[n/2],A2[n/2],Y2[n/2];
      	for(int i=0;i<n/2;++i) A1[i]=A[2*i],A2[i]=A[2*i+1];
      	FFT(Y1,A1,n/2,op);FFT(Y2,A2,n/2,op);
      	complex w1({cos(2*PI/n),sin(2*PI/n)*op}), wk({1,0});
      	for(int i=0;i<n/2;++i,wk*=w1)
      	{
      		Y[i]=Y1[i]+Y2[i]*wk;
      		Y[i+n/2]=Y1[i]-Y2[i]*wk;
      	}
      }
      int main()
      {
          int n,m;scanf("%d%d",&n,&m);n++;m++; 
          for(int i=0;i<n;++i)scanf("%lf",&A[i]);
          for(int i=0;i<m;++i)scanf("%lf",&B[i]);
          n=n+m-1;int lim=1;while(lim<n)lim<<=1;
          FFT(YA,A,lim,1);
      	FFT(YB,B,lim,1);
          for(int i=0;i<lim;++i)YC[i]=YA[i]*YB[i];
          FFT(YCn,YC,lim,-1);
          for(int i=0;i<n;++i) printf("%d ",int(YCn[i].real()/lim+0.5) );
          return 0;
      }
      
      //快速傅里叶变换(非递归版,不用!)(Fast Fourier Transform,FFT)
      #include<bits/stdc++.h>//快速傅里叶变换(非递归版)(Fast Fourier Transform,FFT)
      #define complex complex<double>
      using namespace std;
      const double PI=acos(-1.0);
      const int N=4e5+10;//要开4倍
      complex A[N],YA[N],B[N],YB[N],C[N],YC[N],YCn[N];
      int r[N];
      void FFT(complex Y[],complex A[],int n,int op)
      {
          for(int i=0;i<n;++i)if(i<r[i])swap(A[i],A[r[i]]);
          for(int i=0;i<n;++i)Y[i]=A[i];
          for(int m=2;m<=n;m<<=1)
          {
              complex w1({cos(2*PI/m),sin(2*PI/m)*op});
              for(int i=0;i<n;i+=m)
              {
                  complex wk({1,0});
                  for(int j=0;j<m/2;++j,wk*=w1)
                  {
                      complex x=Y[i+j],y=Y[i+j+m/2]*wk;
                      Y[i+j]=x+y;
      				Y[i+j+m/2]=x-y;
                  }
              }
          }
      }
      int main()
      {
          int n,m;scanf("%d%d",&n,&m);n++;m++;
          for(int i=0;i<n;++i)scanf("%lf",&A[i]);
          for(int i=0;i<m;++i)scanf("%lf",&B[i]);
          n=n+m-1;int lim=1;while(lim<n)lim<<=1;
          for(int i=0;i<lim;++i)r[i]=r[i/2]/2+(i&1)*lim/2;
          FFT(YA,A,lim,1);
      	FFT(YB,B,lim,1);
          for(int i=0;i<lim;++i)YC[i]=YA[i]*YB[i];
          FFT(YCn,YC,lim,-1);
          for(int i=0;i<n;++i) printf("%d ",int(YCn[i].real()/lim+0.5) );
          return 0;
      }
      
      • -1
        @ 2026-8-3 15:06:46

        我的 FFT:

        #include<bits/stdc++.h>
        using namespace std;
        #define int long long
        #define complex complex<double>
        const int N=4e5+10;
        const double pi=acos(-1.0);
        complex a[N],b[N],c[N];
        void fft(complex a[],int n,int op)
        {
        	if(n==1)return ;
        	complex a1[n/2],a2[n/2];
        	for(int i=0;i<n/2;i++)a1[i]=a[i*2],a2[i]=a[i*2+1];
        	fft(a1,n/2,op);fft(a2,n/2,op);
        	complex w1({cos(2*pi/n),sin(2*pi/n)*op}),wk({1,0});
        	for(int i=0;i<n/2;i++,wk*=w1)
        	{
        		a[i]=a1[i]+a2[i]*wk;
        		a[i+n/2]=a1[i]-a2[i]*wk;
        	}
        }
        void calc(complex a[],int al,complex b[],int bl,complex c[])
        {
        	int lim=1;while(lim<al+bl-1)lim<<=1;
        	fft(a,lim,1);fft(b,lim,1);
        	for(int i=0;i<lim;i++)c[i]=a[i]*b[i];
        	fft(c,lim,-1);
        }
        signed main()
        {
        	int n,m;cin>>n>>m;n++,m++;
        	for(int i=0;i<n;i++)cin>>a[i];
        	for(int i=0;i<m;i++)cin>>b[i];
        	calc(a,n,b,m,c);
        	int lim=1;while(lim<n+m-1)lim<<=1;
        	for(int i=0;i<n+m-1;i++)cout<<(int)(c[i].real()/lim+0.5)<<' ';
        	return 0;
        }
        

        我的 NTT:

        #include<bits/stdc++.h>
        using namespace std;
        #define int long long
        const int N=4e5+10,P=(7ll<<26)+1;
        int qpow(int a,int b){int ans=1;for(;b;b>>=1,a=a*a%P)if(b&1)ans=ans*a%P;return ans;}
        int a[N],b[N],c[N];
        void ntt(int s[],int n,int x)
        {
        	if(n==1)return;
        	int s1[n/2],s2[n/2];
        	for(int i=0;i<n/2;i++)s1[i]=s[2*i],s2[i]=s[2*i+1];
        	ntt(s1,n/2,x*x%P);ntt(s2,n/2,x*x%P);
        	for(int i=0,xi=1;i<n/2;i++,xi=xi*x%P)
        	{
        		s[i]=(s1[i]+s2[i]*xi)%P;
        		s[i+n/2]=((s1[i]-s2[i]*xi)%P+P)%P;
        	}
        }
        signed main()
        {
        	int n,m;cin>>n>>m;n++,m++;
        	for(int i=0;i<n;i++)cin>>a[i];
        	for(int i=0;i<m;i++)cin>>b[i];
        	int D=1;while(D<n+m-1)D<<=1;int inv=qpow(D,P-2);
        	int x=qpow(3,(P-1)/D);
        	ntt(a,D,x);ntt(b,D,x);
        	for(int i=0;i<D;i++)c[i]=a[i]*b[i]%P;
        	int inv1=qpow(x,P-2);
        	ntt(c,D,inv1);
        	for(int i=0;i<n+m-1;i++)cout<<c[i]*inv%P<<' ';
        	return 0;
        }
        • -1
          @ 2026-3-8 9:04:06
          #include<bits/stdc++.h>
          #define complex complex<double>
          using namespace std;
          const double PI = acos(-1.0);/*获得pi的值,为赋值w1*/
          const int N = 4e5 + 10;
          complex A[N], B[N];/*存储多项式的系数*/
          complex YA[N], YB[N];/*存储转换后的点值*/
          complex YC[N]/*点值乘积*/, C[N]/*逆变换成系数*/;
          void FFT(complex Y[]/*点值*/, complex A[]/*系数*/, int n/*采样点个数*/, int op/*正逆变换*/)
          {
              if (n == 1/*只剩下常数项*/) {Y[0] = A[0]/*点值即为系数*/; return ;}
              complex A1[n / 2]/*偶数项系数*/, Y1[n / 2]/*偶数项系数*/;
              complex A2[n / 2]/*奇数项系数*/, Y2[n / 2]/*奇数项点值*/;
              for (int i = 0; i < n / 2; i++)/*奇偶项分离*/
                  A1[i] = A[2 * i], A2[i] = A[2 * i + 1];
              FFT(Y1, A1, n / 2, op); FFT(Y2, A2, n / 2, op);/*递归处理*/
              complex w1 = {cos(2 * PI / n), sin(2 * PI / n) * op};
              complex wk = {1, 0};
              //正变换: 从w0到w(k-1)  逆变换:从wk到w1
              //本来要计算n个采样点, 现在只需要计算n/2次 
              for (int i = 0; i < n / 2; i ++, wk *= w1)
              {
                  Y[i] = Y1[i] + Y2[i] * wk;
                  Y[i + n / 2] = Y1[i] - Y2[i] * wk;
              }
          }
          
          int main()
          {
              int n, m; cin >> n >> m; n++; m++;
              for (int i = 0; i < n; ++i) cin >> A[i];
              for (int i = 0; i < m; ++i) cin >> B[i];
              n = n + m - 1;
              int lim = 1;
              while (lim < n)
                  lim <<= 1;/*计算出采样点的数量, 必须是2的幂*/
              FFT(YA, A, lim, 1);/*将A(x)用点值表示*/
              FFT(YB, B, lim, 1);/*将B(x)用点值表示*/
              for (int i = 0; i < lim; ++i)
                  YC[i] = YA[i] * YB[i];/*点值乘法*/
              FFT(C, YC, lim, -1);/*逆变换会系数表示法*/
              for (int i = 0; i < n; ++i)
                  cout << int(C[i].real() / lim + 0.5) << ' ';
              return 0;
          }
          
          • -1
            @ 2025-12-14 16:28:19

            快速数论变换NTT算法

            [TOC]

            一. 问题

            给出两个多项式:一个 nnn1n-1 次 多项式 A(x)A(x)和一个 mmm1m-1 次多项式 B(x)B(x)C(x)=A(x)B(x)C(x)=A(x) * B(x),求C(x)C(x)的各项系数(1n,m1051 \leq n,m \leq {10}^5)。
            形式如下:
            $A(x)=a_0+a_1 * x+a_2 * x^2+ \dots +a_{n-1} * x^{n-1}$
            $B(x)=b_0+b_1 * x+b_2 * x^2+ \dots +b_{m-1} * x^{m-1}$
            $C(x)=c_0+c_1 * x+c_2 * x^2+ \dots +c_{n+m-2} * x^{n+m-2}$
            快速数论变换算法(number-theoretic transform, NTT)是一种计算带模数卷积(convolution)的快速算法。

            二. 多项式的表示

            1. 系数表示:A(x)=a0+a1x+a2x2++an1xn1A(x)=a_0+a_1 * x+a_2 * x^2+\dots+a_{n-1} * x^{n-1}

            2. 点值表示:$A(x)=\{(x_0,y_0),(x_1,y_1), \dots ,(x_{n-1},y_{n-1})\}$,要求 xix_i 的值各不相同

            已知 系数表示 容易得到 点值表示
            已知 点值表示 也能得到 系数表示
            即:给定n个不同的点可以确定n-1次函数曲线方程的系数。
            NTT算法的重点:理解如何快速算点值。

            三. 多项式的点值计算

            1. 取x0,x1,x2,,xn+m2x_0,x_1,x_2,\dots,x_{n+m-2}代入A(x)A(x)B(x)B(x),求点值

            A(x):YA,0,YA,1,YA,2,,YA,n+m2A(x):Y_{A,0},Y_{A,1},Y_{A,2},\dots,Y_{A,n+m-2}
            B(x):YB,0,YB,1,YB,2,,YB,n+m2B(x):Y_{B,0},Y_{B,1},Y_{B,2},\dots,Y_{B,n+m-2}
            C(x)=A(x)B(x)C(x)=A(x) * B(x)YC,i=YA,iYB,iY_{C,i}=Y_{A,i} * Y_{B,i}
            可得:C(x):YC,0,YC,1,YC,2,,YC,n+m2C(x):Y_{C,0},Y_{C,1},Y_{C,2}, \dots ,Y_{C,n+m-2}

            2. xx 如何取?

            xx3imodP3^i \bmod P
            关于 3imodP 3^i \bmod P ,当 P=7226+1=469762049 P=7 * 2^{26}+1=469762049 时,有以下特殊性质:

            • 互异性 3imodP(i[1,P1]) 3^i \bmod P( i \in [1,P-1] ) 的值两两不同。
            • 周期性 因 3P11(modP)3^{P-1} \equiv 1(\mod P),故 3P=31,3P+1=32,3P+2=333^P=3^1,3^{P+1}=3^2,3^{P+2}=3^3(此处及以下省略 modP\mod P ) ,则有: gk=gk+Ng^k=g^{k+N}
            • 对称性 扩大 x x 的选取个数 NN ,使得 N=2kN=2^k ,且 n+m1N226n+m-1 \leq N \leq 2^{26}
              t=P1Nt=\frac{P-1}{N} ,由于 N=2kN=2^kP=7×226+1P=7 \times 2^{26} +1 , 所以 tt 必是整数。
              g=3tg=3^t ,固有:g2=(3t)2,g3=(3t)3gN1=(3t)N1g^2=(3^t)^2,g^3=(3^t)^3 \dots g^{N-1}=(3^t)^{N-1}
              xx 取:g0,g1,g2,,gN1g^0,g^1,g^2,\dots,g^{N-1} ,此时有:gk=gk+N/2g^k=-g^{k+N/2}
              说明:gN=(3t)N=3P1=1g^N=(3^t)^N=3^{P-1}=1,又因 3i3^i 两两不同,故3t)N/2=1(3^t)^{N/2}=-1,而不可能等于 11 。(模意义下的 1-1 等于 P1P-1

            3. 计算 A(x)A(x)

            假设 N=8 N=8 t=P1N=4697620488=58720256 t=\frac{P-1}{N}=\frac{469762048}{8}=58720256
            xx 取:$3^0,3^{58720256},3^{58720256 \times 2},3^{58720256 \times 3},3^{58720256 \times 4},3^{58720256 \times 5},3^{5872025 \times 6},3^{58720256 \times 7}$
            xx 取:g0,g1,g2,,gN1g^0 , g^1 , g^2 , \dots , g^{N-1}
            具体计算 A(x)A(x) 的值如下:
            $Y_0=A(x_0)=A(g^0)= a_0 + a_1 + a_2 + a_3 + a_4 + a_5 + a_6 + a_7$
            $Y_1=A(x_1)=A(g^1)= a_0 + a_1g + a_2g^2 + a_3g^3 + a_4g^4 + a_5g^5 + a_6g^6 + a_7g^7$
            $Y_2=A(x_2)=A(g^2)= a_0 + a_1g^2 + a_2g^4 + a_3g^6 + a_4g^8 + a_5g^{10} + a_6g^{12} + a_7g^{14}$
            $Y_3=A(x_3)=A(g^3)= a_0 + a_1g^3 + a_2g^6 + a_3g^9 + a_4g^{12} + a_5g^{15} + a_6g^{18} + a_7g^{21}$
            $Y_4=A(x_4)=A(g^4)= a_0 + a_1g^4 + a_2g^8 + a_3g^{12} + a_4g^{16} + a_5g^{20} + a_6g^{24} + a_7g^{28}$
            $Y_5=A(x_5)=A(g^5)= a_0 + a_1g^5 + a_2g^{10} + a_3g^{15} + a_4g^{20} + a_5g^{25} + a_6g^{30} + a_7g^{35}$
            $Y_6=A(x_6)=A(g^6)= a_0 + a_1g^6 + a_2g^{12} + a_3g^{18} + a_4g^{24} + a_5g^{30} + a_6g^{36} + a_7g^{42}$
            $Y_7=A(x_7)=A(g^7)= a_0 + a_1g^7 + a_2g^{14} + a_3g^{21} + a_4g^{28} + a_5g^{35} + a_6g^{42} + a_7g^{49}$
            因:$A(x)= a_0 + a_1x + a_2x^2 + a_3x^3 + a_4x^4 + a_5x^5 + a_6x^6 + a_7x^7$
            把偶数项和奇数项分开如下:
            $A(x)= (a_0 + a_2x^2 + a_4x^4 + a_6x^6) + (a_1 + a_3x^2 + a_5x^4 + a_7x^6) \times x$
            设:A1(x2)=a0+a2x2+a4x4+a6x6A_1(x^2)= a_0 + a_2x^2 + a_4x^4 + a_6x^6A2(x2)=a1+a3x2+a5x4+a7x6A_2(x^2)= a_1 + a_3x^2 + a_5x^4 + a_7x^6
            则有:A(x)=A1(x2)+A2(x2)×xA(x)=A_1(x^2)+A_2(x^2) \times x ,具体Y0Y7如下Y_0 \dots Y_7如下
            $Y_0=A(x_0)=A(g^0)= A_1((g^0)^2)+A_2((g^0)^2) \times g^0$
            $Y_1=A(x_1)=A(g^1)= A_1((g^1)^2)+A_2((g^1)^2) \times g^1$
            $Y_2=A(x_2)=A(g^2)= A_1((g^2)^2)+A_2((g^2)^2) \times g^2$
            $Y_3=A(x_3)=A(g^3)= A_1((g^3)^2)+A_2((g^3)^2) \times g^3$
            $Y_4=A(x_4)=A(g^4)= A_1((g^4)^2)+A_2((g^4)^2) \times g^4$
            $Y_5=A(x_5)=A(g^5)= A_1((g^5)^2)+A_2((g^5)^2) \times g^5$
            $Y_6=A(x_6)=A(g^6)= A_1((g^6)^2)+A_2((g^6)^2) \times g^6$
            $Y_7=A(x_7)=A(g^7)= A_1((g^7)^2)+A_2((g^7)^2) \times g^7$
            代码如下:

            for(LL i=0,xi=1;i<N;++i,xi=xi*x%P)
            {
                A(i)=(A1(i)+A2(i)*xi)%P;
            }
            

            注:以上程序块中, A(i)A(i) 表示 A(gi)A(g^i) ,而 A1(i)A1(i) 表示 A1(gi2)A_1({g^i}^2), A2(i)A2(i) 表示 A2(gi2)A_2({g^i}^2)
            考虑周期性:g0=g8,g2=g10,g4=g12,g3=g11g^0=g^8,g^2=g^{10},g^4=g^{12},g^3=g^{11}
            考虑对称性:g0=g4,g1=g5,g2=g6,g3=g7g^0=-g^4,g^1=-g^5,g^2=-g^6,g^3=-g^7
            则有:
            Y0=A(x0)=A(g0)=A1(g0)+A2(g0)×g0Y_0=A(x_0)=A(g^0)= A_1(g^0)+A_2(g^0) \times g^0
            Y1=A(x1)=A(g1)=A1(g2)+A2(g2)×g1Y_1=A(x_1)=A(g^1)= A_1(g^2)+A_2(g^2) \times g^1
            Y2=A(x2)=A(g2)=A1(g4)+A2(g4)×g2Y_2=A(x_2)=A(g^2)= A_1(g^4)+A_2(g^4) \times g^2
            Y3=A(x3)=A(g3)=A1(g6)+A2(g6)×g3Y_3=A(x_3)=A(g^3)= A_1(g^6)+A_2(g^6) \times g^3
            Y4=A(x4)=A(g4)=A1(g0)A2(g0)×g0Y_4=A(x_4)=A(g^4)= A_1(g^0)-A_2(g^0) \times g^0
            Y5=A(x5)=A(g5)=A1(g2)A2(g2)×g1Y_5=A(x_5)=A(g^5)= A_1(g^2)-A_2(g^2) \times g^1
            Y6=A(x6)=A(g6)=A1(g4)A2(g4)×g2Y_6=A(x_6)=A(g^6)= A_1(g^4)-A_2(g^4) \times g^2
            Y7=A(x7)=A(g7)=A1(g6)A2(g6)×g3Y_7=A(x_7)=A(g^7)= A_1(g^6)-A_2(g^6) \times g^3

            对比项目: A(gk)A(g^k)A(gk+N/2)A(g^{k+N/2}) ,
            0k<N/20 \leq k <N/2
            表达式的左边的变量有:g0,g1,g2,g3,g4,g5,g6,g7g^0,g^1,g^2,g^3,g^4,g^5,g^6,g^7
            表达式的右边的变量只有: g0,g1,g2,g3g^0,g^1,g^2,g^3
            对比 A(g0)A(g^0)A(g4)A(g^4) A(g0)=A1((g0)2)+A2((g0)2)×g0A(g^0)= A_1((g^0)^2)+A_2((g^0)^2) \times g^0
            A(g4)=A1((g0)2)A2((g0)2)×g0A(g^4)= A_1((g^0)^2)-A_2((g^0)^2) \times g^0
            对比 A(g1)A(g^1)A(g5)A(g^5) A(g1)=A1((g1)2)+A2((g1)2)×g1A(g^1)= A_1((g^1)^2)+A_2((g^1)^2) \times g^1
            A(g5)=A1((g1)2)A2((g1)2)×g1A(g^5)= A_1((g^1)^2)-A_2((g^1)^2) \times g^1
            对比 A(g2)A(g^2)A(g6)A(g^6) A(g2)=A1((g2)2)+A2((g2)2)×g2A(g^2)= A_1((g^2)^2)+A_2((g^2)^2) \times g^2
            A(g6)=A1((g2)2)A2((g2)2)×g2A(g^6)= A_1((g^2)^2)-A_2((g^2)^2) \times g^2
            对比 A(g3)A(g^3)A(g7)A(g^7) A(g3)=A1((g3)2)+A2((g3)2)×g3A(g^3)= A_1((g^3)^2)+A_2((g^3)^2) \times g^3
            A(g7)=A1((g3)2)A2((g3)2)×g3A(g^7)= A_1((g^3)^2)-A_2((g^3)^2) \times g^3

            发现只要枚举 x=g0,g1,g2,g3x=g^0,g^1,g^2,g^3 ,就可以获得: A(g0),A(g1),,A(g7)A(g^0),A(g^1), \dots ,A(g^7) 代码中循环次数可以减半,如下:

            for(LL i=0,xi=1;i<N/2;++i,xi=xi*x%P)
            {
                A(i)    =(A1(i)+A2(i)*xi)%P;
                A(i+N/2)=(A1(i)-A2(i)*xi+P)%P;
            }
            

            A(x0),A(x1),A(x2),,A(xN1)A(x_0),A(x_1),A(x_2), \dots ,A(x_{N-1}) ,可在 O(nlogn)O(nlogn) 内解决。
            同样可得:B(x0),B(x1),B(x2),,B(xN1)B(x_0),B(x_1),B(x_2), \dots ,B(x_{N-1})
            以及: C(x0),C(x1),C(x2),,C(xN1)C(x_0),C(x_1),C(x_2), \dots ,C(x_{N-1})
            到此,NTT 算法已经学会一半:在 O(nlogn)O(nlogn) 内得到C(x)C(x)的N个不同的点值。另外一半:已知N个点值反推 C(x)C(x) 的各项系数。

            四. 已知多项式的点值求系数

            已知点值表示:$C(x)=\{(x_0,y_0),(x_1,y_1), \dots ,(x_{N-1},y_{N-1})\}$,求C(x)=c0+c1x+c2x2++cN1xN1C(x)=c_0+c_1 * x+c_2 * x^2+\dots+c_{N-1} * x^{N-1}的各项系数。

            1. 把 C(x)C(x)yy 值作为新的多项式C(x)C'(x)的系数

            $C'(x)=y_0+y_1 * x+y_2 * x^2+ \dots +y_{N-1} * x^{N-1}=\sum\limits_{i=0}^{N-1}y_i * x^i$

            2. xx 取:g0,g1,g2,,g(N1)g^{-0},g^{-1},g^{-2}, \dots ,g^{-(N-1)}代入 C(x)C'(x) 得到 NN 个新点值:y0,y1,y2,,yN1y'_0,y'_1,y'_2, \dots ,y'_{N-1}

            3. 分析 yk y'_k 有惊喜: ck=ykN=ykN1 c_k=\frac{y'_k}{N}=y'_k * N^{-1}

            $y'_k=C'(g^{-k})=\sum\limits_{i=0}^{N-1}y_i*(g^{-k})^i$--1式
            又因:yi=C(gi)=j=0N1cj(gi)jy_i=C(g^i)=\sum\limits_{j=0}^{N-1}c_j*(g^i)^j--2式
            把2式的yiy_i代入1式,得到:
            $y'_k=C'(g^{-k})=\sum\limits_{i=0}^{N-1}\sum\limits_{j=0}^{N-1}c_j*(g^i)^j*(g^{-k})^i$
                $=\sum\limits_{i=0}^{N-1}\sum\limits_{j=0}^{N-1}c_j * g^{i*j}*g^{-k*i}$
                $=\sum\limits_{i=0}^{N-1}\sum\limits_{j=0}^{N-1}c_j * g^{i*(j-k)}$
                $=\sum\limits_{i=0}^{N-1}\sum\limits_{j=0}^{N-1}c_j * (g^{(j-k)})^i$
                =j=0N1cj=\sum\limits_{j=0}^{N-1}c_j * i=0N1(g(jk))i\sum\limits_{i=0}^{N-1}(g^{(j-k)})^i
            先分析 i=0N1(g(jk))i\sum\limits_{i=0}^{N-1}(g^{(j-k)})^i
            j=kj=k时,i=0N1(g(jk))i\sum\limits_{i=0}^{N-1}(g^{(j-k)})^i $=\sum\limits_{i=0}^{N-1}(g^0)^i=\sum\limits_{i=0}^{N-1}1=N$
            jkj \ne k时,i=0N1(g(jk))i\sum\limits_{i=0}^{N-1}(g^{(j-k)})^i $=\frac{(g^{(j-k)})^N - 1}{g^{(j-k)} - 1}=\frac{(g^N)^{(j-k)} - 1}{g^{(j-k)} - 1}=\frac{1^{(j-k)} - 1}{g^{(j-k)} - 1}=0$
            故:yk=j=0N1cjy'_k=\sum\limits_{j=0}^{N-1}c_j * i=0N1(g(jk))i\sum\limits_{i=0}^{N-1}(g^{(j-k)})^i
                    $=c_0*0+c_1*0+c_2*0+ \dots +c_k*N+ \dots +c_{N-2}*0+c_{N-1}*0$
                    =ckN=c_k*N
            因: yk=ckNy'_k=c_k*N ,可得: ck=ykN=ykN1c_k=\frac{y'_k}{N}=y'_k * N^{-1}
            也即是:只要求出y0,y1,y2yN1y'_0,y'_1,y'_2 \dots y'_{N-1} ,根据ck=ykNc_k=\frac{y'_k}{N},可以轻松得到 C(x)C(x) 的各项系数 c0,c1,c2,,cN1c_0,c_1,c_2, \dots ,c_{N-1}
            实际上 C(x)C(x) 的系数只有 n+m1n+m-1 项,Nn+m1N \geq n+m-1 ,多出的系数为0(ci=0,i[n+m,N1]c_i=0,i \in [n+m,N-1])。

            五. 代码与改进

            1. 题目:LOJ108多项式乘法

            2. 代码

            代码1:数论变换(递归版,105ms)

            #include<bits/stdc++.h>//快速数论变换(递归版)(NTT)
            #define LL long long
            const int M=4e5+10;//要开4倍
            const LL P=(7ll<<26)+1;//(119ll<<23)+1;//(7ll<<26)+1;//(17ll<<27+1);
            LL qpow(LL a,LL b){LL res=1;for(;b;b>>=1,a=a*a%P)if(b&1)res=res*a%P;return res;}
            LL A[M],B[M],C[M];
            void NTT(LL A[],LL n,LL x)
            {
                if(n==1)return;//省略A[0]=A[0],实际是Y[0]=A[0] 
                LL A1[n/2],A2[n/2];//A1,A2必须在函数内定义 
                for(LL i=0;i<n/2;++i)A1[i]=A[2*i],A2[i]=A[2*i+1];
                NTT(A1,n/2,x*x%P);
                NTT(A2,n/2,x*x%P);
                for(LL i=0,xi=1;i<n/2;++i,xi=xi*x%P)
                {
                    A[i]     =  (A1[i] + A2[i]*xi) % P;
                    A[i+n/2] = ((A1[i] - A2[i]*xi) % P + P) % P;
                }
            }
            int main()
            {
                LL n, m;scanf("%lld%lld", &n, &m);n++;m++;
                for(LL i=0;i<n;i++)scanf("%lld", &A[i]);
                for(LL i=0;i<m;i++)scanf("%lld", &B[i]);
                LL N=1;while(N<n+m-1)N<<=1;LL inv_N=qpow(N,P-2);
                LL x=qpow(3ll,(P-1)/N);
                NTT(A,N,x);
                NTT(B,N,x);
                for(LL i=0;i<N;++i)C[i]=A[i]*B[i]%P;
                LL inv_x=qpow(x,P-2);
                NTT(C,N,inv_x);
                for(LL i=0;i<n+m-1;++i)printf("%lld ",C[i]*inv_N% P);
                return 0;
            }
            

            代码2:快速数论变换(非递归版,90ms)

            #include <bits/stdc++.h>//快速数论变换(非递归版:常用)(NTT)
            #define LL long long
            using namespace std;
            const int M=4e5+10;
            const LL P=(7ll<<26)+1;
            LL qpow(LL a,LL b){LL res=1;for(;b;b>>=1,a=a*a%P)if(b&1)res=res*a%P;return res;}
            LL A[M],B[M],C[M],r[M];
            void NTT(LL A[],LL n,LL x)
            {
                for(int i=0;i<n;++i)if(i<r[i])swap(A[i],A[r[i]]);
                for(LL m=2;m<=n;m<<=1)
                {
                    LL xm=qpow(x,n/m);
                    for(LL i=0;i<n;i+=m)
                    {
                        for(LL j=0,xj=1;j<m/2;++j,xj=xj*xm%P)
                        {
                            LL t1=A[i+j],t2=A[i+j+m/2]*xj%P;
                            A[i+j]    =(t1+t2)%P;
                            A[i+j+m/2]=(t1-t2+P)%P;
                        }
                    }
                }
            }
            int main()
            {
                int n,m;scanf("%d%d",&n,&m);n++;m++;
                for(int i=0;i<n;i++)scanf("%lld",&A[i]);
                for(int i=0;i<m;i++)scanf("%lld",&B[i]);
                LL N=1;while(N<n+m-1)N<<=1;LL inv_N=qpow(N, P-2);
                for(LL i=0; i<N;++i)r[i]=r[i/2]/2+(i&1)*N/2;
                LL x=qpow(3ll,(P-1)/N);
                NTT(A,N,x);
                NTT(B,N,x);
                for(LL i=0;i<N;++i)C[i]=A[i]*B[i]%P;
                LL inv_x=qpow(x,P-2);
                NTT(C,N,inv_x);
                for(LL i=0;i<n+m-1;++i)printf("%lld ",C[i]*inv_N%P);
                return 0;
            }
            
            • 1

            信息

            ID
            564
            时间
            1000ms
            内存
            256MiB
            难度
            8
            标签
            递交数
            329
            已通过
            38
            上传者