给我一份全一点的多项式模板
下面是一份以 998244353 为模数的 C++17 多项式模板,覆盖 ICPC 中最常用的内容:
- NTT / 卷积
- 多项式加减乘
- 导数、积分
- 多项式求逆
- ln
- exp
- 幂函数
- 平方根
- 多项式除法与取模
- 单点求值
下面按模板中的函数分类说明。默认多项式系数下标从 0 开始,n 表示保留的系数个数。
Mint 相关
函数 参数说明
Mint(x) 用整数 x 初始化模数类
pow(n) 当前数的 n 次幂
inv() 当前数的乘法逆元
operator + - * / 模意义下的四则运算
基础操作
函数 参数说明
ntt(a, invert) 对多项式 a 做 NTT;invert=false 为正变换,true 为逆变换
trim(a) 删除末尾的零系数
cut(a, n) 截取或补零,使多项式长度变为 n
convolution(a, b) 计算多项式乘法,不截断结果
add(a, b) 计算 a+b
sub(a, b) 计算 a-b
mul(a, b) 计算 a*b,等价于 convolution(a,b)
微积分操作
函数 参数说明
derivative(a) 求多项式 a 的导数
integral(a) 求多项式 a 的积分,常数项设为 0
例如:
a = {a0, a1, a2};
derivative(a) = {a1, 2*a2};
integral(a) = {0, a0, a1/2, ...};
形式幂级数操作
函数 参数说明
inverse(a, n) 求 a 的逆,满足 a*b ≡ 1 (mod x^n);要求 a[0] != 0
logarithm(a, n) 求 ln(a) mod x^n;要求 a[0] != 0
exponential(a, n) 求 exp(a) mod x^n;要求 a[0] == 0
square_root(a, n) 求 sqrt(a) mod x^n;无平方根时返回空多项式
power(a, k, n) 求 a^k mod x^n;k 为非负整数
其中:
inverse(a, n)
返回前 n 项系数。
power(a, k, n)
中的 n 不是最高次数,而是保留系数个数。
平方根相关
函数 参数说明
sqrt_mod(a) 求单个模数意义下的平方根;无解返回 -1
square_root(a, n) 求多项式平方根的前 n 项
多项式除法
函数 参数说明
divmod(a, b) 返回 {商, 余数},满足 a = b * q + r
modulo(a, b) 只返回 a 除以 b 的余数
例如:
auto [q, r] = divmod(a, b);
单点求值
函数 参数说明
evaluate(a, x) 计算多项式 a 在点 x 处的值,即 a(x)
例如:
Poly a = {1, 2, 3}; // 1 + 2x + 3x^2
Mint y = evaluate(a, 5); // 1 + 2*5 + 3*25
#include <bits/stdc++.h>
using namespace std;
namespace poly {
constexpr int mod = 998244353;
constexpr int G = 3;
struct Mint {
int v;
Mint(long long x = 0) {
x %= mod;
if (x < 0) x += mod;
v = int(x);
}
Mint& operator+=(const Mint& o) {
v += o.v;
if (v >= mod) v -= mod;
return *this;
}
Mint& operator-=(const Mint& o) {
v -= o.v;
if (v < 0) v += mod;
return *this;
}
Mint& operator*=(const Mint& o) {
v = int((long long)v * o.v % mod);
return *this;
}
Mint pow(long long n) const {
Mint a = *this, r = 1;
while (n) {
if (n & 1) r *= a;
a *= a;
n >>= 1;
}
return r;
}
Mint inv() const {
return pow(mod - 2);
}
Mint& operator/=(const Mint& o) {
return *this *= o.inv();
}
friend Mint operator+(Mint a, const Mint& b) {
return a += b;
}
friend Mint operator-(Mint a, const Mint& b) {
return a -= b;
}
friend Mint operator*(Mint a, const Mint& b) {
return a *= b;
}
friend Mint operator/(Mint a, const Mint& b) {
return a /= b;
}
friend Mint operator-(const Mint& a) {
return Mint(-a.v);
}
bool operator==(const Mint& o) const {
return v == o.v;
}
bool operator!=(const Mint& o) const {
return v != o.v;
}
};
using Poly = vector<Mint>;
void ntt(Poly& a, bool invert) {
int n = (int)a.size();
for (int i = 1, j = 0; i < n; i++) {
int bit = n >> 1;
while (j & bit) {
j ^= bit;
bit >>= 1;
}
j ^= bit;
if (i < j) swap(a[i], a[j]);
}
for (int len = 2; len <= n; len <<= 1) {
Mint wn = Mint(G).pow((mod - 1) / len);
if (invert) wn = wn.inv();
for (int i = 0; i < n; i += len) {
Mint w = 1;
for (int j = 0; j < len / 2; j++) {
Mint u = a[i + j];
Mint v = a[i + j + len / 2] * w;
a[i + j] = u + v;
a[i + j + len / 2] = u - v;
w *= wn;
}
}
}
if (invert) {
Mint inv_n = Mint(n).inv();
for (auto& x : a) x *= inv_n;
}
}
Poly trim(Poly a) {
while (!a.empty() && a.back() == Mint(0)) a.pop_back();
return a;
}
Poly cut(const Poly& a, int n) {
Poly b = a;
b.resize(n);
return b;
}
Poly convolution(const Poly& a, const Poly& b) {
if (a.empty() || b.empty()) return {};
if (min(a.size(), b.size()) <= 40) {
Poly c(a.size() + b.size() - 1);
for (int i = 0; i < (int)a.size(); i++) {
for (int j = 0; j < (int)b.size(); j++) {
c[i + j] += a[i] * b[j];
}
}
return c;
}
int n = 1;
while (n < (int)a.size() + (int)b.size() - 1) n <<= 1;
Poly x(a.begin(), a.end());
Poly y(b.begin(), b.end());
x.resize(n);
y.resize(n);
ntt(x, false);
ntt(y, false);
for (int i = 0; i < n; i++) x[i] *= y[i];
ntt(x, true);
x.resize(a.size() + b.size() - 1);
return x;
}
Poly add(Poly a, const Poly& b) {
if (a.size() < b.size()) a.resize(b.size());
for (int i = 0; i < (int)b.size(); i++) {
a[i] += b[i];
}
return a;
}
Poly sub(Poly a, const Poly& b) {
if (a.size() < b.size()) a.resize(b.size());
for (int i = 0; i < (int)b.size(); i++) {
a[i] -= b[i];
}
return a;
}
Poly mul(const Poly& a, const Poly& b) {
return convolution(a, b);
}
Poly derivative(const Poly& a) {
if (a.size() <= 1) return {};
Poly b(a.size() - 1);
for (int i = 1; i < (int)a.size(); i++) {
b[i - 1] = a[i] * Mint(i);
}
return b;
}
Poly integral(const Poly& a) {
static vector<Mint> invs{Mint(0), Mint(1)};
int n = (int)a.size();
while ((int)invs.size() <= n) {
int i = invs.size();
invs.push_back(-Mint(mod / i) * invs[mod % i]);
}
Poly b(n + 1);
for (int i = 0; i < n; i++) {
b[i + 1] = a[i] * invs[i + 1];
}
return b;
}
Poly inverse(const Poly& a, int n) {
assert(!a.empty() && a[0] != Mint(0));
if (n == 0) return {};
Poly r{a[0].inv()};
for (int len = 1; len < n; len <<= 1) {
int lim = len << 1;
Poly f = cut(a, lim);
Poly g = convolution(f, r);
g.resize(lim);
for (auto& x : g) x = -x;
g[0] += Mint(2);
r = convolution(r, g);
r.resize(lim);
}
r.resize(n);
return r;
}
Poly logarithm(const Poly& a, int n) {
assert(!a.empty() && a[0] != Mint(0));
if (n <= 1) return Poly(n, Mint(0));
Poly da = derivative(a);
Poly ia = inverse(a, n - 1);
Poly res = convolution(da, ia);
res.resize(n - 1);
res = integral(res);
res.resize(n);
return res;
}
Poly exponential(const Poly& a, int n) {
assert(a.empty() || a[0] == Mint(0));
if (n == 0) return {};
Poly r{Mint(1)};
for (int len = 1; len < n; len <<= 1) {
int lim = len << 1;
Poly lr = logarithm(r, lim);
Poly f = cut(a, lim);
f = sub(f, lr);
f[0] += Mint(1);
r = convolution(r, f);
r.resize(lim);
}
r.resize(n);
return r;
}
Mint sqrt_mod(Mint a) {
if (a == Mint(0)) return Mint(0);
if (a.pow((mod - 1) / 2) != Mint(1)) {
return Mint(-1);
}
if (mod % 4 == 3) {
return a.pow((mod + 1) / 4);
}
int q = mod - 1;
int s = 0;
while ((q & 1) == 0) {
q >>= 1;
s++;
}
Mint z = 2;
while (z.pow((mod - 1) / 2) == Mint(1)) z += Mint(1);
Mint c = z.pow(q);
Mint x = a.pow((q + 1) / 2);
Mint t = a.pow(q);
int m = s;
while (t != Mint(1)) {
int i = 1;
Mint tt = t * t;
while (tt != Mint(1)) {
tt *= tt;
i++;
}
Mint b = c.pow(1LL << (m - i - 1));
x *= b;
t *= b * b;
c = b * b;
m = i;
}
return x;
}
Poly square_root(const Poly& a, int n) {
assert(!a.empty());
Mint r0 = sqrt_mod(a[0]);
if (r0 == Mint(-1)) return {};
Poly r{r0};
Mint inv2 = Mint(2).inv();
for (int len = 1; len < n; len <<= 1) {
int lim = len << 1;
Poly ir = inverse(r, lim);
Poly t = convolution(cut(a, lim), ir);
t.resize(lim);
r.resize(lim);
for (int i = 0; i < lim; i++) {
t[i] += r[i];
t[i] *= inv2;
}
r = t;
}
r.resize(n);
return r;
}
Poly power(const Poly& a, long long k, int n) {
if (n == 0) return {};
if (k == 0) {
Poly r(n);
r[0] = Mint(1);
return r;
}
int pos = 0;
while (pos < (int)a.size() && a[pos] == Mint(0)) pos++;
if (pos == (int)a.size()) {
return Poly(n, Mint(0));
}
long long shift = 1LL * pos * k;
if (shift >= n) {
return Poly(n, Mint(0));
}
int m = n - (int)shift;
Mint c = a[pos];
Poly b;
for (int i = pos; i < (int)a.size() && (int)b.size() < m; i++) {
b.push_back(a[i] / c);
}
b.resize(m);
Poly lb = logarithm(b, m);
for (auto& x : lb) x *= Mint(k);
Poly r = exponential(lb, m);
Mint factor = c.pow(k);
for (auto& x : r) x *= factor;
Poly ans(n, Mint(0));
for (int i = 0; i < (int)r.size() && i + shift < n; i++) {
ans[i + shift] = r[i];
}
return ans;
}
Poly mul_cut(const Poly& a, const Poly& b, int limit) {
Poly c = convolution(a, b);
if ((int)c.size() > limit) c.resize(limit);
return c;
}
Poly binpow(Poly base, long long exponent, int limit) {
Poly result{Mint(1)};
while (exponent > 0) {
if (exponent & 1) {
result = mul_cut(result, base, limit);
}
exponent >>= 1;
if (exponent > 0) {
base = mul_cut(base, base, limit);
}
}
return result;
}
pair<Poly, Poly> divmod(Poly a, Poly b) {
a = trim(a);
b = trim(b);
if (a.size() < b.size()) {
return {Poly{Mint(0)}, a};
}
int n = a.size();
int m = b.size();
int k = n - m + 1;
Poly ra = a;
Poly rb = b;
reverse(ra.begin(), ra.end());
reverse(rb.begin(), rb.end());
Poly q = convolution(ra, inverse(rb, k));
q.resize(k);
reverse(q.begin(), q.end());
Poly r = sub(a, convolution(b, q));
r.resize(m - 1);
r = trim(r);
return {trim(q), r};
}
Poly modulo(Poly a, const Poly& b) {
return divmod(a, b).second;
}
Mint evaluate(const Poly& a, Mint x) {
Mint res = 0;
for (int i = (int)a.size() - 1; i >= 0; i--) {
res = res * x + a[i];
}
return res;
}
}
using namespace poly;
int main() {
Poly a = {1, 2, 3};
Poly b = {4, 5};
Poly c = mul(a, b);
Poly inv_a = inverse(a, 10);
Poly ln_a = logarithm(a, 10);
Poly exp_a = exponential(Poly{0, 1}, 10);
Poly sqrt_a = square_root(a, 10);
Poly pow_a = power(a, 5, 10);
return 0;
}
需要注意:
- logarithm(a, n) 要求 a[0] != 0。
- exponential(a, n) 要求 a[0] == 0。
- square_root(a, n) 要求常数项存在模意义下的平方根。
- 所有多项式均按低次到高次存储。
- 998244353 支持长度不超过 (2^{23}) 的标准 NTT。
- 这份模板的核心复杂度通常是 (O(n\log n)) 到 (O(n\log^2 n))。
haze

浙公网安备 33010602011771号