CC++ & Algorithm

NTT 数论变换

极难3
语言版本:通用
概述:NTT是FFT在整数模世界里的“孪生兄弟”,用原根代替复数单位根,完全避免浮点误差,适合计算精确的多项式乘法。

NTT 数论变换——整数世界里的“超快速”多项式乘法

你是不是遇到过这样的问题:两个很大的整数相乘,或者两个多项式相乘,算得手都酸了,还容易出错?比如说,你想计算
(3x² + 2x + 1) × (4x + 5)
逐项乘开要算 3×2×2 = 12 次运算,如果次数更高,比如几百次方,手工根本算不过来。

之前我们学过 FFT(快速傅里叶变换),它能把多项式乘法加速到 O(n log n)。但 FFT 用的是浮点数复数,计算时会产生舍入误差,比如本来该是整数 42,结果算出 41.99999998…… 四舍五入后虽能恢复,但在要求精确整数结果的场景(比如大数乘法、模素数下的乘法)中,我们更希望所有运算都在整数范围内完成,没有任何误差。

NTT(数论变换,Number Theoretic Transform) 就是 FFT 在整数模世界里的“孪生兄弟”。它把复数单位根换成了模素数 p 下的原根,于是所有运算都变成模 p 的整数运算,结果百分百精确,速度也同样是 O(n log n)。

下面我们就来认识这位“整数超人”吧!


1. 为什么要用 NTT?—— 精确、快速、无误差

  • FFT 的“小烦恼”:用复数计算时,每个乘法都会有小误差,多次累加后结果可能差一点点。虽然可以四舍五入到最近整数,但如果结果很大(比如几百位的数字),误差累积会变得不可控制。
  • 小例子:计算 (x² + 1) × (x² + 1) 的精确结果应该是 x⁴ + 2x² + 1。如果用 FFT 算,可能得到 x⁴ + 2.0000001x² + 0.9999998,虽然最后能圆回来,但心里总觉得不踏实。
  • NTT 的“大优势”:所有加法和乘法都在整数模 p 下进行,没有浮点数,没有误差。尤其适合密码学(比如 RSA 加密里的大数乘法)、大整数乘法(比如算 123456789101112… × 111213141516…),以及竞赛编程中的多项式运算。

2. 数学原理 —— 原根是如何模仿单位根的?

回忆 FFT 用到的复数单位根 ω,它有以下神奇性质:

  • ωⁿ = 1,且对于 k = 1,2,…,n-1,ωᵏ ≠ 1(即 ω 的阶是 n)
  • ω^(n/2) = -1(在复数里 -1 是实数)
  • 折半定理:ω^(2k) 是 n/2 次单位根

在模素数 p 的世界里,我们需要找一个“整数版的 ω”,让它也满足这些性质。

2.1 原根是什么?

  • 定义:对于素数 p,如果存在一个整数 g,使得 g¹, g², …, g^(p-1) 在模 p 下都不相等(也就是 g 的阶是 p-1),那么 g 就是 p 的一个原根
  • 小例子:p = 7,我们试试 g = 3:
    3¹ ≡ 3, 3² ≡ 2, 3³ ≡ 6, 3⁴ ≡ 4, 3⁵ ≡ 5, 3⁶ ≡ 1(模 7)
    确实 3 的阶是 6 = p-1,所以 3 是 7 的一个原根。

2.2 从原根到单位根

如果多项式长度 n 能整除 p-1(即 n | (p-1)),那么令
ω = g^((p-1)/n)
则 ω 的阶正好是 n。而且,根据费马小定理,g^(p-1) ≡ 1,所以
ω^(n/2) = g^((p-1)/2) ≡ -1 (mod p)
这里的 -1 就是 p-1 啦!比如 p=7,n=2,取 g=3,则 ω = 3^((7-1)/2) = 3³ = 27 ≡ 6 ≡ -1 (mod 7),完美。

这样,复数单位根的所有性质——阶为 n、折半、旋转——都在模 p 下实现了。

2.3 常用的好模数

不是随便找个素数就能用的,我们需要 n 是 2 的幂(比如 1024、2048),而且 n 能整除 p-1。所以最常用的模数是:

  • p = 998244353 = 119 × 2²³ + 1,它的原根是 3,支持 n 最大到 2²³(约 800 万),足够大多数应用。
  • p = 1004535809,原根也是 3,支持 n 最大到 2²¹。

你只要记住这两兄弟,大部分 NTT 题目都能搞定。


3. NTT 算法的步骤 —— 和 FFT 几乎一样,只是把复数换成了模运算

