AIGC标识 给我一份全一点的多项式模板

下面是一份以 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))。
posted @ 2026-10-04 17:35  haze1231  阅读(5)  评论(0)    收藏  举报