【笔记】多项式 GCD 与 Half-GCD 算法

· · 算法·理论

:::info[AI 使用说明] 本文的文章润色与卡常工作由 AI 完成。 :::

怎么全世界都会多项式了。

那就做个笔记记录一下。

多项式 GCD 问题

给定两个度数分别为 n,m 的多项式 f(x)=a_0+a_1x+a_2x^2+\ldots+a_nx^n,g(x)=b_0+b_1x+b_2x^2+\ldots+b_mx^m

求出一个多项式 h(x)=c_0+c_1x+c_2x^2+\ldots+c_kx^k,使得 hf,g 的因式且 k 尽量大,如果有多个结果,则任意一个均可以。

那么该怎么做呢?

暴力

显然,我们很容易想到欧几里得算法,它在多项式 gcd 上仍然是对的。

但是,它的时间复杂度是多少?

众所周知,多项式取模可以做到 O(n \log n)

但是,我们发现,做一次取模,最坏情况下只会让多项式的度数减 O(1)

比如 f(x)=x^5+2x^4+x^3+5x^2-x+6,g(x)=x^4,这个时候 f \bmod g=x^3+5x^2-x+6,度数只减了 2
因此,我们最坏情况下会做 O(n) 次取模。(而且不需要精心构造,随机数据都容易卡到)
所以,欧几里得算法计算多项式 gcd 的最坏时间复杂度为 O(n^2\log n),非常的不优秀。

而且我们很容易想到一个更优的做法,因为对于 f,gf 的度大于等于 g)我们肯定能够找到一个 c 使得 f 减去 cg 后度数至少减一!这样我们就只需要做 O(n) 次多项式减法,时间复杂度变成了 O(n^2)

那么,该如何突破 O(n^2) 的壁垒呢?

Half-GCD

Half-GCD 提供了一种在 O(n \log^2 n) 的时间复杂度下计算两个多项式的 gcd(exgcd)的方法。(据说还可以改造成 O(n \log^2 n) 的最短线性递推式算法)

我们注意到一个事实,这一切的问题都在于取模是 O(\max(n+m)\log\max(n+m)) 的,即使这样,答案也只能缩小到 O(\min(n,m)) 级别。

如果我们能够像除法一样可以做到 O((n-m)\log(n-m)) 的时间复杂度,那么每在计算上花费的时间,都会转换成数据规模的下降,这样的话,复杂度就可以均摊掉了!

这使我们回顾欧几里得算法的过程。

具体地,我们构造了一个多项式序列 r_{-2},r_{-1},r_0,r_1,\ldots r_k,其中:

而这个可以写成矩阵的形式:

\begin{bmatrix}0&1\\1&-q_i\end{bmatrix}\begin{bmatrix}r_{i-2}\\r_{i-1}\end{bmatrix}=\begin{bmatrix}r_{i-1}\\r_{i}\end{bmatrix}

看上去非常优美!于是我们定义:

[Q]=\begin{bmatrix}Q&1\\1&0\end{bmatrix},[Q]^{-1}=\begin{bmatrix}0&1\\1&-Q\end{bmatrix}

然后就可以:

[q_k]^{-1}[q_{k-1}]^{-1}\ldots[q_1]^{-1}[q_0]^{-1}\begin{bmatrix}f\\g\end{bmatrix}=\begin{bmatrix}\gcd(f,g)\\0\end{bmatrix}

因此我们只需要计算 [q_k][q_{k-1}]\ldots[q_1][q_0] 即可,记它为 M^{-1},就可以快速算出答案,这就是 Half-GCD 的思想。

但是直接算整个 M 还是难度太大了,于是我们选择退而求其次,尝试计算出一个 M 使得 M^{-1}\begin{bmatrix}A\\B\end{bmatrix}=\begin{bmatrix}C\\D\end{bmatrix}\deg(C) \ge \dfrac{\deg(A)}{2} > \deg(D),然后再计算 E=C \bmod D,这个时候 \dfrac{\deg(A)}{2} > \deg(D) > \deg(E),也就是说,我们将问题的规模减半了!

