NTT 数论变换
极难3NTT 数论变换——整数世界里的“超快速”多项式乘法
你是不是遇到过这样的问题:两个很大的整数相乘,或者两个多项式相乘,算得手都酸了,还容易出错?比如说,你想计算
(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可能大于MOD,u - 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 就像一个整数世界的“魔法加速器”,掌握了它,你就能轻松应对各种精确、高速的多项式计算啦!
例题精讲
在NTT中,以下哪个模数最适合用于长度不超过2^23的多项式乘法?
NTT中,对于模数p,如果存在原根g,则g^(p-1) ≡ 1 mod p,且对于p-1的任意真因子d,g^d ≠ 1 mod p。
在NTT逆变换中,通常最后一步需要对每个系数乘以n的模逆。请填补以下代码中的空白:
for (int i = 0; i < n; i++) a[i] = (long long) a[i] * ___ % mod;关于NTT与FFT,下列哪项说法正确?
在NTT中,如果模数p不是素数,则无法进行NTT变换。