基础多项式模板

· · 个人记录

使用 DIF-FFT 实现。

#include <cstdio>
#include <iostream>
#include <cstring>
#include <cmath>
#include <algorithm>
#include <cassert>
#include <ctime>
#include <queue>
#include <vector>
#include <set>

const int N = 400011;
const int LOG = log2(N) + 1;
const int p = 998244353;
typedef std::vector <int> poly;

char buf[1 << 25] ,*p1 = buf ,*p2 = buf;
#define getchar() (p1 == p2 && (p2 = (p1 = buf) + fread(buf ,1 ,1 << 21 ,stdin) ,p1 == p2) ? EOF : *p1++)
inline int read() {
    int x = 0, f = 1; char ch = getchar();
    while(ch > '9' || ch < '0') { if(ch == '-') f = -1; ch = getchar(); }
    while(ch >= '0' && ch <= '9') x = x * 10 + ch - 48, ch = getchar();
    return x * f;
}

inline int add(int a,int b) { return a + b >= p ? a + b - p : a + b; }
inline int dec(int a,int b) { return a - b < 0 ? a - b + p : a - b; }
inline int mul(int a,int b) { return 1ll * a * b % p; }
inline int fastpow(int a,int b) {
    int r = 1;
    while(b) {
        if(b & 1) r = mul(r,a);
        a = mul(a,a), b >>= 1;
    }
    return r;
}

int rt[N],inv[N];
void Init() {
    rt[0] = 1, rt[1 << (LOG - 1)] = fastpow(3,(p - 1) >> (LOG + 1));
    for(int i = LOG - 2;i >= 0;i--) rt[1 << i] = mul(rt[1 << (i + 1)],rt[1 << (i + 1)]);
    for(int i = 1;i < N;i++) rt[i] = mul(rt[i & (i - 1)],rt[i & -i]);
    inv[1] = 1;
    for(int i = 2;i < N;i++) inv[i] = mul(inv[p % i],p - p / i);
}

poly resz(poly a,int n) { return a.resize(n), a; }

void NTT(poly &a) {
    int N = a.size(), t;
    for(int n = N, m = n >> 1;m;n = m, m >>= 1)
        for(int l = 0, x = 0;l < N;l += n, x++)
            for(int i = l;i < l + m;i++) {
                t = mul(rt[x],a[i + m]);
                a[i + m] = dec(a[i],t), a[i] = add(a[i],t);
            }
}

void INTT(poly &a) {
    int N = a.size(), t;
    for(int n = 2, m = 1;n <= N;m = n, n <<= 1)
        for(int l = 0, x = 0;l < N;l += n, x++)
            for(int i = l;i < l + m;i++) {
                t = dec(a[i],a[i + m]);
                a[i] = add(a[i],a[i + m]);
                a[i + m] = mul(t,rt[x]);
            }
    std::reverse(a.begin() + 1,a.end());
    for(int i = 0;i < N;i++) a[i] = mul(a[i],inv[N]);
}

poly Mul(poly a,poly b,int modxn = 1) {
    int n = a.size(), m = b.size();
    if(n <= 40 && m <= 40) {
        poly c(n + m - 1);
        for(int i = 0;i < n;i++) for(int j = 0;j < m;j++)
            c[i + j] = add(c[i + j],mul(a[i],b[j]));
        if(modxn) c.resize(n);
        return c;
    }
    int N = 1; while(N <= n + m) N <<= 1;
    a.resize(N), b.resize(N), NTT(a), NTT(b);
    for(int i = 0;i < N;i++) a[i] = mul(a[i],b[i]);
    INTT(a), a.resize(modxn ? n : n + m - 1);
    return a;
}

