多项式入门
数学
多项式乘法
单位根
我们称方程 \(x^n = 1\) 的 \(n\) 个在复数域上的解为单位根,记作 \(\omega_n^1, \omega_n^2, \omega_n^3, ..., \omega_n^n\),我们发现这些复数都满足模长为 \(1\),辐角为 \(\frac{2\pi}{n}\) 的倍数,我们认为 \(\omega_n^k\) 的辐角为 \(\frac{2k\pi}{n}\),推广一下我们就有 \(\omega_n^i = \omega_n^{i + nk}\)。
这个有一些优美的性质:
第一个第二个很显然啊,第三个就是转 \(180\) 度复数变成原来的相反数,第四个大概就是每个单位根的和为 \(0\)。
DFT 和 IDFT
DFT 又称离散傅里叶变换, IDFT 就是他的逆运算。
DFT 是指做这样一个事情,其中 \(c_i\) 是原多项式的对应项的系数:
可以理解为是把对应的单位根带进去的结果。
IDFT 就是这个式子:
带进去就可以证明互为逆运算了。
为什么这个对我们计算多项式乘法有用呢?
因为两个多项式 \(A, B\),他们的乘法就为
至于为什么,我也不会,感觉跟 FWT 很类似。
FFT
FFT 就是就是一种快速计算 DFT 和 IDFT 的方法。
我们钦定多项式次数为二的整次幂:
然后我们设
就有:
我们把 \(\omega_n\) 带进去
对于 \(0 \le 2k < n\) 的数
然后对于其他的:
我们发现后面的是递归子问题。
但是这样直接做常数直接起飞
于是我们列一下这个数组的变化:
我们发现最后一行的数组就是下标 \(i\) 的颠倒过来,于是就可以从下往上做,同时维护目前数组下标就可以了。
//奇跡を信じて,願いを叶えて
#include<bits/stdc++.h>
using namespace std;
#define int long long
#define double long double
#define uint unsigned long long
#define Air
namespace io{
inline int read(){
int f = 1, t = 0; char ch = getchar();
while(ch < '0' || ch > '9'){if(ch == '-') f = -f; ch = getchar();}
while(ch >= '0' && ch <= '9'){t = t * 10 + ch - '0'; ch = getchar();}
return t * f;
}
inline void write(int x){
if(x < 0){putchar('-'); x = -x;}
if(x >= 10){write(x / 10);}
putchar(x % 10 + '0');
}
}
using namespace io;
int n, m;
const int N = 2.2e6 + 10;
const double Pi = acos(-1);
struct Complex{
double x, y;
Complex (double xx = 0, double yy = 0){
x = xx;
y = yy;
return ;
}
friend Complex operator + (Complex x, Complex y){
return {x.x + y.x, x.y + y.y};
}
friend Complex operator - (Complex x, Complex y){
return {x.x - y.x, x.y - y.y};
}
friend Complex operator * (Complex x, Complex y){
return {x.x * y.x - x.y * y.y, x.y * y.x + x.x * y.y};
}
};
#define vec vector<Complex>
int rev[N];
vec a, b;
void fft(vec &x, int lim, bool op){
//op = 1 为正操作,op = 0 为逆操作
for(int i = 0; i < lim; i++){
if(i < rev[i]) swap(x[i], x[rev[i]]);
}
for(int len = 1; len < lim; len <<= 1){
Complex ompr = {cos(Pi / len), (2 * op - 1) * sin(Pi / len)};
//求出每次增多的单位根
for(int i = 0; i < lim; i += 2 * len){
Complex ome = {1, 0};
//最初的单位根
for(int j = i; j < i + len; j++){
//枚举对应的合并位置
Complex l = x[j], r = x[j + len];
x[j] = l + ome * r;
x[j + len] = l - ome * r;
//按照公式合并
ome = ome * ompr;
}
}
}
// return ;
if(!op){
for(int i = 0; i < lim; i++){
x[i].x /= lim;
x[i].y /= lim;
}
}
}
vec mul(vec a, vec b){
int x = a.size() + b.size() - 1;
int lim = 1;
while(lim < x) lim <<= 1;
for(int i = 0; i < lim; i++){
rev[i] = rev[i / 2] / 2 + (i & 1) * lim / 2;
}
while(a.size() < lim) a.push_back(0);
while(b.size() < lim) b.push_back(0);
// return a;
fft(a, lim, 1);
// cerr << "!!!\n";
// return a;
fft(b, lim, 1);
// return a;
for(int i = 0; i < lim; i++) a[i] = a[i] * b[i];
fft(a, lim, 0);
while(a.size() > x) a.pop_back();
return a;
}
signed main() {
#ifndef Air
freopen(".in","r",stdin);
freopen(".out","w",stdout);
#endif
ios::sync_with_stdio(false);cin.tie(0);cout.tie(0);
n = read();
m = read();
for(int i = 0; i <= n; i++){
a.push_back(read());
}
for(int i = 0; i <= m; i++){
b.push_back(read());
}
vec c = mul(a, b);
for(int i = 0; i <= n + m; i++){
cout << (int)(c[i].x + 0.5) << ' ';
}
return 0;
}
NTT
但是 FFT 有几个问题,一是精度会有误差,二是无法处理有模数的情况。
当要对某个东西取模的时候,就需要使用 NTT。
NTT 的过程就是把单位根换成了模数的原根。
具体来说,假如模数的原根是 \(G\),那么把 \(\omega_n^k\) 替换成 \(G^{\frac{(p - 1)k}{n}}\) 就是对的,我不会证,大概就是原根和单位根性质类似。
//奇跡を信じて,願いを叶えて
#include<bits/stdc++.h>
using namespace std;
#define int long long
#define double long double
#define uint unsigned long long
#define Air
namespace io{
inline int read(){
int f = 1, t = 0; char ch = getchar();
while(ch < '0' || ch > '9'){if(ch == '-') f = -f; ch = getchar();}
while(ch >= '0' && ch <= '9'){t = t * 10 + ch - '0'; ch = getchar();}
return t * f;
}
inline void write(int x){
if(x < 0){putchar('-'); x = -x;}
if(x >= 10){write(x / 10);}
putchar(x % 10 + '0');
}
}
using namespace io;
int n;
const int N = 5e5 + 10, MOD = 167772161, G = 3, IG = 55924054;
int quick_pow(int a, int b){
if(!b) return 1;
if(b & 1){
return a * quick_pow(a, b - 1) % MOD;
}
else{
int tmp = quick_pow(a, b / 2);
return tmp * tmp % MOD;
}
}
#define vec vector<int>
int rev[N];
void ntt(vec &x, int lim, bool op){
for(int i = 0; i < lim; i++){
if(i < rev[i]) swap(x[i], x[rev[i]]);
}
for(int len = 1; len < lim; len <<= 1){
int omepr = quick_pow(op ? G : IG, (MOD - 1) / (len << 1));
for(int i = 0; i < lim; i += len * 2){
int ome = 1;
for(int j = i; j < i + len; j++){
int l = x[j], r = x[j + len];
x[j] = l + r * ome % MOD; x[j] %= MOD;
x[j + len] = l - r * ome % MOD + MOD; x[j + len] %= MOD;
ome *= omepr; ome %= MOD;
}
}
}
if(!op){
for(int i = 0; i < lim; i++){
// cerr << x[i] << ' ';
x[i] = x[i] * quick_pow(lim, MOD - 2) % MOD;
}
// cerr << '\n';
}
}
vec mul(vec a, vec b){
int x = a.size() + b.size() - 1;
int lim = 1;
while(lim < x) lim <<= 1;
for(int i = 0; i < lim; i++){
rev[i] = rev[i / 2] / 2 + (i & 1) * lim / 2;
}
while(a.size() < lim) a.push_back(0);
while(b.size() < lim) b.push_back(0);
ntt(a, lim, 1); ntt(b, lim, 1);
for(int i = 0; i < lim; i++){
a[i] = a[i] * b[i] % MOD;
}
ntt(a, lim, 0);
while(a.size() > x) a.pop_back();
return a;
}
int fac[N];
signed main() {
#ifdef Air
freopen(".in","r",stdin);
freopen(".out","w",stdout);
#endif
ios::sync_with_stdio(false);cin.tie(0);cout.tie(0);
n = read();
fac[0] = 1;
for(int i = 1; i < N; i++){
fac[i] = fac[i - 1] * i % MOD;
}
vec a, b; a.clear(); b.clear();
for(int i = 0; i <= n; i++){
int p1 = quick_pow(i, n), p2 = ((i & 1) ? (MOD - 1) : 1), inv = quick_pow(fac[i], MOD - 2);
a.push_back(p1 * inv % MOD);
b.push_back(p2 * inv % MOD);
}
a = mul(a, b);
for(int i = 0; i <= n; i++){
cout << a[i] << ' ';
}
return 0;
}
多项式求逆
主要就是一个式子,我们假如知道前 \(2^i\) 项的逆元 \(B'\) 那么如何推出前 \(2^{i + 1}\) 项的逆元 \(B\) 呢?
我们有,方便起见我们称 \(z = 2^i\) :
就有
就能得到:
直接递推 NTT 就做完了,复杂度好像是 \(n \log n\) 的。
//奇跡を信じて,願いを叶えて
#include<bits/stdc++.h>
using namespace std;
#define int long long
#define double long double
#define uint unsigned long long
#define Air
namespace io{
inline int read(){
int f = 1, t = 0; char ch = getchar();
while(ch < '0' || ch > '9'){if(ch == '-') f = -f; ch = getchar();}
while(ch >= '0' && ch <= '9'){t = t * 10 + ch - '0'; ch = getchar();}
return t * f;
}
inline void write(int x){
if(x < 0){putchar('-'); x = -x;}
if(x >= 10){write(x / 10);}
putchar(x % 10 + '0');
}
}
using namespace io;
int n;
const int N = 4e5 + 10, MOD = 998244353, G = 3, IG = 332748118;
#define vec vector<int>
int rev[N];
int quick_pow(int a, int b){
int res = 1;
while(b){
if(b & 1) res = res * a % MOD;
a = a * a % MOD;
b >>= 1;
}
return res;
}
void ntt(vec &x, int lim, bool op){
for(int i = 0; i < lim; i++){
if(i < rev[i]) swap(x[i], x[rev[i]]);
}
for(int len = 1; len < lim; len <<= 1){
int omepr = quick_pow(op ? G : IG, (MOD - 1) / (len << 1));
for(int i = 0; i < lim; i += len * 2){
int ome = 1;
for(int j = i; j < i + len; j++){
int l = x[j], r = x[j + len];
x[j] = l + ome * r % MOD; x[j] -= (x[j] >= MOD ? MOD : 0);
x[j + len] = l - ome * r % MOD + MOD; x[j + len] -= (x[j + len] >= MOD ? MOD : 0);
ome = ome * omepr % MOD;
}
}
}
if(!op){
int tmp = quick_pow(lim, MOD - 2);
for(int i = 0; i < lim; i++){
x[i] = x[i] * tmp % MOD;
}
}
}
vec mul(vec a, vec b){
int x = a.size() + b.size() - 1;
int lim = 1;
while(lim < x) lim <<= 1;
for(int i = 0; i < lim; i++){
rev[i] = (rev[i >> 1] >> 1) + (i & 1) * (lim >> 1);
}
while(a.size() < lim) a.push_back(0);
while(b.size() < lim) b.push_back(0);
ntt(a, lim, 1); ntt(b, lim, 1);
for(int i = 0; i < lim; i++){
a[i] = a[i] * b[i] % MOD;
}
ntt(a, lim, 0);
while(a.size() > x) a.pop_back();
return a;
}
// int a[N];
vector<int> a;
vec b[30];
signed main() {
#ifndef Air
freopen(".in","r",stdin);
freopen(".out","w",stdout);
#endif
ios::sync_with_stdio(false);cin.tie(0);cout.tie(0);
n = read();
for(int i = 1; i <= n; i++){
a.push_back(read());
}
b[0].push_back(quick_pow(a[0], MOD - 2));
int id = 0;
for(int i = 0; (1ll << i) <= n; i++){
vec tmp = b[i];
vec tt; tt.clear();
id ++;
for(int j = 0; j < tmp.size(); j++){
tmp[j] *= 2; tmp[j] %= MOD;
}
for(int j = 0; j < (1ll << i); j++){
tmp.push_back(0);
tt.push_back(a[j]);
}
for(int j = (1ll << i); j < (1ll << (i + 1)); j++){
tt.push_back(j < n ? a[j] : 0);
}
b[i] = mul(b[i], mul(b[i], tt));
for(int j = 0; j < (1ll << (i + 1)); j++){
tmp[j] = tmp[j] - b[i][j] + MOD;
tmp[j] -= (tmp[j] >= MOD ? MOD : 0);
}
b[i + 1] = tmp;
}
for(int i = 0; i < n; i++){
cout << b[id][i] << ' ';
}
return 0;
}

浙公网安备 33010602011771号