题解:CF2034F2 Khayyam's Royal Decree (Hard Version)

· · 题解

我怎么只会 k^2

题意

有一个计数器,初始为 0。人要从 (n,m) 走到 (0,0),每次等概率地向下或向左走,向左走计数器 +2,向下走计数器 +1。其中有 k 关键点 (r_i,b_i),若走到了这些点计数器 \times 2。求最后计数器期望为多少。

思路

首先我们要知道:从 (sx,sy) 走到 (ex,ey) 方案数为 C_{sx-ex+sy-ey}^{sx-ex}。之后有用。

由于每条路径概率相等,所以期望等于总和除以总方案数,而我们知道总方案数为 C_{n+m}^n。那么我们来求总和。

实际上我们不能直接算整张图每个点的总和,我们会发现最终影响答案或者被影响需要我们单独考虑计算的只有这 k 个关键点和 (0,0)。那么我们接下来只考虑这几个点。

假设不考虑这些点 (r_i,b_i) 的限制,那我们 (0,0) 的总和就是总方案数 \times 计数器的值,也就是 C_{n+m}^n\times (2n+m)

那么现在看看假设我们多了一个 (r_i,b_i) 答案会发生什么变化。

首先我们计算 (n,m) 走到 (r_i,b_i) 的总和记为 dp_idp_i 相对于原来没有这个关键点的时候由于 \times 2 增加了 (2 - 1)\times C_{n-r_i+m-b_i}^{n-r_i}\times (2(n-r_i)+(m-b_i))=C_{n-r_i+m-b_i}^{n-r_i}\times (2(n-r_i)+(m-b_i)),记为 \Delta dp_i

然后我们接着计算对 (0,0) 的影响。所有被影响的路径即经过 (r_i,b_i) 的路径可以与 (r_i,b_i)(0,0) 的路径一一对应,方案数为 C_{r_i+b_i}^{r_i},而增加的值就是 \Delta dp_i 了。

假设再增加一个 (r_j,b_j)(可以走到 (r_i,b_i)),那么我们可以先计算出 \Delta dp_j,再把 (r_i,b_i) 当做刚刚的 (0,0) 那样来统计 \Delta dp_i

发现我们只需要维护 \Delta dp_i 了,那么我们把原来的 dp_i 的含义踢掉变成 \Delta dp_i

借此我们得到了这题的完整做法:

  1. (r_i,b_i) 从大到小排序,保证如果按顺序从前往后计算的话能够更新到 dp_idp_j 都被计算过。

  2. 发现 (0,0) 的答案计算方式是和 dp_i 一模一样的,于是我们为了方便可以把 (0,0) 也当成一个关键点 (r_{k+1},b_{k+1}) 统计答案。最后答案即为 \dfrac{dp_{k+1}}{C_{n+m}^n}

时间复杂度 O(n+m+k^2)。前半部分在于预处理组合数,后半部分在于 dp。

代码:

#include <bits/stdc++.h>
using namespace std;
typedef long long ll;
const ll mod = 998244353;
ll n, m, T, jc[500002], inv[500002], dp[502], k;
struct node {
    ll r, b;
}p[502];
ll c(ll x, ll y) { // x 选 y 组合数
    return jc[x] * inv[y] % mod * inv[x - y] % mod;
}
ll qmi(ll x, ll y) {
    ll res = 1;
    while (y) res = res * (y & 1 ? x : 1) % mod, x = x * x % mod, y >>= 1;
    return res;
}
bool cmp(node az, node bz) {
    return az.r > bz.r || (az.r == bz.r && az.b > bz.b);
}
int main() {
    cin >> T;
    jc[0] = 1;
    for (ll i = 1; i <= 500000; i ++ ) jc[i] = jc[i - 1] * i % mod;
    inv[500000] = qmi(jc[500000], mod - 2);
    for (ll i = 499999; i >= 0; i -- ) inv[i] = inv[i + 1] * (i + 1) % mod;
    //预处理阶乘及其逆元
    while (T -- ) {
        cin >> n >> m >> k;
        for (ll i = 1; i <= k; i ++ ) dp[i] = 0, cin >> p[i].r >> p[i].b;
        sort(p + 1, p + k + 1, cmp);
        p[++ k] = {0, 0};
        for (ll i = 1; i <= k; i ++ ) {
            dp[i] = c(n + m - p[i].r - p[i].b, n - p[i].r) * (2 * (n - p[i].r) + (m - p[i].b)) % mod;
            for (ll j = 1; j < i; j ++ ) if (p[j].r >= p[i].r && p[j].b >= p[i].b) 
                dp[i] = (dp[i] + c(p[j].r + p[j].b - p[i].r - p[i].b, p[j].r - p[i].r) * dp[j] % mod) % mod;
        }
        cout << dp[k] * qmi(c(n + m, n), mod - 2) % mod << "\n";
    }
}