1 条题解
-
0

#include <bits/stdc++.h> #define lg2 std::__lg #define EB emplace_back using std::cin; using std::cout; typedef long long ll; const int N = 530000, mod = 998244353, iv2 = (mod + 1) / 2, root = 31; typedef int vec[N], *pvec; typedef std::vector <int> vector; vec inv, harm; inline void up(int &x, const int y) {x < y ? x = y : 0;} inline void add(int &x, const int y) {x += y - mod, x += x >> 31 & mod;} inline int & reduce(int &x) {return x += x >> 31 & mod;} inline int & neg(int &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] = harm[1] = 1, i = 2; i < N; ++i) inv[i] = ll(mod - mod / i) * inv[mod % i] % mod, add(harm[i] = harm[i - 1], inv[i]); } 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(vector &a, vector &b, vector &ret) { int degA = a.size() - 1, degB = b.size() - 1; if (!(degA || degB)) {ret.EB((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, pvec a, pvec b) {for (int i = 1; i <= deg; ++i) b[i - 1] = (ll)a[i] * i % mod;} void Exp(int deg, pvec a, pvec b) { int i, len; ll iv = iv2; pvec c = B7; assert(!*a); if (*b = 1, deg <= 1) return; if (b[1] = a[1], deg == 2) return; memset(b + 2, 0, i = 8 << lg2(deg - 1)), memset(c, 0, i), memset(B1, 0, i), *c = 1, neg(c[1] = b[1]); for (len = 1; 1 << len < deg; ++len) { NTT_init(len + 1), iv = (iv >> 1) + iv2; DNTT(c, B2), DNTT(b, B3); for (i = 0; i < n; ++i) B4[i] = (ll)B3[i] * B2[i] % mod; DNTT(B4, B5); for (i = n >> 1; i < n; ++i) B5[i] = B5[n - i] * iv % mod; memset(B5, 0, n << 1), DNTT(B5, B4); for (i = 0; i < n; ++i) B4[i] = (ll)B4[i] * B2[i] % mod; DNTT(B4, B5); for (i = n >> 1; i < n; ++i) B5[i] = B5[n - i] * (mod - iv) % mod; memcpy(B5, c, n << 1), DNTT(B5, B6); Diff(n >> 1, b, B1), DNTT(B1, B4); for (i = 0; i < n; ++i) B4[i] = (ll)B4[i] * B6[i] % mod; DNTT(B4, B6); for (i = n >> 1; i < n; ++i) reduce(B5[i] = (a[i] - B6[n - i + 1] * iv % mod * inv[i]) % mod); memset(B5, 0, n << 1), DNTT(B5, B4); for (i = 0; i < n; ++i) B4[i] = (ll)B4[i] * B3[i] % mod; DNTT(B4, B5); for (i = n >> 1; i < n; ++i) b[i] = B5[n - i] * iv % mod; if (2 << len >= deg) return; DNTT(b, B3); for (i = 0; i < n; ++i) B3[i] = (ll)B3[i] * B2[i] % mod; DNTT(B3, B4); for (i = n >> 1; i < n; ++i) B4[i] = B4[n - i] * iv % mod; memset(B4, 0, n << 1), DNTT(B4, B3); for (i = 0; i < n; ++i) B3[i] = (ll)B3[i] * B2[i] % mod; DNTT(B3, B4); for (i = n >> 1; i < n; ++i) c[i] = B4[n - i] * (mod - iv) % mod; } } } int n, U, K, limit, cnt = 0; vec b, c, f, I, J; vec E0, E1, E2, C1, C2, C3, C4, C5, C6, C7; vector g[N]; void CDQ(int L, int w) { int i, R = L + (1 << w), M; ll iv; if (!w) { if (L) { reduce(I[L] = (ll(E1[L] - E2[L]) * inv[L] - (L == K) + (L == K + 1)) % mod); E0[L] = (ll)L * I[L] % mod; E1[L] = ((ll)E1[L] * inv[L] + (ll)K * I[L]) % mod; E2[L] = ((ll)E2[L] * inv[L] + (K + 1ll) * I[L]) % mod; } else *E1 = *E2 = 1; return; } CDQ(L, w - 1); if ((M = (1 << (w - 1)) + L) > limit) return; Poly::NTT_init(w), iv = mod - (mod - 1) / Poly::n; if (L) { memcpy(C2, E0, 4 << w), memcpy(C5, E0 + L, 2 << w), memset(C5 + (1 << (w - 1)), 0, 2 << w), memcpy(C3, E1, 4 << w), memcpy(C6, E1 + L, 2 << w), memset(C6 + (1 << (w - 1)), 0, 2 << w), memcpy(C4, E2, 4 << w), memcpy(C7, E2 + L, 2 << w), memset(C7 + (1 << (w - 1)), 0, 2 << w), Poly::DNTT(C2, C1), Poly::DNTT(C3, C2), Poly::DNTT(C4, C3), Poly::DNTT(C5, C4), Poly::DNTT(C6, C5), Poly::DNTT(C7, C6); for (i = 0; i < Poly::n; ++i) C2[i] = ((ll)C1[i] * C5[i] + (ll)C2[i] * C4[i]) % mod, C3[i] = ((ll)C1[i] * C6[i] + (ll)C3[i] * C4[i]) % mod; Poly::DNTT(C2, C1), Poly::DNTT(C3, C2); for (i = M; i < R; ++i) E1[i] = (E1[i] + C1[Poly::n - (i - L)] * iv % mod * K) % mod, E2[i] = (E2[i] + C2[Poly::n - (i - L)] * iv % mod * (K + 1)) % mod; } else { memcpy(C2, E0, 2 << w), memset(C2 + M, 0, 2 << w), memcpy(C3, E1, 2 << w), memset(C3 + M, 0, 2 << w), memcpy(C4, E2, 2 << w), memset(C4 + M, 0, 2 << w), Poly::DNTT(C2, C1), Poly::DNTT(C3, C2), Poly::DNTT(C4, C3); for (i = 0; i < Poly::n; ++i) C2[i] = (ll)C1[i] * C2[i] % mod, C3[i] = (ll)C1[i] * C3[i] % mod; Poly::DNTT(C2, C1), Poly::DNTT(C3, C2); for (i = M; i < R; ++i) E1[i] = (E1[i] + C1[Poly::n - i] * iv % mod * K) % mod, E2[i] = (E2[i] + C2[Poly::n - i] * iv % mod * (K + 1)) % mod; } CDQ(M, w - 1); } 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 get_poly(int n, vector &ret) { int i, deg = n + K; memset(J, 0, 1 << (lg2(deg - 1) + 3)); for (i = 0; i < deg; ++i) J[i] = (ll)I[i] * deg % mod; Poly::Exp(deg, J, f), ret.clear(), ret.reserve(deg); for (i = 0; i < deg; ++i) ret.EB(f[i] * ll(deg - i) % mod * inv[deg] % mod); } int main() { int i, j, x, top = 0, ans = 0, id; init(); std::ios::sync_with_stdio(false), cin.tie(NULL); cin >> n >> K; for (i = 0; i < n; ++i) cin >> x, up(U, x), ++c[x]; for (i = 1; i <= U; ++i) c[i] += c[i - 1]; for (i = K; i <= U; ++i) if (c[i] == c[i - K] + K) b[top++] = i; for (j = 0, i = 1; i <= top; ++i) if (b[i] != b[i - 1] + 1) up(limit, i - j), j = i; limit += K - 1, CDQ(0, 19); for (j = 0, i = 1; i <= top; ++i) if (b[i] != b[i - 1] + 1) get_poly(i - j, g[cnt++]), j = i; id = solve(0, cnt), j = g[id].size(); for (i = 1; i < j; ++i) ans = (ans + (ll)harm[i] * g[id][i]) % mod; cout << int(ans * ll(mod - n) % mod) << '\n'; return 0; }
- 1
信息
- ID
- 1836
- 时间
- 5000ms
- 内存
- 512MiB
- 难度
- 10
- 标签
- 递交数
- 4
- 已通过
- 3
- 上传者