P4007 [清华集训 2017] 小 Y 和恐怖的奴隶主 题解

笔记列表

  • 发现 \(m\leq 3, k \leq 8\)
  • 首先你要算期望,正推是很麻烦的,所以考虑倒推
  • \(f_{i,a,b,c}\) 表示攻击 \(i\) 次后(a个1血仆从,b个2血2b ,c个3血),之后的进攻的期望攻击次数
  • 注意,这里是的 不是 \(i\) 次进攻 从 第一次到第 \(i\) 次的期望数,而是第\(i\)次到最后一次第 \(n\) 次的期望数 (倒推)
  • 这里初始状态和答案最后考虑,因为还要优化
  • 现在这个简单\(dp\) 的转移如下
    • 设当前单位数 \(s = (a+b+c+1)\) (还有boss自己)
    • \(\frac{1}{s} 的概率攻击boss\), \(f_{i,a,b,c}\gets f_{i,a,b,c}+ \frac{f_{i+1,a,b,c} + 1}{s}\)
    • \(\frac{a}{s}\)的概率攻击1血仆从, \(f_{i,a,b,c}\gets f_{i,a,b,c}+f_{i+1,a-1,b,c} \times \frac{a}{s}\)
    • 攻击到2血和3血仆从类似,但是在攻击完后如果仆从数可以增加,那要从 多一个 \(m\)血仆从的状态转移
  • 发现现在 \(i\) 的转移和 \(i-1\) 相关,而且这个转移发现可以用 矩阵快速幂优化
  • 然后状态数可以类似离散化一下,发现有效的状态数很少
  • \(id_{a,b,c} 表示a个1\)血仆从, \(b\) 个2血仆从, \(c\)个3血仆从是第几个状态
  • \(node_i\) 表示 \(i\) 这个状态对应了几个1血,2血,3血仆从
  • \(idx\) 表示状态总数
  • 好的现在 \(f_{i,j}\) 表示 \(i\) 次攻击后的状态是第\(j\) 个状态,之后的期望攻击次数
  • 定义一个向量

\[V_i = [f_{i,1},f_{i,2},...,f_{i,idx}, 1] \]

递推关系式可以写成 \(V_i = V_{i-1} \times M\)

  • \(M\) 就是要重点要构造的转移矩阵
  • 对应我们的转移
    • 攻击了 \(boss\)
      \(M[s][s]\)加上了攻击 \(boss\) 的概率 \(\frac{1}{a+b+c+1}\)
      然后多攻击了一次 \(M[idx+1][s]\)加上 \(\frac{1}{a+b+c+1}\)
      想不清除可以回忆矩阵乘法是一行乘一列,这里相当与是把上面推的式子拆了以下,变成了\(\frac{f_{i+1,a,b,c}}{s}+\frac{1}{s}\)
    • 攻击了某一个仆从
      可以通过上面的方法算出攻击这个仆从的概率 \(p\),和攻击这个仆从需要从 \(s'\) 转移, \(M[s'][s]加上点p\)
  • 最后别忘了常数项1也要转移:
    \(M[idx+1][idx+1]=1\)

然后还是太慢了,因为是查询多次所以你预处理\(M\)\(2^i\)次幂
然后查的时候直接乘起来

哦哦哦等以下忘记说初始化和答案了
初始时从 \(V_0\)开始乘,因为还剩下\(0\)次攻击所以未来的期望数是0
所以 \(V_0初始矩阵[0,0,..,0,1]\) 最后的1是常数项
答案是\(V_n\)中, 1个\(m\)血仆从所对应的那一个状态的值
真是一个大好题

#include<bits/stdc++.h>
#define int long long
#define fore(i,a,b) for(int i=(a);i<=(b);++i)
#define repe(i,a,b) for(int i=(a);i>=(b);--i)
using namespace std;
const int N = 3010;
const int mod = 998244353;
int T, m, k;
int id[10][10][10], idx;
struct Node {
    int A, B, C;
} node[N];
struct Matrix {
    int a[210][210];
    Matrix() {
        memset(a, 0, sizeof(a));
    }
    Matrix operator *(const Matrix& o) const {
        Matrix res;
        fore(t, 1, idx + 1) {
            fore(i, 1, idx + 1) {
                if (a[i][t] == 0) continue;
                fore(j, 1, idx + 1) {
                    res.a[i][j] = (res.a[i][j] + a[i][t] * o.a[t][j] % mod) % mod;
                }
            }
        }
        return res;
    }
} base, pre[64];
int qpow(int x, int y) {
    int res = 1;
    while (y) {
        if (y & 1) res = res * x % mod;
        x = x * x % mod;
        y >>= 1;
    }
    return res;
}
signed main() {
    ios::sync_with_stdio(false);
    cin.tie(0);
    cout.tie(0);
    cin >> T >> m >> k;
    fore(a, 0, k) {
        fore(b, 0, (m >= 2 ? k - a : 0)) {
            fore(c, 0, (m >= 3 ? k - a - b : 0)) {
                id[a][b][c] = ++idx;
                node[idx] = {a, b, c};
            }
        }
    }
    fore(i, 1, idx) {
        int a = node[i].A, b = node[i].B, c = node[i].C;
        int s = qpow(a + b + c + 1, mod - 2);
        base.a[i][i] = s;
        base.a[idx + 1][i] = s;
        if (a != 0) {
            base.a[id[a - 1][b][c]][i] = s * a % mod;
        }
        if (b != 0) {
            int x = a + 1, y = b - 1, z = c;
            if (x + y + z < k) {
                if (m == 1) x++;
                else if (m == 2) y++;
                else z++;
            }
            base.a[id[x][y][z]][i] = s * b % mod;
        }
        if (c != 0) {
            int x = a, y = b + 1, z = c - 1;
            if (x + y + z < k) {
                if (m == 1) x++;
                else if (m == 2) y++;
                else z++;
            }
            base.a[id[x][y][z]][i] = s * c % mod;
        }
    }
    base.a[idx + 1][idx + 1] = 1;
    pre[0] = base;
    for (int i = 1; i < 64; ++i) {
        pre[i] = pre[i - 1] * pre[i - 1];
    }
    while (T--) {
        int n;
        cin >> n;
        int ans[210] = {0};
        ans[idx + 1] = 1;
        for (int i = 0; i < 64; ++i) {
            if (n & (1LL << i)) {
                int new_ans[210] = {0};
                for (int k = 1; k <= idx + 1; ++k) {
                    if (ans[k] == 0) continue;
                    for (int j = 1; j <= idx + 1; ++j) {
                        new_ans[j] = (new_ans[j] + ans[k] * pre[i].a[k][j]) % mod;
                    }
                }
                for (int j = 1; j <= idx + 1; ++j) ans[j] = new_ans[j];
            }
        }
        cout << ans[id[m == 1][m == 2][m == 3]] << '\n';
    }
    return 0;
}
posted @ 2026-09-18 07:56  wmq2012  阅读(7)  评论(1)    收藏  举报