1 条题解

  • 0
    @ 2026-1-16 23:18:39

    #include <bits/stdc++.h>
    #define EB emplace_back
    #define lg2(x) (31 - __builtin_clz(x))
    
    const int N = 135400, mod = 998244353, half_mod = (mod + 1) / 2, root = 31;
    typedef long long ll;
    typedef std::vector <int> vector;
    typedef int vec[N], *pvec;
    
    vec fact, inv, finv;
    
    inline int & reduce(int &x) {return x += (x >> 31 & mod);}
    inline ll & reduce(ll &x) {return x += (x >> 63 & mod);}
    inline int & half(int &x) {return x = (x >> 1) + (-(x & 1) & half_mod);}
    inline ll & half(ll &x) {return x = (x >> 1) + (-(x & 1) & half_mod);}
    inline int & neg(int &x) {return x = (!x - 1) & (mod - x);}
    inline ll & neg(ll &x) {return x = (!x - 1) & (mod - x);}
    ll PowerMod(ll a, int n, ll c = 1) {for (; n; n >>= 1, a = a * a % mod) if (n & 1) c = c * a % mod; return c;}
    
    void init() {
    	int i;
    	for (inv[1] = 1, i = 2; i < N; ++i) inv[i] = (ll)(mod - mod / i) * inv[mod % i] % mod;
    	for (*finv = *fact = i = 1; i < N; ++i) fact[i] = (ll)fact[i - 1] * i % mod, finv[i] = (ll)finv[i - 1] * inv[i] % mod;
    }
    
    namespace Poly {
    	int l, n;
    	vec rev, x, y;
    
    	void NTT_init(int len) {
    		if (l == len) return; n = 1 << (l = len);
    		ll g = PowerMod(root, 1 << (23 - l));
    		*x = 1, *rev = 0;
    		for (int i = 1; i < n; ++i)
    			x[i] = x[i - 1] * g % mod, rev[i] = rev[i >> 1] >> 1 | (i & 1) << (l - 1);
    	}
    
    	void DNTT(int *d, int *t) {
    		int i, *j, *k, len = 1, delta = n, R;
    		for (i = 0; i < n; ++i) t[rev[i]] = d[i];
    		for (i = 0; i < l; ++i) {
    			delta >>= 1;
    			for (k = x, j = y; j < y + len; k += delta, ++j) *j = *k;
    			for (j = t; j < t + n; j += len << 1)
    				for (k = j; k < j + len; ++k) {
    					R = (ll)y[k - j] * k[len] % mod;
    					k[len] = (*k - R < 0 ? *k - R + mod : *k - R);
    					*k = (*k + R >= mod ? *k + R - mod : *k + R);
    				}
    			len <<= 1;
    		}
    	}
    
    	vec B1, B2, B3, B4, B5, B6, B7;
    
    	void Mul(int deg, pvec a, pvec b, pvec c) {
    		if (!deg) {*c = (ll)*a * *b % mod; return;}
    		NTT_init(lg2(deg) + 1);
    		int i; ll iv = mod - (mod - 1) / n;
    		DNTT(a, c), DNTT(b, B1);
    		for (i = 0; i < n; ++i) B1[i] = (ll)B1[i] * c[i] % mod;
    		DNTT(B1, c), std::reverse(c + 1, c + n);
    		for (i = 0; i < n; ++i) c[i] = c[i] * iv % mod;
    	}
    
    	void Inv(int deg, pvec a, pvec b) {
    		int len, i; ll iv = half_mod;
    		*b = PowerMod(*a, mod - 2), b[1] = 0;
    		*B1 = *a, B1[1] = a[1];
    		for (len = 0; 1 << len < deg; ++len) {
    			NTT_init(len + 2);
    			for (i = n >> 1; i < n; ++i) b[i] = B1[i] = 0;
    			DNTT(b, B2), DNTT(B1, B3);
    			for (i = 0; i < n; ++i) reduce(B2[i] = B2[i] * (2ll - (ll)B2[i] * B3[i] % mod) % mod);
    			DNTT(B2, B3), std::reverse(B3 + 1, B3 + n), half(iv);
    			for (i = 0; i < n >> 1; ++i) b[i] = B3[i] * iv % mod;
    			for (; i < n; ++i) b[i] = 0, B1[i] = a[i];
    		}
    	}
    
    	void Mul(vector &a, vector &b, vector &ret) {
    		int degA = a.size() - 1, degB = b.size() - 1;
    		if (!(degA || degB)) {ret.emplace_back((ll)a[0] * b[0] % mod); return;}
    		NTT_init(lg2(degA + degB) + 1);
    		int i; ll iv = mod - (mod - 1) / n;
    		memcpy(B1, a.data(), (degA + 1) << 2), memset(B1 + (degA + 1), 0, (n - degA - 1) << 2);
    		memcpy(B2, b.data(), (degB + 1) << 2), memset(B2 + (degB + 1), 0, (n - degB - 1) << 2);
    		DNTT(B1, B3), DNTT(B2, B1);
    		for (i = 0; i < n; ++i) B1[i] = (ll)B1[i] * B3[i] % mod;
    		DNTT(B1, B3), std::reverse(B3 + 1, B3 + n); ret.clear(), ret.reserve(degA + degB + 1);
    		for (i = 0; i <= degA + degB; ++i) ret.EB(B3[i] * iv % mod);
    	}
    
    	void Diff(int deg, vec a, vec b) {for (int i = 1; i <= deg; ++i) b[i - 1] = (ll)a[i] * i % mod;}
    
    	void Intg(int deg, vec a, vec b, int ct = 0) {for (int i = 1; i <= deg; ++i) b[i] = (ll)a[i - 1] * inv[i] % mod, *b = ct;}
    
    	void Ln(int deg, vec a, vec b, bool need_reintegrate = true) {
    		deg -= need_reintegrate, assert(deg);
    		int i, j = deg * 2 - 1; NTT_init(lg2(j) + 1);
    		Diff(deg, a, B4), Inv(deg, a, B5);
    		for (i = deg; i < n; ++i) B4[i] = B5[i] = 0;
    		Mul(j, B4, B5, B6);
    		if (need_reintegrate) Intg(deg, B6, b);
    		else memcpy(b, B6, deg << 2);
    	}
    
    	void Exp(int deg, vec a, vec b) {
    		int len, i, n = 2;
    		*b = 1, b[1] = 0;
    		for (len = 0; 1 << len < deg; ++len, n <<= 1) {
    			Ln(n, b, B7); *B7 = 1;
    			for (i = 1; i < n; ++i) reduce(B7[i] = a[i] - B7[i]);
    			for (; i < n << 1; ++i) B7[i] = b[i] = 0;
    			Mul((n << 1) - 1, b, B7, B6);
    			for (i = 0; i < n; ++i) b[i] = B6[i];
    			for (; i < n << 1; ++i) b[i] = 0;
    		}
    	}
    }
    
    int n, K, cnt;
    vec A, B, invA, lnA;
    vec a, F, G, lhs, rhs;
    vector g[N];
    
    int solve(int L, int R) {
    	if (L + 1 == R) return L;
    	int M = (L + R) / 2, id = cnt++, lp = solve(L, M), rp = solve(M, R);
    	return Poly::Mul(g[lp], g[rp], g[id]), id;
    }
    
    void calc_deg_sum() {
    	int i, id;
    	for (i = 0; i < n; ++i) g[i].EB(1), g[i].EB(mod - a[i]);
    	cnt = n, id = solve(0, n);
    	memcpy(G, g[id].data(), (n + 1) << 2);
    	Poly::Ln(n, G, F + 1, false);
    	for (*F = n, i = 1; i <= n; ++i) neg(F[i]);
    }
    
    int main() {
    	int i, j, ans = 0;
    	scanf("%d%d", &n, &K), init();
    	for (i = 0; i < n; ++i) scanf("%d", a + i);
    	if (n <= 1) return putchar(48 + !K), putchar(10), 0;
    	if (n == 2) return printf("%lld\n", 2ll * *a * a[1] % mod);
    	calc_deg_sum();
    	for (i = 0; i < n - 1; ++i) A[i] = PowerMod(i + 1, K, finv[i]), B[i] = PowerMod(i + 1, 2 * K, finv[i]);
    	Poly::Ln(n, A, lnA), lnA[n - 1] = 0;
    	memcpy(invA, Poly::B5, (n - 1) << 2);
    	for (i = 0; i < n - 1; ++i) lnA[i] = (ll)lnA[i] * F[i] % mod;
    	Poly::Exp(n - 1, lnA, rhs);
    	Poly::Mul(2 * (n - 2), B, invA, lhs);
    	for (i = 0; i < n - 1; ++i) lhs[i] = (ll)lhs[i] * F[i] % mod;
    	for (i = 0, j = n - 2; i < n - 1; ++i, --j) ans = (ans + (ll)lhs[i] * rhs[j]) % mod;
    	for (i = 0; i < n; ++i) ans = (ll)ans * a[i] % mod;
    	printf("%lld\n", (ll)ans * fact[n - 2] % mod);
    	return 0;
    }
    
    
    • 1

    信息

    ID
    4713
    时间
    5000ms
    内存
    1024MiB
    难度
    10
    标签
    递交数
    1
    已通过
    1
    上传者