1 条题解
-
0

#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
- 上传者