题解:P5299 [PKUWC2018] Slay the Spire

· · 题解

[PKUWC2018] Slay the Spire

T4 - P5299 [PKUWC2018] Slay the Spire - 洛谷

开始九条可怜的卡组有 2n 张牌,每张牌上都写着一个数字w_i,一共有两种类型的牌,每种类型各 n 张:攻击牌,造成等于牌上的数字的伤害。强化牌,其他剩下的攻击牌的数字都会乘上 x保证强化牌上的数字都大于 1。现在等概率随机从卡组中抽出 m 张牌,最多打出 k 张牌,永远都会采取能造成最多伤害的策略,求期望造成多少伤害。

假设答案为 \text{ans} ,你只需要输出:\left (\text{ans}\times \frac{(2n)!}{m!(2n-m)!}\right) ~\bmod 998244353

--- 首先我们先观察答案,发现这是个假期望,真计数题(QwQ,赛时被骗了),答案就是所有方案的答案之和。 考虑出牌的顺序,首先我们先把 $a_i$ 和 $b_i$ 按顺序分别从大到小排序,那么先把能力牌打完再打攻击牌一定最优,证明如下: 假设有两张攻击牌 $b_1 \ge b_2$,对于能力牌 $a_1 \ge a_2 \ge 2$,我们只能打出 $3$ 张牌。 比较 $a_1(b_1 + b_2)$ 和 $a_1a_2b_1$ 的大小,$a_1b_1 + a_1b_2 \le a_1a_2b_1$,这个取等条件是苛刻的。 那么我们就可以尝试进行 DP 转移了,对于这种有限制的转移,我们可以钦定一些固定点,也就是某些点的选择情况,我们设如下的状态: $f_{i, j, 0/1}$ 表示在前 $i$ 张牌中打了 $j$ 张,当前的牌 $i$ 是否打了的强化牌的乘积之和。 $g_{i, j, 0/1}$ 表示在前 $i$ 张牌中打了 $j$ 张,当前的牌 $i$ 是否打了的攻击牌的和的和。 这里可以简单的对这种设法进行解释,为什么是乘积之和与和的和。 考虑不同的方案对答案造成的贡献,是类似于这种 $f_1g_1 + f_1g_2 \cdots f_2g_1 + \cdots$,那么对于一种强化牌的方式,可以对多个攻击牌贡献,有类似 $(f_1 + f_2 + f_3 \cdots)(g_1 + g_2 + g_3 \cdots)$ 的类背包贡献,因此我们这样做。 考虑转移,显然有: $$ f_{i, j, 0} \gets f_{i - 1, j, 1} + f_{i - 1, j, 0}\\ f_{i, j, 1} \gets a_i \times (f_{i - 1, j - 1, 1} + f_{i - 1,j - 1, 0})\\ g_{i, j, 0} \gets g_{i - 1, j, 0} + g_{i - 1, j - 1, 1}\\ g_{i, j, 1} \gets g_{i - 1, j - 1, 0} + g_{i - 1, j - 1, 1} + b_i \times \binom{i - 1}{j - 1} $$ 难点在于对于 $g_{i, j, 1}$ 的贡献,对于组合计数的题,加法的贡献我们不能简单的相加,而要求组合数系数,在这个里面,我们的组合数系数就是 $\binom{i - 1}{j - 1}$。 考虑统计答案,我们可以枚举抽中的牌中强化牌的数量为 $i$,根据我们刚刚的结论,需要进行分类讨论: - 当 $i < k$ 时,对答案造成的贡献是: $$ (\sum_{j = 1} ^ n f_{j, i, 1}) (\sum_{j = 1} ^ n g_{j, k - i, 1} \binom{n - j}{m - k}) $$ - 当 $i \ge k$ 时,对答案造成的贡献是: $$ (\sum_{j = 1} ^ n f_{j, k - 1, 1} \binom{n - j}{i - (k - 1)})(\sum_{j = 1} ^ n g_{j, 1, 1} \binom{n - j}{m - i - 1}) $$ 直接做统计即可。 时间复杂度 $\mathcal{O}(n ^ 2)$。 但是我被骗了: ![mx](https://cdn.luogu.com.cn/upload/image_hosting/ym3gy4hv.png) 这个 mx 的 OJ 上显示的 $n$ 达到了 $30000$,于是我认为无法通过。我们考虑进一步的优化。 我们其实不难注意到当 $i < k$ 和 $i \ge k$ 的时候,不止答案的统计是不同的,我们的 $f$ 和 $g$ 事实上也是不同的。 强化牌的贡献,我们可以设为 $E_i$,当 $i < k$ 的时候,我们的强化牌一定可以直接全选,$E_i = \sum_{j = 1} ^ n f_{j, i, 1}$。当 $i \ge k$ 的时候,$E_i = \sum_{j = 1} ^ n f_{j, k - 1, 1} \binom{n - j}{i - k + 1}$。对于攻击牌的贡献,我们可以记为 $A_i$,当 $i < k$ 的时候,$A_i = \sum_{j = 1} ^ n g_{j, k - i, 1} \binom{n - j}{m - k}$。而当 $i \ge k$ 的时候,$A_i = \sum_{j = 1} ^ n g_{j, 1, 1} \binom{n - j}{m - i - 1}$。 那么我们的最终答案可以记为 $\sum E_i \times A_i$。 那么我们考虑用更低的复杂度去计算 $E_i$ 和 $A_i$。 **对于攻击牌:** 我们先考虑 $A_i$ 的优化,我们开始拆式子,定义 $g_{j, 1, 1} = b_j$,**那么 $i \ge k$ 的时候**,拆式子: $$ A_i = \sum_{j = 1} ^ n b_j \binom{n - j}{m - i - 1} \\ =\sum_{j = 1} ^ n b_j \frac{(n - j) !}{(m - i - 1) !(n - j - (m - i - 1)) !} $$ 我们把与式子无关的 $\frac{1}{(m - i - 1) !}$ 提出来,那么可以得到: $$ A_i = \frac{1}{(m - i - 1) !} \sum_{j = 1} ^ n b_j \frac{(n - j) !}{((n - j) - (m - i - 1))!}\\ = \frac{1}{(m - i - 1) !} \sum_{j = 1} ^ n (b_j(n - j)!) \frac{1}{((n - j) - (m - i - 1))!}\\ $$ 对于右边,我们发现 $(n - j) - ((n - j) - (m - i - 1)) = m - i - 1$,这个部分的下标相减为差值,因此我们直接做一次减法卷积,最后第 $n - (m - i)$ 项就是求和的值,要乘上前面的 $\frac{1}{(m - i - 1) !}$。这里我们把算出来的数列记为 $c$,我们记其第 $t$ 项为 $\sum_j b_j \binom{n - j}{t}$,或者说 $A_i = c_{m - i -1}$。 **对于 $i < k$ 的时候**,拆式子: $$ A_i = \sum_{j = 1} ^ n g_{j, k - i, 1} \binom{n - j}{m - k} $$ 这个式子就是在 $m - i$ 张牌中选择最大的 $k - i$ 张,我们可以用刚刚的 $c$ 数列改写式子: $$ A_i = \sum_{j = m - k} ^ {m - i - 1} c_j \times (-1) ^ {j - (m - k)} \binom{j - 1}{(m - k) - 1}\binom{n - i - j}{(m - i -1) - j}\\ $$ 我们拆一下后面的那个组合数: $$ \binom{n - 1 - j}{(m - i -1) - j} = \frac{(n - i - j) !}{(m - i - 1 - j) ! (n - m + i)!} $$ 我们整理一下整个式子,则: $$ A_i = \frac{1}{(n - m + i) !} \sum_{j = m -k} ^ {m - i - 1} c_j \times (-1) ^ {j - (m - k)} \binom{j - 1}{m - k - 1}(n - 1 - j)! \times \frac{1}{(m - i - 1 - j)!} $$ $j + (m - i - 1 - j) = m - i - 1$,这个部分直接对后面做加法卷积即可,最后还是乘一下前面的系数,提取第 $m - i -1 $ 项。 **对于强化牌:** **当 $i < k$ 的时候**,对于 $E_i = \sum_{j = 1} ^ n f_{j,i,1}$,我们看之前 $n ^ 2$ 的转移式子,每一个当前的值都是前面一项选或者不选当前的牌 $j$ 而得到的系数 $a_j$,这个选或不选直接用我们二项式定理的本质即可改写,选 $i$ 张牌的乘积和,就是把所有的 $(1 + a_jx)$ 乘开之后第 $x ^ i$ 前面的系数: $$ E_i = [x ^ i] \prod_{j = 1} ^n (1 + a_jx) $$ 这个式子直接分治 NTT 就做完了。 **当 $i \ge k$ 的时候**,我们观察原式子 $E_i = \sum_{j = 1} ^ n f_{j, k - 1, 1} \binom{n - j}{i - k + 1}$,根据定义,这个 $f_{j, k - 1, 1}$ 指在第 $j$ 个位置选择这张牌,且一共选择 $k - 1$ 张牌,所以其应该是前 $j - 1$ 张牌的乘积当中 $x ^ {k - 2}$ 的系数。 $$ f_{j, k - 1, 1} = a_j \times [x ^ { k - 2}] \prod_{l = 1} ^ {j - 1} (1 + a_lx) $$ 好的我们发现无从下手,但是对于多项式的推导,有一位古人说过,不是求导就是差分,这个前缀的形式显然我们可以尝试选择出差分: $$ a_j x \prod _{l = 1} ^ { j - 1} (1 + a_lx) = (1 + a_j x) \prod_{l = 1} ^ {j - 1} - \prod_{l = 1} ^{j - 1}(1 + a_l x) \\ = \prod_{l = 1} ^ j (1 + a_lx) - \prod _{l = 1} ^ {j - 1} (1 + a_l x) $$ 我们定义 $C_j = [x ^ {k - 1}] \prod_{l = 1} ^ j (1 + a_l x)$,那么我们就可以得出 $f_{j, k - 1, 1} = C_j - C_{j - 1}$。对于这个 $C_j$ 来说,我们在线段树自顶向底打标记就能求。这时候我们把求出来的 $f_{j, k - 1, 1}$ 带回原式子,先把右边组合数拆一下: $$ \binom{n - j}{i - k+ 1} = \frac{(n - j) !}{(i - k + 1) ! ((n - j) - (i - k + 1))!}\\ $$ 稍微整理一下各式子: $$ E_i = \frac{1}{(i - k + 1) !}\sum_{j = 1} ^ n (f_{j, k - 1, 1} \times (n - j) !) \frac{1}{((n - j) - (i - k + 1))!} $$ 依旧看下标,$(n - j) - ((n - j) - (i - k + 1)) = i - k+ 1$,这个还是对这玩意做一次减法卷积,最后还是乘一下前面的系数。 我们现在的所有 $A_i$ 和 $E_i$ 已经全部求出来了,直接合并答案即可。 时间复杂度 $\mathcal{O}(n \log ^ 2 n)$。 :::success[CODE 参考实现] ```cpp #include<bits/stdc++.h> using namespace std; #ifdef LOCAL #include<algo/debug.h> #else #define debug(...) 42 #endif using ll = long long; using ull = unsigned long long; using f64 = double; using f128 = long double; using pii = pair<int, int>; using pll = pair<ll, ll>; using vi = vector<int>; using vll = vector<ll>; #define pb emplace_back #define mk make_pair #define all(x) (x).begin(), (x).end() #define rall(x) (x).rbegin(), (x).rend() #define sz(x) (int)((x).size()) #define ciallo(x) cerr << (x) << '\n'; template <typename T, typename U> inline bool chmin(T& a, const U& b){return (b < a ? a = b, true : false);} template <typename T, typename U> inline bool chmax(T& a, const U& b){return (a < b ? a = b, true : false);} const int N = 3005; const int M = 1 << 19; const int P = 998244353; int n, m, k; int a[N], b[N]; int E[N], A[N], c[N]; namespace base{ int qpow(int a, int b){ int res = 1; while(b){ if(b & 1) res = (1ll * res * a) % P; a = (1ll * a * a) % P; b >>= 1; } return res; } int fac[N], inv[N]; void init(){ fac[0] = 1; for(int i = 1;i < N;i ++) fac[i] = 1ll * fac[i - 1] * i % P; inv[N - 1] = qpow(fac[N - 1], P - 2); for(int i = N - 2;i >= 0;i --) inv[i] = 1ll * inv[i + 1] * (i + 1) % P; } int C(int n, int m){ if(n < 0 || m < 0 || m > n) return 0; return 1ll * fac[n] * inv[m] % P * inv[n - m] % P; } } namespace muti{ int rev[M]; int tmpA[M], tmpB[M]; void NTT(int *f, int lim, int type){ for(int i = 0;i < lim;i ++){ if(i < rev[i]){ swap(f[i], f[rev[i]]); } } for(int mid = 1;mid < lim;mid <<= 1){ int wn = base::qpow(3, (P - 1) / (mid << 1)); if(type == -1) wn = base::qpow(wn, P - 2); for(int j = 0;j < lim;j += (mid << 1)){ for(int k = 0, w = 1;k < mid;k ++, w = 1ll * w * wn % P){ int x = f[j + k], y = 1ll * w * f[j + k + mid] % P; f[j + k] = (x + y) % P; f[j + k + mid] = (x - y + P) % P; } } } if(type == -1){ int iv = base::qpow(lim, P - 2); for(int i = 0;i < lim;i ++){ f[i] = 1ll * f[i] * iv % P; } } } vi mul(const vi &A, const vi &B){ int len = sz(A) + sz(B) - 1; if(len <= 40){ vi res(len, 0); for(int i = 0;i < sz(A);i ++){ for(int j = 0;j < sz(B);j ++){ res[i + j] = (res[i + j] + 1ll * A[i] * B[j]) % P; } } return res; } int lim = 1, l = 0; while(lim < len) lim <<= 1, l ++; for(int i = 0;i < lim;i ++){ rev[i] = (rev[i >> 1] >> 1) | ((i & 1) << (l - 1)); } for(int i = 0;i < lim;i ++){ tmpA[i] = (i < sz(A)) ? A[i] : 0; tmpB[i] = (i < sz(B)) ? B[i] : 0; } NTT(tmpA, lim, 1); NTT(tmpB, lim, 1); for(int i = 0;i < lim;i ++){ tmpA[i] = 1ll * tmpA[i] * tmpB[i] % P; } NTT(tmpA, lim, -1); return vi(tmpA, tmpA + len); } } namespace seg{ vi tr[N << 2]; int C[N]; void build(int u, int l, int r){ if(l == r){ tr[u] = {1, a[l]}; return; } int mid = (l + r) >> 1; build(u << 1, l, mid); build(u << 1 | 1, mid + 1, r); tr[u] = muti::mul(tr[u << 1], tr[u << 1 | 1]); } void down(int u, int l, int r, vi U){ if(l == r){ if(sz(U) > 1) C[l] = (U[0] + 1ll * U[1] * a[l] % P) % P; else C[l] = U[0] % P; return; } int mid = (l + r) >> 1; int len1 = mid - l + 2; vi L(U.begin(), U.begin() + min(sz(U), len1)); L.resize(len1, 0); down(u << 1, l, mid, L); int d = sz(tr[u << 1]) - 1; int len2 = r - mid + 1; vi p = tr[u << 1]; reverse(all(p)); vi res = muti::mul(U, p); vi R(len2, 0); for(int i = 0;i < len2;i ++){ if(d + i < sz(res)){ R[i] = res[d + i]; } } down(u << 1 | 1, mid + 1, r, R); } } inline void solve(){ cin >> n >> m >> k; for(int i = 1;i <= n;i ++) cin >> a[i]; for(int i = 1;i <= n;i ++) cin >> b[i]; sort(a + 1, a + n + 1, greater<int>()); sort(b + 1, b + n + 1, greater<int>()); if(k == 1){ for(int i = 0;i <= n;i ++){ E[i] = base::C(n, i); } } else{ seg::build(1, 1, n); vi U(k, 0); U[k - 1] = 1; seg::C[0] = 0; seg::down(1, 1, n, U); vi X(n), Y(n); for(int i = 0;i < n;i ++){ int F = (seg::C[n - i] - seg::C[n - i - 1] + P) % P; X[i] = 1ll * F * base::fac[i] % P; Y[i] = base::inv[i]; } reverse(all(X)); vi res = muti::mul(X, Y); for(int i = k;i <= n;i ++){ int w = i - k + 1; if(n - 1 - w >= 0 && n - 1 - w < sz(res)){ E[i] = 1ll * res[n - 1 - w] * base::inv[w] % P; } else E[i] = 0; } for(int i = 0;i < k;i ++){ if(i < sz(seg::tr[1])) E[i] = seg::tr[1][i]; else E[i] = 0; } } int d = m - k; if(d <= 0){ int sum = 0; for(int i = 1;i <= n;i ++) sum = (sum + b[i]) % P; for(int i = 0;i <= m;i ++){ int y = m - i; if(y >= 1 && y <= n) A[i] = 1ll * sum * base::C(n - 1, y - 1) % P; else A[i] = 0; } } else{ vi F(n), G(n); for(int i = 0;i < n;i ++){ F[i] = 1ll * b[n - i] * base::fac[i] % P; G[i] = base::inv[i]; } reverse(all(F)); vi H = muti::mul(F, G); for(int i = 0;i < n;i ++){ c[i] = 1ll * H[n - 1 - i] * base::inv[i] % P; } for(int i = k;i <= m;i ++){ int p = m - i - 1; if(p >= 0 && p < n) A[i] = c[p]; else A[i] = 0; } F.assign(n, 0); G.assign(n, 0); for(int i = d;i < n;i ++){ int s = base::C(i - 1, d - 1); if((i - d) & 1) s = (P - s) % P; F[i] = 1ll * c[i] * s % P * base::fac[n - 1 - i] % P; } for(int i = 0;i < n;i ++) G[i] = base::inv[i]; vi V = muti::mul(F, G); for(int i = 0;i < k;i ++){ int p = m - i - 1; if(p >= 0 && p < sz(V)){ A[i] = 1ll * V[p] * base::inv[n - m + i] % P; } else A[i] = 0; } } int ans = 0; for(int i = max(0, m - n);i <= min(n, m);i ++){ ans = (ans + 1ll * E[i] * A[i]) % P; } cout << ans << '\n'; } int main(){ #ifdef LOCAL //freopen("test.txt", "r", stdin); #endif cin.tie(0) -> ios::sync_with_stdio(0); base::init(); int T = 1; cin >> T; while(T --) solve(); #ifdef LOCAL cout << "Time: " << 1.0 * clock() / CLOCKS_PER_SEC << " s\n "; #endif return 0; } ``` ::: 成功拿下你谷[最优解](https://www.luogu.com.cn/record/297090248)。