好的,现在的问题是,怎么算?

具体过程

前提

定义函数 \operatorname{HalfGCD}(A, B),它接受两个多项式 A, B\deg A \ge \deg B),返回一个矩阵 M,使得:

M^{-1} \begin{bmatrix} A \\ B \end{bmatrix} = \begin{bmatrix} C \\ D \end{bmatrix}

且满足 规模减半条件

\deg C \ge \frac{\deg A}{2} > \deg D

选定分治点 m

取:

m = \left\lfloor \frac{\deg A}{2} \right\rfloor

A, Bx^m 拆分为高位和低位:

A = x^m A_1 + A_0, \quad B = x^m B_1 + B_0

其中:

计算:

M_1 = \operatorname{HalfGCD}(A_1, B_1)

B_1 = 0,则直接构造单位矩阵或做特殊处理。通常要求 A_1, B_1 满足递归前提(\deg A_1 \ge \deg B_1,否则交换或补零)。

于是有:

M_1^{-1} \begin{bmatrix} A_1 \\ B_1 \end{bmatrix} = \begin{bmatrix} C_1 \\ D_1 \end{bmatrix}, \quad \deg C_1 \ge \frac{\deg A_1}{2} > \deg D_1

M_1 作用到原始 A, B

由于 A = x^m A_1 + A_0,且矩阵乘法是线性的,得到:

M_1^{-1} \begin{bmatrix} A \\ B \end{bmatrix} = \begin{bmatrix} A' \\ B' \end{bmatrix}

具体展开(以第一行为例):

\begin{aligned} A'&=(M_1^{-1})_{11}A+(M_1^{-1})_{12}B\\ &= x^m ((M_1^{-1})_{11} A_1 + (M_1^{-1})_{12} B_1+ ((M_1^{-1})_{11} A_0 + (M_1^{-1})_{12} B_0)\\ &= x^m C_1 + A_0' \end{aligned}

同理:

B' = x^m D_1 + B_0'

此时 (A', B') 的次数相比原始 A, B 并没有显著下降,因为 A' 仍由 x^m C_1 主导。

少量调整步骤

虽然 (A', B') 次数未降,但 C_1 的次数远大于 D_1,这意味着高次项已经“对齐”,可以通过少量的经典除法步骤快速降低次数。

计算:

A' = Q B' + R, \quad \deg R < \deg B'

由于 \deg D_1 远小于 \deg C_1,可以证明 Q 的次数很小,因此这个除法是快速的。

于是得到新的矩阵变换:

\begin{bmatrix} 0 & 1 \\ 1 & -Q \end{bmatrix} \begin{bmatrix} A' \\ B' \end{bmatrix} = \begin{bmatrix} B' \\ R \end{bmatrix}

合并答案

将所有变换合成总矩阵 M

M^{-1} = \begin{bmatrix} 0 & 1 \\ 1 & -Q \end{bmatrix} \cdot M_1^{-1}

最终得到:

M^{-1} \begin{bmatrix} A \\ B \end{bmatrix} = \begin{bmatrix} B' \\ R \end{bmatrix}

时间复杂度

假设我们的时间复杂度为 T(n),观察过程容易发现我们将问题分成了两个子问题,也就是 2T\left(\dfrac{n}{2}\right),每合并两个子问题显然是 O(n \log n) 的,因此有:

T(n)=2T\left(\dfrac{n}{2}\right)+O(n \log n)

主定理一下就可以得到时间复杂度为 O(n \log^2 n)

实现

什么,居然没有模板题?!

我出的模板题。

结果写起来才发现这个东西有多难写,常数又大又难看。

只好借助 AI 才在 n,m \le 10^5 的时候卡到了五秒内。

所以有没有多项式大佬能继续卡常。

这个是 AI 卡常出的完整代码:

