题解:B4337 [中山市赛 2023] 简单数学题

· · 题解

Analysis

递推推式子好题。

我们设第 i 次操作后,第一个盒子里的白球期望概率为 X_i,那么 X_0 也就是 a1。首先设 A=a1+a2,B=b1+b2,S=a1+b1(原题面中就没加下标)。

先看从盒一转移到盒二点过程。我们有 \frac{X_{i-1}}{A} 的概率抽中白球,即盒一白球数期望减少 \frac{X_{i-1}}{A}。而从盒二到盒一,可以分两类讨论。

综合来看,由加法原理可得盒一白球数期望增加 \frac{X_{i-1}}{A}\cdot\frac{S-X_{i-1}+1}{B+1}+\frac{A-X_{i-1}}{A}\cdot\frac{S-X_{i-1}}{B+1}。再将上述过程结合起来,可以得到 X_i 的表达式。

X_i&=X_{i-1}-\frac{X_{i-1}}{A}+\frac{X_{i-1}}{A}\cdot\frac{S-X_{i-1}+1}{B+1}+\frac{A-X_{i-1}}{A}\cdot\frac{S-X_{i-1}}{B+1}\\ &=X_{i-1}-\frac{X_{i-1}}{A}+\frac{X_{i-1}S-{X_{i-1}}^2+X_{i-1}+AS-X_{i-1}S-X_{i-1}A+{X_{i-1}}^2}{A\cdot(B+1)}\\ &=X_{i-1}-\frac{X_{i-1}}{A}+\frac{X_{i-1}+AS-X_{i-1}A}{A\cdot(B+1)}\\ &=X_{i-1}+\frac{X_{i-1}+AS-X_{i-1}A-X_{i-1}B-X_{i-1}}{A\cdot(B+1)}\\ &=X_{i-1}+\frac{AS-X_{i-1}\cdot(A+B)}{A\cdot(B+1)}\\ &=\frac{A\cdot(B+1)-(A+B)}{A\cdot(B+1)}X_{i-1}+\frac{AS}{A\cdot(B+1)}\\ &=\frac{AB-B}{A\cdot(B+1)}X_{i-1}+\frac{S}{B+1} \end{aligned}

我们令 p=\frac{AB-B}{A\cdot(B+1)},q=\frac{S}{B+1},可以得到一个非常珂爱的线性递推式 X_i=p\cdot X_{i-1}+q

由于操作次数 n 可达 10^{18},我们可以利用矩阵快速幂来 \mathcal O(\log n) 求解。构造状态转移矩阵可得:

\begin{pmatrix} X_n \\ 1 \end{pmatrix} = \begin{pmatrix} p & q \\ 0 & 1 \end{pmatrix} \begin{pmatrix} X_{n-1} \\ 1 \end{pmatrix}

最终答案即为 \frac{X_n}{A}\bmod998244353

Code

#include"bits/stdc++.h"
#define int long long
using namespace std;
const int mod = 998244353;
int n, a1, a2, b1, b2;
struct matrix {
    int mat[3][3];
    int r, c;
    matrix () {
        memset(mat, 0, sizeof mat);
    }
    int *operator[] (int i) {
        return mat[i];
    }
    matrix operator* (matrix &b) const {
        matrix re;
        re.r = r, re.c = b.c;
        for (int i = 1; i <= r; i++)
            for (int k = 1; k <= c; k++)
                for (int j = 1; j <= b.c; j++)
                    re[i][j] = (re[i][j] + mat[i][k] * b[k][j]) % mod;
        return re;
    }
};
matrix qpow(matrix a, int b) {
    matrix re;
    re[1][1] = re[2][2] = 1;
    re.r = re.c = 2;
    while (b) {
        if (b & 1)
            re = re * a;
        a = a * a;
        b >>= 1;
    }
    return re;
}
int fpow(int a, int b) {
    a %= mod;
    int re = 1;
    while (b) {
        if (b & 1)
            re = re * a % mod;
        a = a * a % mod;
        b >>= 1;
    }
    return re;
}
signed main() {
    cin.tie(0);
    cout.tie(0);
    ios::sync_with_stdio(0);
    cin >> n >> a1 >> a2 >> b1 >> b2;
    int a = (a1 + a2) % mod, b = (b1 + b2) % mod, s = (a1 + b1) % mod;
    matrix t;
    t[1][1] = (a - 1) % mod * b % mod * fpow(a * (b + 1) % mod, mod - 2) % mod;
    t[1][2] = s * fpow(b + 1, mod - 2) % mod;
    t[2][2] = 1;
    t.r = t.c = 2;
    t = qpow(t, n);
    int ans = (t[1][1] * (a1 % mod) % mod + t[1][2]) % mod;
    ans = ans * fpow(a, mod - 2) % mod;
    cout << ans;
    return 0;
}