多项式乘法逆与NTT的综合应用
极难6多项式乘法逆与NTT:从零搭建高效计算积木
你是否想过,如果有一个“万能工具”,能让你快速算出多项式的逆、除法、开方等运算,就像用魔方公式还原魔方一样?这个工具就是 NTT(数论变换),而今天我们要用它来实现一个超级实用的功能——多项式求逆。学会了它,你就能像搭积木一样,轻松搞定更复杂的数学计算。
什么是多项式乘法逆?
想象一下,你在做数学题时,老师要求你找到一个数 ,使得 ,这个 就是 的倒数(乘法逆元)。对于多项式来说,道理类似:
对于多项式 ,如果存在另一个多项式 ,使得它们的乘积在 模 意义下等于常数1,即:
那么 就是 在模 下的 乘法逆元。这里的“模 ”是什么意思呢?简单说,就是只保留乘积中次数小于 的项,更高次的项全部丢掉。就像你只关心前几个数字,后面的忽略不计。
举个生活中的例子
假设你每个月的零花钱是 ,其中 代表月份(比如 表示第一个月、第二个月...)。你想找一个“反向公式” ,使得 的结果恰好等于1(常数),而且只计算前两个月的效果(因为模 )。经过尝试,你发现 就满足:
因为 项已经超出“模 ”的范围(我们只关心 和 项),所以结果就是1,完美!所以 就是 的逆元。
注意:求逆的前提是常数项不能为0(否则没有逆元,就像0没有倒数一样)。
核心原理:牛顿迭代法——从粗糙到精确
你可能觉得,直接算逆元很麻烦,尤其是多项式很长的时候。但别怕,我们可以用 牛顿迭代法,它是一种逐步逼近的方法,就像你玩射击游戏时,先瞄准一个大概的方向,然后不断修正,最终精准命中。
数学公式(不怕,我们用人话解释)
假设我们已经知道了 在模 下的逆元 (精度较低),想要求模 下的更精确逆元 。牛顿迭代给出了一个公式:
这个公式可以这样理解:
- 先计算 ,看它离1有多远(如果完美,结果应该等于1)。
- 然后用 减去这个结果(因为 ,如果结果偏离,这个差会修正回去)。
- 最后再乘上 ,就得到了更精确的逆元。
迭代过程:我们从常数项开始(模 ),只要常数项非零,它的逆元就是它的倒数(比如 的逆元是 )。然后每次将精度加倍(从 到 ,再到 ,直到 ),每次迭代用一次牛顿公式。每次迭代需要做两次多项式乘法,而乘法可以用NTT加速,所以总的时间复杂度是 ,相当快!
动手编程:用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))
新手常见错误与避坑指南
-
常数项为0的情况
如果多项式 的常数项 ,那么它没有逆元(因为任何多项式乘以它后,常数项都是0,不可能得到1)。解决方法是先判断常数项是否非零,如果为零,需要先提取 因子(例如 ),再对 求逆。 -
NTT长度选择不对
牛顿迭代中,每次卷积需要长度为 ,我们向上取到2的幂。如果长度只取 ,会导致结果被截断,无法正确迭代。一定要确保卷积长度足够大。 -
忘记模运算
所有中间结果都要取模,尤其是减法时要注意先加MOD再取模,防止负数出现。 -
逆变换后忘记除以n
NTT逆变换后,每个元素需要乘以 的逆元,否则结果会多一个因子 。检查代码中是否正确处理了if invert部分。
应用:多项式求逆能做什么?
学会了多项式求逆,就像拥有了一把万能钥匙,可以打开很多高级运算的大门:
-
多项式除法:已知多项式 和 ,求商 和余数 ,使得 。方法:先计算 的逆元,然后 。这就像做整数除法时,用除数的倒数乘以被除数。
-
多项式开方:求 使得 ,也需要用到求逆(牛顿迭代法类似)。
-
多项式对数、指数:在生成函数、组合数学中,经常需要计算 或 ,这些公式里都包含多项式求逆或乘法。
你可以把多项式求逆想象成乐高积木中的基础砖块,有了它,就能搭出更复杂的结构,比如在算法竞赛中解决快速计算、图像处理、信号处理等领域的问题。
总结
从多项式的基本概念出发,我们学习了用NTT加速多项式乘法,然后利用牛顿迭代法实现了多项式求逆。整个过程就像一步步升级装备:先学会两位数乘法,再学速算技巧,最后能轻松完成百位数乘法。现在你已经掌握了“多项式求逆”这个高级技能,接下来可以尝试挑战多项式除法、开方、指数等更酷的运算!
相关指引:如果你对FFT、NTT还不太熟悉,建议先复习一下“快速傅里叶变换”和“数论变换”的原理。之后可以学习“多项式乘法”、“多项式除法”等专题。加油,下一个编程数学高手就是你!
例题精讲
使用NTT实现多项式乘法逆,若多项式次数为n,则求逆的渐近时间复杂度为?
在模数为998244353时,可以使用NTT进行多项式乘法,因为该模数是质数且存在原根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
}模数998244353下,NTT支持的最大变换长度(2的幂次)为多少?
利用牛顿迭代法求多项式乘法逆时,需要保证常数项a[0] ≠ 0。