P4007 小 Y 和恐怖的奴隶主 题解
题目描述
初始 Boss 有一个生命值为 \(m\) 的随从,你可以进行 \(n\) 次攻击,每次从 Boss 及其随从中等概率选择一个,并扣减一点生命值。
如果当前随从数量小于 \(k\) ,攻击某个随从且没有造成死亡,那么会重新召唤一个生命值为 \(m\) 的随从。
\(T\) 次询问,给定 \(n\) ,求对 Boss 造成的伤害的期望值。
数据范围
- \(1\le T\le 10^3,1\le n\le 10^{18},1\le m\le 3,1\le k\le 8\) 。
时间限制 \(\texttt{2s}\) ,空间限制 \(\texttt{512MB}\) 。
分析
先考虑一个期望 \(\texttt{dp}\) 。
\(f_{i,a,b,c}\) 表示当前进行 \(i\) 次攻击,血量为 \(1,2,3\) 的随从分别有 \(a,b,c\) 个的概率。
为统计答案,新开一个变量 res 记录对 Boss 造成伤害的期望即可。
注意到 \(a+b+c\le8\) ,因此合法的三元组 \((a,b,c)\) 数量不超过 \(\binom{11}3=165\) 。
由于相邻层之间的转移固定,考虑矩阵快速幂。
我们需要记录所有三元组 \((a,b,c)\) 和 res ,总状态数 \(166\) 。
使用向量乘矩阵的 trick ,时间复杂度 \(\mathcal O(166^3\log n+T166^2\log n)\) 。
卡常小技巧:矩乘乘法时,用
__int128存储中间变量,可以大幅减少取模次数。
#include<bits/stdc++.h>
#define ll long long
using namespace std;
const int maxn=170,mod=998244353;
int k,m,t,u;
int id[9][9][9];
ll n,f[maxn][maxn];
struct vec
{
int v[maxn];
}cur,res;
struct mat
{
int v[maxn][maxn];
}pw[60];
inline vec operator*(const vec &a,const mat &b)
{
static vec c;
for(int i=1;i<=u;i++)
{
__int128 res=0;
for(int j=1;j<=u;j++) res+=1ll*a.v[j]*b.v[j][i];
c.v[i]=res%mod;
}
return c;
}
inline mat operator*(const mat &a,const mat &b)
{
static mat c;
for(int i=1;i<=u;i++)
for(int j=1;j<=u;j++)
{
__int128 res=0;
for(int k=1;k<=u;k++) res+=1ll*a.v[i][k]*b.v[k][j];
c.v[i][j]=res%mod;
}
return c;
}
inline int qpow(int a,int k)
{
int res=1;
while(k)
{
if(k&1) res=1ll*res*a%mod;
a=1ll*a*a%mod,k>>=1;
}
return res;
}
int main()
{
scanf("%d%d%d",&t,&m,&k);
for(int a=0;a<=k;a++)
for(int b=0;b<=(m>=2?k-a:0);b++)
for(int c=0;c<=(m>=3?k-a-b:0);c++)
id[a][b][c]=++u;
++u,f[u][u]=1;
for(int a=0;a<=k;a++)
for(int b=0;b<=k;b++)
for(int c=0;c<=k;c++)
{
int x=id[a][b][c],flg=a+b+c<k,inv=qpow(a+b+c+1,mod-2);
if(!x) continue;
f[x][x]+=inv,f[x][u]+=inv;
if(a) f[x][id[a-1][b][c]]+=1ll*a*inv%mod;
if(b) f[x][id[a+1][b-1+(flg&&m==2)][c+(flg&&m==3)]]+=1ll*b*inv%mod;
if(c) f[x][id[a][b+1][c-1+flg]]+=1ll*c*inv%mod;
}
for(int i=1;i<=u;i++) for(int j=1;j<=u;j++) pw[0].v[i][j]=f[i][j]%mod;
for(int i=1;i<60;i++) pw[i]=pw[i-1]*pw[i-1];
cur.v[id[m==1][m==2][m==3]]=1;
while(t--)
{
scanf("%lld",&n),res=cur;
for(int i=0;i<60;i++) if(n>>i&1) res=res*pw[i];
printf("%d\n",res.v[u])%mod;
}
return 0;
}
本文来自博客园,作者:peiwenjun,转载请注明原文链接:https://www.cnblogs.com/peiwenjun/p/17255757.html
浙公网安备 33010602011771号