NTT 的迭代实现和 FFT 长得一模一样,只不过每次加减乘都要 取模。下面我们一步步看。

3.1 位反转排序(蝴蝶飞舞前的准备)

FFT 需要先把系数排成“位反转”顺序(也就是二进制数的低位到高位反过来)。比如长度为 8 的多项式,原始顺序是 0,1,2,3,4,5,6,7,位反转后变成 0,4,2,6,1,5,3,7。这步 NTT 也完全一样,因为这只和数组索引有关,不涉及数学运算。

3.2 蝶形操作(核心计算)

假设当前层长度为 len,我们要把两个长度为 len/2 的小多项式合并成一个长度为 len 的。蝶形公式是:

a[i]           = (u + w * v)            % p
a[i + len/2]   = (u - w * v + p)        % p

其中:

  • u = a[i](偶数索引部分)
  • v = a[i + len/2](奇数索引部分)
  • w 是当前旋转因子,从 1 开始,每次乘以 wlen(wlen 是这一层的单位根)

注意:u 和 v 都是模 p 下的数,所以加完之后要 % p,减法要加上 p 再 % p(防止负数)。

3.3 逆变换(INTT)

逆变换时,把 wlen 换成它的逆元(即 wlen^(p-2) 模 p),最后所有结果再乘以 n 的逆元 n^(p-2) 模 p。这样就能从点值表示变回系数表示。


4. 编程实现 —— 手把手写一个“整数快速乘法”

下面给出 C++ 和 Python 的完整实现。代码中所有变量都用简短英文单词,每行变量定义后面有中文注释,方便你理解。

4.1 C++ 实现(迭代版本)

#include <iostream>
#include <vector>
#include <algorithm>
using namespace std;

typedef long long ll;               // 用 long long 防止乘法溢出
const int MOD = 998244353;          // 常用模数
const int ROOT = 3;                 // 模数 MOD 的原根

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

// NTT 迭代实现
// a: 输入多项式系数(长度为 n,n 必须是 2 的幂)
// invert: false 表示 NTT(正变换),true 表示 INTT(逆变换)
void ntt(vector<ll>& a, bool invert) {
    int n = a.size();               // 多项式的长度

    // ---------- 位反转排序 ----------
    for (int i = 1, j = 0; i < n; ++i) {
        int bit = n >> 1;           // 最高位的 bit 值
        for (; j & bit; bit >>= 1)  // 从高位往低位找 1 的位置
            j ^= bit;               // 把那个 1 变成 0
        j ^= bit;                   // 最后一个 bit 取反(添加 1)
        if (i < j) swap(a[i], a[j]); // 交换两个位置
    }

    // ---------- 蝶形操作(自底向上迭代) ----------
    for (int len = 2; len <= n; len <<= 1) {   // len 从 2 开始,每次翻倍
        // 计算当前层的单位根 wlen = omega_len
        ll wlen = qpow(ROOT, (MOD - 1) / len);
        if (invert) wlen = qpow(wlen, MOD - 2); // 逆变换时用逆元

        for (int i = 0; i < n; i += len) {     // 每个长度为 len 的段
            ll w = 1;                          // 旋转因子初始为 1
            for (int j = 0; j < len / 2; ++j) {
                ll u = a[i + j];               // 偶数部分
                ll v = a[i + j + len/2] * w % MOD; // 奇数部分乘上 w
                a[i + j] = (u + v) % MOD;            // 新左半
                a[i + j + len/2] = (u - v + MOD) % MOD; // 新右半(防止负数)
                w = w * wlen % MOD;            // 更新下一个 w
            }
        }
    }

    // ---------- 逆变换时乘 n 的逆元 ----------
    if (invert) {
        ll n_inv = qpow(n, MOD - 2); // n 的逆元(费马小定理)
        for (int i = 0; i < n; ++i)
            a[i] = a[i] * n_inv % MOD;
    }
}

// 多项式乘法:使用 NTT 计算 A * B(系数模 MOD)
vector<ll> multiplyNTT(const vector<ll>& a, const vector<ll>& b) {
    // 计算需要的长度 n(不小于 a.size() + b.size() - 1 的 2 的幂)
    int n = 1;
    while (n < a.size() + b.size() - 1) n <<= 1;

    // 把系数复制到长度为 n 的数组中,后面补 0
    vector<ll> fa(n, 0), fb(n, 0);
    copy(a.begin(), a.end(), fa.begin());
    copy(b.begin(), b.end(), fb.begin());

    // 三次 NTT:正变换 -> 点乘 -> 逆变换
    ntt(fa, false);          // 正变换 A
    ntt(fb, false);          // 正变换 B
    for (int i = 0; i < n; ++i)
        fa[i] = fa[i] * fb[i] % MOD;   // 点乘
    ntt(fa, true);           // 逆变换回到系数

    // 去掉末尾多余的 0(如果有)
    while (fa.size() > 1 && fa.back() == 0)
        fa.pop_back();
    return fa;
}

