CC++ & Algorithm

多项式乘法逆与NTT的综合应用

极难6
语言版本:通用
概述:学会了NTT,我们就可以用它来实现多项式求逆、除法等高级运算,就像用万能工具做出更复杂的积木城堡。

多项式乘法逆与NTT:从零搭建高效计算积木

你是否想过,如果有一个“万能工具”,能让你快速算出多项式的逆、除法、开方等运算,就像用魔方公式还原魔方一样?这个工具就是 NTT(数论变换),而今天我们要用它来实现一个超级实用的功能——多项式求逆。学会了它,你就能像搭积木一样,轻松搞定更复杂的数学计算。

什么是多项式乘法逆?

想象一下,你在做数学题时,老师要求你找到一个数 bb,使得 a×b=1a \times b = 1,这个 bb 就是 aa 的倒数(乘法逆元)。对于多项式来说,道理类似:

对于多项式 A(x)A(x),如果存在另一个多项式 B(x)B(x),使得它们的乘积在 xnx^n 意义下等于常数1,即:

A(x)B(x)1(modxn)A(x) \cdot B(x) \equiv 1 \pmod{x^n}

那么 B(x)B(x) 就是 A(x)A(x) 在模 xnx^n 下的 乘法逆元。这里的“模 xnx^n”是什么意思呢?简单说,就是只保留乘积中次数小于 nn 的项,更高次的项全部丢掉。就像你只关心前几个数字,后面的忽略不计。

举个生活中的例子

假设你每个月的零花钱是 A=1+2xA = 1 + 2x,其中 xx 代表月份(比如 xx 表示第一个月、第二个月...)。你想找一个“反向公式” B(x)B(x),使得 A×BA \times B 的结果恰好等于1(常数),而且只计算前两个月的效果(因为模 x2x^2)。经过尝试,你发现 B=12xB = 1 - 2x 就满足:

(1+2x)×(12x)=14x2(1+2x) \times (1-2x) = 1 - 4x^2

因为 x2x^2 项已经超出“模 x2x^2”的范围(我们只关心 x0x^0x1x^1 项),所以结果就是1,完美!所以 B(x)=12xB(x) = 1 - 2x 就是 A(x)A(x) 的逆元。

注意:求逆的前提是常数项不能为0(否则没有逆元,就像0没有倒数一样)。

核心原理:牛顿迭代法——从粗糙到精确

你可能觉得,直接算逆元很麻烦,尤其是多项式很长的时候。但别怕,我们可以用 牛顿迭代法,它是一种逐步逼近的方法,就像你玩射击游戏时,先瞄准一个大概的方向,然后不断修正,最终精准命中。

数学公式(不怕,我们用人话解释)

假设我们已经知道了 A(x)A(x) 在模 xn/2x^{n/2} 下的逆元 B0(x)B_0(x)(精度较低),想要求模 xnx^n 下的更精确逆元 B(x)B(x)。牛顿迭代给出了一个公式:

B=B0(2AB0)(modxn)B = B_0 \cdot (2 - A \cdot B_0) \pmod{x^n}

这个公式可以这样理解:

  • 先计算 AB0A \cdot B_0,看它离1有多远(如果完美,结果应该等于1)。
  • 然后用 22 减去这个结果(因为 21=12 - 1 = 1,如果结果偏离,这个差会修正回去)。
  • 最后再乘上 B0B_0,就得到了更精确的逆元。

迭代过程:我们从常数项开始(模 x1x^1),只要常数项非零,它的逆元就是它的倒数(比如 a0a_0 的逆元是 a01a_0^{-1})。然后每次将精度加倍(从 x1x^1x2x^2,再到 x4x^4,直到 xnx^n),每次迭代用一次牛顿公式。每次迭代需要做两次多项式乘法,而乘法可以用NTT加速,所以总的时间复杂度是 O(nlogn)O(n \log n),相当快!

动手编程:用NTT实现多项式求逆

下面我们给出C++和Python两种语言的完整实现。代码中包含了NTT函数,以及求逆的核心逻辑。注意:代码中每个变量都加了中文注释,方便你理解每一步在做什么。

C++ 实现