#include<bits/stdc++.h>
using namespace std;
const int MOD=998244353,G=3,Gi=332748118;
const int MAXL=1<<19;
inline int qpow(int a,int b){int res=1;for(;b;b>>=1,a=1LL*a*a%MOD)if(b&1)res=1LL*res*a%MOD;return res;}
namespace NTT{
    vector<int>roots[2];
    bool ntt_ready=false;
    void init_ntt(){
        if(ntt_ready)return;
        ntt_ready=true;
        roots[0].resize(MAXL);
        roots[1].resize(MAXL);
        int wn_fwd=qpow(G,(MOD-1)/MAXL);
        int wn_inv=qpow(Gi,(MOD-1)/MAXL);
        roots[0][0]=1;roots[1][0]=1;
        for(int i=1;i<MAXL;++i){
            roots[0][i]=1LL*roots[0][i-1]*wn_fwd%MOD;
            roots[1][i]=1LL*roots[1][i-1]*wn_inv%MOD;
        }
    }
    int rev[MAXL];
    void ntt(vector<int>&a,bool inv){
        init_ntt();
        int n=a.size();
        for(int i=0;i<n;++i)if(i<rev[i])swap(a[i],a[rev[i]]);
        for(int len=2;len<=n;len<<=1){
            int step=MAXL/len;
            for(int i=0;i<n;i+=len){
                for(int j=0;j<len/2;++j){
                    int w=roots[inv][j*step];
                    int u=a[i+j];
                    int v=1LL*a[i+j+len/2]*w%MOD;
                    a[i+j]=u+v;
                    if(a[i+j]>=MOD)a[i+j]-=MOD;
                    a[i+j+len/2]=u-v;
                    if(a[i+j+len/2]<0)a[i+j+len/2]+=MOD;
                }
            }
        }
        if(inv){
            int invn=qpow(n,MOD-2);
            for(int&x:a)x=1LL*x*invn%MOD;
        }
    }
    vector<int>mul(vector<int>a,vector<int>b){
        int n=1,len=0;
        int sz=a.size()+b.size()-1;
        while(n<sz)n<<=1,++len;
        for(int i=0;i<n;++i)rev[i]=(rev[i>>1]>>1)|((i&1)<<(len-1));
        a.resize(n,0);b.resize(n,0);
        ntt(a,false);ntt(b,false);
        for(int i=0;i<n;++i)a[i]=1LL*a[i]*b[i]%MOD;
        ntt(a,true);
        a.resize(sz);
        return a;
    }
}
using Poly=vector<int>;
inline bool isZero(const Poly&p){return p.empty()||(p.size()==1&&p[0]==0);}
inline bool isOne(const Poly&p){return p.size()==1&&p[0]==1;}
Poly operator+(Poly a,const Poly&b){
    if(a.size()<b.size())a.resize(b.size());
    for(size_t i=0;i<b.size();++i){
        a[i]+=b[i];
        if(a[i]>=MOD)a[i]-=MOD;
    }
    return a;
}
Poly operator-(Poly a,const Poly&b){
    if(a.size()<b.size())a.resize(b.size());
    for(size_t i=0;i<b.size();++i){
        a[i]-=b[i];
        if(a[i]<0)a[i]+=MOD;
    }
    return a;
}
Poly operator*(Poly a,int k){
    if(k==1)return a;
    if(k==0)return{0};
    for(int&x:a)x=1LL*x*k%MOD;
    return a;
}
Poly operator*(const Poly&a,const Poly&b){
    if(isZero(a)||isZero(b))return{0};
    if(isOne(a))return b;
    if(isOne(b))return a;
    if(a.size()<=64&&b.size()<=64){
        int n=a.size(),m=b.size();
        Poly c(n+m-1,0);
        for(int i=0;i<n;++i){
            long long ai=a[i];
            for(int j=0;j<m;++j){
                c[i+j]=(c[i+j]+ai*b[j])%MOD;
            }
        }
        return c;
    }
    return NTT::mul(a,b);
}
Poly inv(const Poly&a,int n){
    Poly res={qpow(a[0],MOD-2)};
    int cur=1;
    while(cur<n){
        cur<<=1;
        Poly f(a.begin(),a.begin()+min((int)a.size(),cur));
        Poly r=res*res*f;r.resize(cur);
        res=res*2-r;
    }
    res.resize(n);
    return res;
}
Poly rev(Poly a){reverse(a.begin(),a.end());return a;}
Poly div(Poly a,Poly b){
    int n=a.size()-1,m=b.size()-1;
    if(n<m)return{0};
    a=rev(a);b=rev(b);
    a.resize(n-m+1);b.resize(n-m+1);
    Poly invb=inv(b,n-m+1);
    Poly q=a*invb;q.resize(n-m+1);
    return rev(q);
}
Poly mod(const Poly&a,const Poly&b){
    if(a.size()<b.size())return a;
    Poly q=div(a,b);
    Poly r=a-b*q;
    while(r.size()>1&&r.back()==0)r.pop_back();
    return r;
}
int deg(const Poly&a){return(int)a.size()-1;}
pair<Poly,Poly>matMul(const array<Poly,4>&M,const Poly&a,const Poly&b){
    Poly a2=M[0]*a+M[1]*b;
    Poly b2=M[2]*a+M[3]*b;
    while(!a2.empty()&&a2.back()==0)a2.pop_back();
    while(!b2.empty()&&b2.back()==0)b2.pop_back();
    if(a2.empty())a2={0};
    if(b2.empty())b2={0};
    return{a2,b2};
}
array<Poly,4>matMul(const array<Poly,4>&A,const array<Poly,4>&B){
    return{A[0]*B[0]+A[1]*B[2],
            A[0]*B[1]+A[1]*B[3],
            A[2]*B[0]+A[3]*B[2],
            A[2]*B[1]+A[3]*B[3]};
}
array<Poly,4>hgcd(Poly a,Poly b){
    int n=deg(a),m_deg=deg(b);
    if(m_deg<=n/2)return{Poly{1},Poly{0},Poly{0},Poly{1}};
    int m=(n+1)/2;
    Poly a0,b0;
    if(deg(a)>=m)a0=Poly(a.begin()+m,a.end());
    else a0={0};
    if(deg(b)>=m)b0=Poly(b.begin()+m,b.end());
    else b0={0};
    auto M1=hgcd(a0,b0);
    tie(a,b)=matMul(M1,a,b);
    if(deg(b)<m)return M1;
    Poly q=div(a,b),r=mod(a,b);
    auto c=b,d=r;
    if(deg(d)<m){
        array<Poly,4>M2={Poly{0},Poly{1},Poly{1},Poly{0}-q};
        return matMul(M2,M1);
    }
    Poly c0,d0;
    if(deg(c)>=m)c0=Poly(c.begin()+m,c.end());
    else c0={0};
    if(deg(d)>=m)d0=Poly(d.begin()+m,d.end());
    else d0={0};
    auto M2=hgcd(c0,d0);
    array<Poly,4>Q={Poly{0},Poly{1},Poly{1},Poly{0}-q};
    return matMul(M2,matMul(Q,M1));
}
Poly gcd(Poly a,Poly b){
    while(!b.empty()&&!(b.size()==1&&b[0]==0)){
        if(deg(a)<deg(b))swap(a,b);
        auto M=hgcd(a,b);
        tie(a,b)=matMul(M,a,b);
        if(!b.empty()&&!(b.size()==1&&b[0]==0)){
            Poly r=mod(a,b);
            a=b;b=r;
        }
    }
    return a;
}
Poly normalize(Poly a){
    while(!a.empty()&&a.back()==0)a.pop_back();
    if(a.empty())return{0};
    int lead=a.back();
    if(lead!=1){
        int inv=qpow(lead,MOD-2);
        for(int&x:a)x=1LL*x*inv%MOD;
    }
    return a;
}
int main(){
    ios::sync_with_stdio(false);cin.tie(0);
    int n,m;cin>>n>>m;
    Poly f(n+1),g(m+1);
    for(int i=0;i<=n;++i)cin>>f[i];
    for(int i=0;i<=m;++i)cin>>g[i];
    Poly ans=normalize(gcd(f,g));
    if(ans.size()==1){
        cout<<"0\n1";
        return 0;
    }
    cout<<ans.size()-1<<'\n';
    for(int x:ans)cout<<x<<' ';
    return 0;
}