int main() {
    // 多项式 A: 1 + 2x + 3x^2  相当于系数 {1,2,3}
    vector<ll> A = {1, 2, 3};
    // 多项式 B: 5 + 4x         相当于系数 {5,4}
    vector<ll> B = {5, 4};

    vector<ll> C = multiplyNTT(A, B);

    // 输出结果(从高次到低次)
    cout << "A * B (mod " << MOD << ") = ";
    for (int i = C.size() - 1; i >= 0; --i) {
        cout << C[i] << (i > 0 ? " " : "");
    }
    cout << endl;

    // 期望输出:15 22 23 4  即 15x^3 + 22x^2 + 23x + 4
    // 验证: (3x^2+2x+1)*(4x+5) = 12x^3+15x^2+8x^2+10x+4x+5 = 12x^3+23x^2+14x+5???
    // 哦,这里我算错了,应该是 (3x^2+2x+1)*(4x+5) = 12x^3 + 15x^2 + 8x^2 + 10x + 4x + 5 = 12x^3 + 23x^2 + 14x + 5
    // 但是系数是 {5,14,23,12},从高到低是 12 23 14 5。我们的代码输出是?
    // 注意系数顺序:vector 里下标 0 对应常数项,下标 1 对应 x 项……所以输出从后往前时会逆转。
    // 我们来手动计算一下:
    // A = [1,2,3] 表示 3x^2 + 2x + 1
    // B = [5,4]   表示 4x + 5
    // 乘积: (3x^2+2x+1)*(4x+5) = 12x^3 + 15x^2 + 8x^2 + 10x + 4x + 5 = 12x^3 + 23x^2 + 14x + 5
    // 系数 [5,14,23,12]  -> 从高到低打印:12 23 14 5
    return 0;
}

4.2 Python 实现(迭代版本)

MOD = 998244353    # 常用模数
ROOT = 3           # 模数 MOD 的原根

def qpow(x, e, mod=MOD):
    """快速幂:计算 x^e mod MOD"""
    res = 1
    while e:
        if e & 1:
            res = res * x % mod
        x = x * x % mod
        e >>= 1
    return res

def ntt(a, invert):
    """NTT 迭代实现
       a: 列表,长度 n(2 的幂)
       invert: False 为正变换,True 为逆变换
    """
    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):            # 每个长度为 length 的段
            w = 1
            half = length // 2
            for j in range(i, i + half):
                u = a[j]                         # 偶数部分
                v = a[j + half] * w % MOD         # 奇数部分乘上 w
                a[j] = (u + v) % MOD              # 新左半
                a[j + half] = (u - v + MOD) % MOD # 新右半
                w = w * wlen % MOD                # 更新 w

        length <<= 1

    # ---------- 逆变换时乘 n 的逆元 ----------
    if invert:
        n_inv = qpow(n, MOD - 2)
        for i in range(n):
            a[i] = a[i] * n_inv % MOD

def multiply_ntt(a, b):
    """使用 NTT 计算两个多项式乘法(系数列表 a 和 b)"""
    # 计算需要的长度 n(不小于 len(a)+len(b)-1 的 2 的幂)
    n = 1
    while n < len(a) + len(b) - 1:
        n <<= 1

    # 补成长度为 n 的列表,多余位置填 0
    fa = a + [0] * (n - len(a))
    fb = b + [0] * (n - len(b))

    # 三次 NTT
    ntt(fa, False)          # 正变换
    ntt(fb, False)
    for i in range(n):
        fa[i] = fa[i] * fb[i] % MOD   # 点乘
    ntt(fa, True)           # 逆变换

    # 去掉末尾多余的 0
    while len(fa) > 1 and fa[-1] == 0:
        fa.pop()
    return fa

# 测试
A = [1, 2, 3]      # 3x^2 + 2x + 1
B = [5, 4]         # 4x + 5
C = multiply_ntt(A, B)
print("A * B (mod {}):".format(MOD), C[::-1])  # 从高次到低次打印

5. 常见错误与调试小技巧

5.1 模数选错了

  • 错误:用了素数 p,但 n 不能整除 p-1。比如 p=7,n=4,4 不能整除 6,那么原根构造不出阶为 4 的元素,算法会出问题。
  • 解决:用 998244353 或 1004535809 这种“好模数”,它们都能支持 2 的幂到很大。