#include <iostream>
#include <vector>
#include <algorithm>
using namespace std;
typedef long long ll;                // 长整型别名
const int MOD = 998244353;           // 常用模数,保证NTT可用
const int ROOT = 3;                  // 原根

// 快速幂:计算 x^e % MOD
ll qpow(ll x, ll e) {
    ll res = 1;
    while (e) {
        if (e & 1) res = res * x % MOD;  // 如果当前位为1,乘上x
        x = x * x % MOD;                 // x 自乘
        e >>= 1;                         // e 右移
    }
    return res;
}

// NTT:对向量 a 进行数论变换(invert=false为正变换,true为逆变换)
void ntt(vector<ll>& a, bool invert) {
    int n = a.size();
    // 位逆序重排
    for (int i = 1, j = 0; i < n; ++i) {
        int bit = n >> 1;
        for (; j & bit; bit >>= 1) j ^= bit;
        j ^= bit;
        if (i < j) swap(a[i], a[j]);
    }
    // 迭代合并
    for (int len = 2; len <= n; len <<= 1) {
        ll wlen = qpow(ROOT, (MOD - 1) / len); // 计算单位根
        if (invert) wlen = qpow(wlen, MOD - 2); // 逆变换用逆元
        for (int i = 0; i < n; i += len) {
            ll w = 1;
            for (int j = 0; j < len / 2; ++j) {
                ll u = a[i + j];
                ll v = a[i + j + len/2] * w % MOD;
                a[i + j] = (u + v) % MOD;
                a[i + j + len/2] = (u - v + MOD) % MOD;
                w = w * wlen % MOD;
            }
        }
    }
    // 如果是逆变换,每个元素要除以 n
    if (invert) {
        ll n_inv = qpow(n, MOD - 2);
        for (int i = 0; i < n; ++i) a[i] = a[i] * n_inv % MOD;
    }
}

// 多项式求逆:给定系数向量 a(常数项非零),返回模 x^n 下的逆元 b
// 返回的向量长度至少为 n
vector<ll> polyInv(const vector<ll>& a, int n) {
    vector<ll> inv_a(n, 0);           // 逆元系数,初始全0
    inv_a[0] = qpow(a[0], MOD - 2);  // 常数项的逆元

    // 当前精度 cur,从1开始,每次加倍直到达到 n
    for (int cur = 1; cur < n; cur <<= 1) {
        int new_cur = min(cur * 2, n);   // 下一轮需要达到的精度
        // 取 f 为 a 的前 new_cur 项,其余补0
        vector<ll> f(new_cur, 0);
        copy(a.begin(), a.begin() + min((int)a.size(), new_cur), f.begin());

        // 计算 f * inv_a,需要卷积长度至少 new_cur + cur - 1
        int size_fft = 1;
        while (size_fft < new_cur + cur) size_fft <<= 1; // 找到2的幂
        vector<ll> fa(size_fft, 0), fb(size_fft, 0);    // 两个工作数组
        copy(f.begin(), f.end(), fa.begin());             // 复制 f
        copy(inv_a.begin(), inv_a.begin() + cur, fb.begin()); // 复制当前逆元

        ntt(fa, false);  // 对 fa 做正变换
        ntt(fb, false);  // 对 fb 做正变换
        for (int i = 0; i < size_fft; ++i) fa[i] = fa[i] * fb[i] % MOD; // 点乘
        ntt(fa, true);   // 逆变换得到乘积

        // 现在 fa 的前 new_cur 项是 f * inv_a (模 x^{new_cur})
        // 计算 (2 - f * inv_a)
        for (int i = 0; i < new_cur; ++i) {
            fa[i] = (MOD - fa[i]) % MOD;   // 先取负
        }
        fa[0] = (fa[0] + 2) % MOD;         // 再加2,得到 2 - f*inv_a

        // 清空 fa 超过 new_cur 的部分,因为后续只用到前 new_cur
        fill(fa.begin() + new_cur, fa.end(), 0);
        // 再乘以 inv_a,得到 inv_a * (2 - f*inv_a)
        fill(fb.begin(), fb.end(), 0);
        copy(inv_a.begin(), inv_a.begin() + cur, fb.begin());

        ntt(fa, false);  // 正变换
        ntt(fb, false);
        for (int i = 0; i < size_fft; ++i) fa[i] = fa[i] * fb[i] % MOD;
        ntt(fa, true);

        // 复制结果到 inv_a 的前 new_cur 项
        for (int i = 0; i < new_cur; ++i) inv_a[i] = fa[i] % MOD;
    }
    return inv_a;
}

