题解:P5299 [PKUWC2018] Slay the Spire
终焉折枝
·
2026-09-08 13:09:13
·
题解
[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 的 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)。