poly Inv(poly a) {
    int m = a.size();
    poly res(1,fastpow(a[0],p - 2)), tmp;
    for(int n = 2;n < 2 * m;n <<= 1) {
        if(n > m) n = m;
        tmp = poly(a.begin(),a.begin() + n);
        int N = 1; while(N <= 2 * n) N <<= 1;
        tmp.resize(N), res.resize(N), NTT(tmp), NTT(res);
        for(int i = 0;i < N;i++) res[i] = dec(mul(2,res[i]),mul(tmp[i],mul(res[i],res[i])));
        INTT(res), res.resize(n);
    }
    return res;
}

poly deri(poly a) {
    int n = a.size();
    for(int i = 0;i < n - 1;i++) a[i] = mul(a[i + 1],i + 1);
    return a.resize(n - 1), a;
}

poly intg(poly a) {
    int n = a.size(); a.resize(n + 1);
    for(int i = n;i >= 1;i--) a[i] = mul(a[i - 1],inv[i]);
    return a[0] = 0, a;
}

poly Ln(poly a) {
    assert(a[0] == 1);
    return intg(Mul(Inv(a),deri(a)));
}

poly Exp(poly a) {
    assert(a[0] == 0);
    int m = a.size();
    poly res(1,1), tmp;
    for(int n = 2;n < 2 * m;n <<= 1) {
        if(n > m) n = m;
        res.resize(n);
        tmp = res, res = Ln(res), res[0] = dec(res[0],1);
        assert(res[0] == 998244352);
        assert(a[0] == 0);
        for(int i = 0;i < n;i++) res[i] = dec(a[i],res[i]);
        assert(res[0] == 1);
        assert(tmp[0] == 1);
        res = Mul(tmp,res);
    }
    return res;
}

poly reverse(poly a) { return std::reverse(a.begin(),a.end()), a; }
poly Mod(poly a,poly b) {
    int n = a.size(), m = b.size();
    if(n < m) return a;
    poly f = reverse(a), g = reverse(b);
    f.resize(n - m + 1), g.resize(n - m + 1);
    poly q = Mul(f,Inv(g)); q = reverse(q);
    b = Mul(q,b);
    for(int i = 0;i < n;i++) a[i] = dec(a[i],b[i]);
    if(!a.empty()) while(!*--a.end()) a.erase(--a.end());
    return a;
}

namespace Multipoint {

    #define lc(k) k << 1
    #define rc(k) k << 1 | 1
    poly Q[N];

    poly MulT(poly a,poly b) {
        int n = a.size(), m = b.size(); 
        std::reverse(b.begin(),b.end()), b = Mul(a,b,0);
        for(int i = 0;i < n;i++) a[i] = b[i + m - 1];
        return a;
    }

    void Init(poly &a,int k,int l,int r) {
        if(l == r) return void(Q[k] = poly{1,dec(0,a[l])});
        int m = (l + r) / 2;
        Init(a,lc(k),l,m), Init(a,rc(k),m + 1,r);
        Q[k] = Mul(Q[lc(k)],Q[rc(k)],0);
    }

    void Multipoint(int k,int l,int r,poly F,poly &g) {
        F.resize(r - l + 1);
        if(l == r) return void(g[l] = F[0]);
        int m = (l + r) / 2;
        Multipoint(lc(k),l,m,MulT(F,Q[rc(k)]),g);
        Multipoint(rc(k),m + 1,r,MulT(F,Q[lc(k)]),g);
    }

    poly Solve(poly f,poly a) {
        int n = a.size(), m = f.size();
        Init(a,1,0,n - 1);
        poly v(n);
        Multipoint(1,0,n - 1,MulT(f,Inv(resz(Q[1],m))),v);
        return v;
    }
}

poly f,a,res;

int main() {
    Init();
    int n = read() + 1, m = read();
    for(int i = 0;i < n;i++) f.push_back(read());
    for(int i = 0;i < m;i++) a.push_back(read());
    res = Multipoint::Solve(f,a);
    for(int i = 0;i < m;i++) std::printf("%d\n",res[i]);
    return 0;
}