1 条题解

  • 0
    @ 2026-1-16 22:22:19

    #include <bits/stdc++.h>
    #define lg2 std::__lg
    #define EB emplace_back
    
    typedef long long ll;
    const int N = 530000, mod = 998244353, root = 31;
    typedef int vec[N], *pvec;
    typedef std::vector <int> vector;
    
    vec fact, finv;
    
    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 (*fact = i = 1; i < N; ++i) fact[i] = (ll)fact[i - 1] * i % mod;
    	--i, finv[i] = PowerMod(fact[i], mod - 2);
    	for (; i; --i) finv[i - 1] = (ll)finv[i] * i % mod;
    }
    
    inline ll C(int n, int r) {return (ll)fact[n] * finv[r] % mod * finv[n - r] % 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;
    
    	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 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);
    	}
    }
    
    int n, q, K, ty, cnt;
    vec a, b;
    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;
    }
    
    int main() {
    	int i, j, id, cur, z = 1; ll ans = 0;
    	scanf("%d%d%d", &ty, &n, &K), init();
    	for (i = 0; i < ty; a[i] ? ++i : --ty) scanf("%d", a + i);
    	for (i = 0; i < ty; z = (ll)z * finv[a[i++]] % mod)
    		for (j = 1; j <= a[i]; ++j)
    			g[i].EB(C(a[i] - 1, j - 1) * finv[j] % mod * fact[a[i]] % mod);
    	cnt = ty, id = solve(0, ty), n -= ty;
    	memcpy(b, g[id].data(), (n + 1) << 2);
    	for (i = 0; i <= n; ++i) b[i] = (ll)b[i] * fact[i + ty] % mod * z % mod;
    	for (i = K; i <= n; ++i)
    		cur = (ll)C(i, K) * b[n - i] % mod, (i ^ K) & 1 ? ans -= cur : ans += cur;
    	ans %= mod, ans += ans >> 63 & mod, printf("%d\n", int(ans));
    	return 0;
    }
    
    
    • 1

    信息

    ID
    4679
    时间
    1000ms
    内存
    128MiB
    难度
    9
    标签
    递交数
    14
    已通过
    3
    上传者