int main() {
    vector<ll> A = {1, 2, 3};          // 多项式 1 + 2x + 3x^2
    int n = 4;                         // 求模 x^4 下的逆元
    auto B = polyInv(A, n);
    cout << "A(x) = 1 + 2x + 3x^2, 求模 x^4 下的逆: ";
    for (int i = 0; i < n; ++i) {
        if (i > 0) cout << " + ";
        cout << B[i] << (i==0 ? "" : "x^" + to_string(i));
    }
    cout << endl;
    // 可以自行添加验证代码:计算 A*B mod x^4 是否等于1
    return 0;
}

Python 实现

MOD = 998244353   # 模数
ROOT = 3          # 原根

# 快速幂
def qpow(x, e):
    res = 1
    while e:
        if e & 1:
            res = res * x % MOD
        x = x * x % MOD
        e >>= 1
    return res

# NTT函数
def ntt(a, invert):
    n = len(a)
    j = 0
    # 位逆序重排
    for i in range(1, n):
        bit = n >> 1
        while j & bit:
            j ^= bit
            bit >>= 1
        j ^= bit
        if i < j:
            a[i], a[j] = a[j], a[i]
    # 迭代合并
    length = 2
    while length <= n:
        wlen = qpow(ROOT, (MOD - 1) // length)
        if invert:
            wlen = qpow(wlen, MOD - 2)
        for i in range(0, n, length):
            w = 1
            half = length // 2
            for j in range(i, i + half):
                u = a[j]
                v = a[j + half] * w % MOD
                a[j] = (u + v) % MOD
                a[j + half] = (u - v + MOD) % MOD
                w = w * wlen % MOD
        length <<= 1
    # 逆变换除以n
    if invert:
        n_inv = qpow(n, MOD - 2)
        for i in range(n):
            a[i] = a[i] * n_inv % MOD

# 多项式求逆:返回a在模 x^n 下的逆元(长度至少为n)
def poly_inv(a, n):
    inv = [0] * n                     # 逆元系数初始全0
    inv[0] = qpow(a[0], MOD - 2)      # 常数项逆元

    cur = 1
    while cur < n:
        new_cur = min(cur * 2, n)      # 新精度
        # 取 a 的前 new_cur 项作为 f,长度不够补0
        f = a[:new_cur] + [0] * (new_cur - len(a[:new_cur]))
        # 计算 f * inv,需要卷积长度至少 new_cur + cur - 1
        size_fft = 1
        while size_fft < new_cur + cur:
            size_fft <<= 1
        fa = f + [0] * (size_fft - len(f))
        fb = inv[:cur] + [0] * (size_fft - cur)
        ntt(fa, False)
        ntt(fb, False)
        for i in range(size_fft):
            fa[i] = fa[i] * fb[i] % MOD
        ntt(fa, True)
        # 此时 fa[0..new_cur-1] 是 f*inv 的系数
        # 计算 2 - f*inv
        for i in range(new_cur):
            fa[i] = (-fa[i]) % MOD
        fa[0] = (fa[0] + 2) % MOD
        # 清空 fa 超过 new_cur 的部分(其实后面会重新填充,但为了明确)
        # 再乘以 inv,得到 inv*(2 - f*inv)
        fb = inv[:cur] + [0] * (size_fft - cur)
        ntt(fa, False)
        ntt(fb, False)
        for i in range(size_fft):
            fa[i] = fa[i] * fb[i] % MOD
        ntt(fa, True)
        # 复制到 inv 的前 new_cur 项
        for i in range(new_cur):
            inv[i] = fa[i] % MOD
        cur = new_cur
    return inv[:n]

# 测试
A = [1, 2, 3]           # 多项式 1 + 2x + 3x^2
n = 4                    # 求模 x^4 下的逆元
inv_A = poly_inv(A, n)
print("A(x) = 1 + 2x + 3x^2")
print(f"模 x^{n} 下的逆元: ", end="")
terms = []
for i, c in enumerate(inv_A):
    if c == 0:
        continue
    if i == 0:
        terms.append(str(c))
    elif i == 1:
        terms.append(f"{c}x")
    else:
        terms.append(f"{c}x^{i}")
print(" + ".join(terms))

新手常见错误与避坑指南

  1. 常数项为0的情况
    如果多项式 A(x)A(x) 的常数项 a0=0a_0 = 0,那么它没有逆元(因为任何多项式乘以它后,常数项都是0,不可能得到1)。解决方法是先判断常数项是否非零,如果为零,需要先提取 xkx^k 因子(例如 A(x)=xA(x)A(x) = x \cdot A'(x)),再对 A(x)A'(x) 求逆。

  2. NTT长度选择不对
    牛顿迭代中,每次卷积需要长度为 new_cur+cur1new\_cur + cur - 1,我们向上取到2的幂。如果长度只取 new_curnew\_cur,会导致结果被截断,无法正确迭代。一定要确保卷积长度足够大。

  3. 忘记模运算
    所有中间结果都要取模,尤其是减法时要注意先加MOD再取模,防止负数出现。

  4. 逆变换后忘记除以n
    NTT逆变换后,每个元素需要乘以 nn 的逆元,否则结果会多一个因子 nn。检查代码中是否正确处理了 if invert 部分。

应用:多项式求逆能做什么?

学会了多项式求逆,就像拥有了一把万能钥匙,可以打开很多高级运算的大门:

  • 多项式除法:已知多项式 A(x)A(x)B(x)B(x),求商 Q(x)Q(x) 和余数 R(x)R(x),使得 A=BQ+RA = BQ + R。方法:先计算 BB 的逆元,然后 Q=AB1(modxdeg(A)deg(B)+1)Q = A \cdot B^{-1} \pmod{x^{deg(A)-deg(B)+1}}。这就像做整数除法时,用除数的倒数乘以被除数。

  • 多项式开方:求 C(x)C(x) 使得 C2A(modxn)C^2 \equiv A \pmod{x^n},也需要用到求逆(牛顿迭代法类似)。

  • 多项式对数、指数:在生成函数、组合数学中,经常需要计算 ln(A(x))\ln(A(x))eA(x)e^{A(x)},这些公式里都包含多项式求逆或乘法。

你可以把多项式求逆想象成乐高积木中的基础砖块,有了它,就能搭出更复杂的结构,比如在算法竞赛中解决快速计算、图像处理、信号处理等领域的问题。

总结

从多项式的基本概念出发,我们学习了用NTT加速多项式乘法,然后利用牛顿迭代法实现了多项式求逆。整个过程就像一步步升级装备:先学会两位数乘法,再学速算技巧,最后能轻松完成百位数乘法。现在你已经掌握了“多项式求逆”这个高级技能,接下来可以尝试挑战多项式除法、开方、指数等更酷的运算!

相关指引:如果你对FFT、NTT还不太熟悉,建议先复习一下“快速傅里叶变换”和“数论变换”的原理。之后可以学习“多项式乘法”、“多项式除法”等专题。加油,下一个编程数学高手就是你!

例题精讲

1单选题

使用NTT实现多项式乘法逆,若多项式次数为n,则求逆的渐近时间复杂度为?

AO(n log n)
BO(n log^2 n)
CO(n^2)
DO(n)
2判断题

在模数为998244353时,可以使用NTT进行多项式乘法,因为该模数是质数且存在原根3。

3填空题
给定多项式a,实现递归求逆。补全代码:void poly_inv(const vector<int>& a, vector<int>& b, int n) {
    if (n == 1) { b[0] = ___; return; }
    int m = (n+1)/2;
    poly_inv(a, b, m);
    // 后续利用NTT计算b = b*(2 - a*b) mod x^n
}
4单选题

模数998244353下,NTT支持的最大变换长度(2的幂次)为多少?

A2^23
B2^24
C2^25
D2^26
5判断题

利用牛顿迭代法求多项式乘法逆时,需要保证常数项a[0] ≠ 0。