5.2 多项式长度不是 2 的幂

  • 错误:直接对长度为 3 的数组做 NTT,结果随机的。
  • 解决:一定要把长度补成大于等于 (lenA+lenB-1) 的最小 2 的幂,并且后面填 0。

5.3 忘记取模导致溢出

  • 错误:在蝶形操作中,u + v 可能大于 MODu - v 可能为负数,不取模会得到错误结果(甚至负数)。
  • 解决:每一步加减乘后都 % MOD,减法要 (u - v + MOD) % MOD

5.4 逆元计算错误

  • 错误:乘了 n 的逆元,但用的是 1/n 或者 n 直接除,模运算下除法要转换成乘逆元。
  • 解决:用快速幂计算 n^(MOD-2) % MOD(费马小定理,要求 MOD 是素数)。

5.5 输入系数太大

  • 错误:系数大于等于 MOD,比如 MOD=998244353,但系数是 1000000000,直接运算会错误。
  • 解决:确保所有系数都在 [0, MOD-1] 范围内,可以预先取模。

6. 完整示例:计算大数乘法

还记得文章开头说的大数乘法吗?比如计算 123456789 × 987654321,我们可以把数字看成多项式:
123456789 = 1×10⁸ + 2×10⁷ + 3×10⁶ + 4×10⁵ + 5×10⁴ + 6×10³ + 7×10² + 8×10¹ + 9
然后利用 NTT 在模一个大素数下(比如 998244353)运算,最后处理进位。因为 NTT 的结果是精确的整数,所以进位后就能得到准确的大数乘积。

以下代码展示如何用 NTT 计算两个大整数的乘积(不考虑进位,只演示多项式乘法):

#include <iostream>
#include <vector>
#include <string>
using namespace std;

// 把数字字符串转换成系数列表(低位在前,高位在后)
vector<ll> str_to_vec(const string& s) {
    vector<ll> res(s.size());
    for (int i = 0; i < s.size(); ++i)
        res[s.size()-1-i] = s[i] - '0';   // 低位在前
    return res;
}

int main() {
    string sa = "123456789";
    string sb = "987654321";
    vector<ll> A = str_to_vec(sa);
    vector<ll> B = str_to_vec(sb);

    vector<ll> C = multiplyNTT(A, B);  // 多项式乘法,结果在模 MOD 下

    // 注意:这里的结果是模 MOD 的,要得到真正的大数还需要处理进位,
    // 但作为演示,我们只输出多项式系数(模 MOD 后)
    cout << "系数(模 MOD): ";
    for (int i = C.size()-1; i >= 0; --i)
        cout << C[i] << " ";
    cout << endl;
    return 0;
}

7. 总结与相关指引

项目说明
NTT 是什么在整数模 p 下实现的多项式乘法,无浮点误差,速度 O(n log n)
核心公式用原根 g 构造 ω = g^((p-1)/n),代替复数单位根
常用模数998244353(原根 3)、1004535809(原根 3)
适用场景大数乘法、密码学(RSA、ECC)、竞赛编程、精确多项式运算
相关知识点快速幂、逆元、原根、FFT、多项式求逆、多项式除法、卷积

如果你想继续深入,下一步可以学习:

  • 多项式求逆:利用 NTT 计算一个多项式的逆元(用于解方程)。
  • 快速数论变换的变种:如任意模数 NTT(用三个模数后用中国剩余定理还原)。
  • NTT 在密码学中的应用:比如 RSA 中快速计算大数模乘。

NTT 就像一个整数世界的“魔法加速器”,掌握了它,你就能轻松应对各种精确、高速的多项式计算啦!

例题精讲

1单选题

在NTT中,以下哪个模数最适合用于长度不超过2^23的多项式乘法?

A1000000007
B998244353
C1000000009
D1000003
2判断题

NTT中,对于模数p,如果存在原根g,则g^(p-1) ≡ 1 mod p,且对于p-1的任意真因子d,g^d ≠ 1 mod p。

3填空题
在NTT逆变换中,通常最后一步需要对每个系数乘以n的模逆。请填补以下代码中的空白:
for (int i = 0; i < n; i++) a[i] = (long long) a[i] * ___ % mod;
4单选题

关于NTT与FFT,下列哪项说法正确?

ANTT只能处理模数为2的幂
BNTT是FFT在有限域上的模拟,使用原根代替单位根
CNTT不需要位逆序置换
DNTT的旋转因子是复数
5判断题

在NTT中,如果模数p不是素数,则无法进行NTT变换。