求助
查看原帖
求助
731608
TeraniRetZiger楼主2022/6/29 13:04

样例都过不去......

#include <bits/stdc++.h> 

using namespace std;

typedef long long ll;

const int MAXN = 3e5 + 10;

const int mod = 998244353;
const int gw = 3, invgw = 332748118;

inline 
ll qpow(ll b, ll p) {
	ll res = 1;
	while (p) {
		if (p & 1) res = res * b % mod;
		b = b * b % mod, p >>= 1;
	}
	return res;
}

int rev[MAXN];

inline 
int getrev(int n) {
	int l = 1;
	while (l < n << 1) l <<= 1;
	for (int i = 1; i < l; i++) rev[i] = (rev[i >> 1] >> 1) | (i & 1 ? l >> 1 : 0);
	return l;
}

inline 
void ntt(ll *f, int n, int t) {
	for (int i = 0; i < n; i++) {
		if (i < rev[i]) swap(f[i], f[rev[i]]);
	}
	for (int i = 1; i < n; i <<= 1) {
		int w = qpow(t ? invgw : gw, (mod - 1) / (i << 1));
		for (int j = 0; j < n; j += (i << 1)) {
			ll wn = 1, x;
			for (int k = j; k < i + j; k++) {
				x = wn * f[i + k] % mod;
				f[i + k] = (f[k] - x + mod) % mod;
				f[k] = (f[k] + x) % mod;
				wn = wn * w % mod;
			}
		}
	}
	if (t) {
		int inv = qpow(n, mod - 2);
		for (int i = 0; i < n; i++) f[i] = f[i] * inv % mod;
	}
}

ll tf[MAXN];

void finv(ll *f, ll *g, int n) {
	if (n == 1) return g[0] = qpow(f[0], mod - 2), void();
	finv(f, g, n + 1 >> 1);
	int l = getrev(n);
	for (int i = 0; i < n; i++) tf[i] = f[i];
	for (int i = n; i < l; i++) tf[i] = 0;
	ntt(tf, l, 0), ntt(g, l, 0);
	for (int i = 0; i < l; i++) g[i] = (2 - tf[i] * g[i] % mod + mod) % mod * g[i] % mod;
	ntt(g, l, 1);
	for (int i = n; i < l; i++) g[i] = 0;
}

inline 
void drt(ll *f, ll *g, int n) {
	for (int i = 1; i < n; i++) g[i - 1] = i * f[i] % mod;
	g[n - 1] = 0;
}

inline 
void igt(ll *f, ll *g, int n) {
	g[0] = 0;
	for (int i = 1; i < n; i++) g[i] = f[i - 1] * qpow(i, mod - 2) % mod;
}

ll a[MAXN], b[MAXN];

inline 
void fln(ll *f, ll *g, int n) {
	drt(f, a, n), finv(f, b, n);
	int l = getrev(n);
	ntt(a, l, 0), ntt(b, l, 0);
	for (int i = 0; i < l; i++) a[i] = a[i] * b[i] % mod;
	ntt(a, l, 1), igt(a, g, n);
}

ll lng[MAXN];

inline 
void fexp(ll *f, ll *g, int n) {
	if (n == 1) return g[0] = 1, void();
	fexp(f, g, n + 1 >> 1), fln(g, lng, n);
	int l = getrev(n);
	for (int i = 0; i < n; i++) lng[i] = (f[i] + mod - lng[i]) % mod;
	for (int i = n; i < l; i++) g[i] = lng[i] = 0;
	lng[0]++;
	ntt(lng, l, 0), ntt(g, l, 0);
	for (int i = 0; i < l; i++) g[i] = g[i] * lng[i] % mod;
	ntt(g, l, 1);
	for (int i = n; i < l; i++) g[i] = lng[i] = 0;
}

int n;

ll f[MAXN], g[MAXN];

int main() {
    scanf("%d", &n);
    for (int i = 0; i < n; i++) scanf("%lld", &f[i]);
    fexp(f, g, n);
    for (int i = 0; i < n; i++) printf("%lld ", g[i]);
}
2022/6/29 13:04
加载中...