peiwenjun's blog 没有知识的荒原

CF1184D2 Parallel Universes (Hard) 题解

题目描述

\(n\) 个球排成一排,第 \(k\) 个球为黑色,其余均为白色。

对序列进行如下操作直到黑球在开头或末尾:

设当前序列有 \(l\) 个球。

  • \((1-\frac lm)\cdot\frac1{l+1}\) 的概率在第 \(i\)\(0\le i\le l\) )个球后面插入一个白球。
  • \(\frac lm\cdot\frac1{l-1}\) 的概率将整个序列分为 \([1,i],[i+1,l]\)\(1\le i\lt l\) )两个部分,仅保留黑球所在部分。

求操作结束后序列长度的期望,对 \(10^9+7\) 取模。

数据范围

  • \(1\le k\le n\le m\le 250\)

时间限制 \(\texttt{4s}\) ,空间限制 \(\texttt{250MB}\)

分析

\(f_{i,j}\) 表示初始序列长度为 \(i\) ,黑球为第 \(j\) 个的答案。

初始值 \(f_{i,1}=f_{i,i}=i\)

转移方程如下:

\[f_{i,j}=\frac{m-i}m\frac{i-j+1}{i+1}f_{i+1,j} +\frac{m-i}m\frac j{i+1}f_{i+1,j+1} +\sum_{k=1}^{j-1}\frac im\frac1{i-1}f_{i-k,j-k} +\sum_{k=j}^{i-1}\frac im\frac1{i-1}f_{k,j}\\ \]

\(f_{n,k}\) 即为所求。

显然转移有环,直接高斯消元可以获得一个 \(\mathcal O(m^6)\) 的做法。

将上式移项:

\[\frac{m-i}m\frac j{i+1}f_{i+1,j+1}=f_{i,j} -\frac{m-i}m\frac{i-j+1}{i+1}f_{i+1,j} -\sum_{k=1}^{j-1}\frac im\frac1{i-1}f_{i-k,j-k} -\sum_{k=j}^{i-1}\frac im\frac1{i-1}f_{k,j}\\ \]

从小到大枚举 \(j\) ,维护 \(f_{i,j}\) 关于 \(f_{[3,m],2}\) 的线性组合。

\(j\le 2\) 是容易的,其余可以根据上式递推,需要用前缀和优化。

边界 \(\forall 3\le j\le m,f_{m+1,j}=0\) ,恰好 \(m-2\) 个方程。

高斯消元解出 \(f_{*,2}\) 的值,再带入 \(f_{n,k}\) 的线性组合即可,时间复杂度 \(\mathcal O(m^3)\)

#include<bits/stdc++.h>
#define poly array<int,maxn>
using namespace std;
const int maxn=255,mod=1e9+7;
int k,m,n;
int a[maxn][maxn],v[maxn][maxn];
poly sum[maxn],g[maxn][maxn];///f[i][j]=\sum_{x=3}^m g_{i,j,x}*f_{x,2}+g_{i,j,m+1}
inline int qpow(int a,int k)
{
    int res=1;
    for(;k;a=1ll*a*a%mod,k>>=1) if(k&1) res=1ll*res*a%mod;
    return res;
}
inline void add(int &x,int y)
{
    if((x+=y)>=mod) x-=mod;
}
inline void dec(int &x,int y)
{
    if((x-=y)<0) x+=mod;
}
inline poly operator+(poly a,poly b)
{
    for(int i=3;i<=m+1;i++) add(a[i],b[i]);
    return a;
}
inline poly operator-(poly a,poly b)
{
    for(int i=3;i<=m+1;i++) dec(a[i],b[i]);
    return a;
}
inline poly operator*(poly a,int b)
{
    for(int i=3;i<=m+1;i++) a[i]=1ll*a[i]*b%mod;
    return a;
}
inline void operator+=(poly &a,poly b)
{
    a=a+b;
}
void gauss()
{
    for(int i=3;i<=m;i++)
    {
        int p=0;
        for(int j=m;j>=i;j--) if(a[j][i]) p=j;
        assert(p);
        if(p!=i) for(int k=i;k<=m+1;k++) swap(a[i][k],a[p][k]);
        for(int j=3,inv=qpow(a[i][i],mod-2);j<=m;j++)
        {
            if(j==i) continue;
            int t=1ll*a[j][i]*inv%mod;
            for(int k=i;k<=m+1;k++) a[j][k]=(a[j][k]-1ll*a[i][k]*t)%mod;
        }
    }
}
int main()
{
    scanf("%d%d%d",&n,&k,&m);
    if(k==1||k==n) printf("%d\n",n),exit(0);
    for(int i=1;i<=m;i++) for(int j=1,inv=qpow(i,mod-2);j<=i;j++) v[i][j]=1ll*j*inv%mod;
    for(int i=1;i<=m;i++) g[i][1][m+1]=i;
    for(int i=2;i<=m;i++) i==2?g[2][2][m+1]=2:g[i][2][i]=1;
    for(int j=2;j<m;j++)
    {///计算g[j+1~m+1][j+1]
        g[j+1][j+1][m+1]=j+1;
        poly cur{};
        for(int i=j+1;i<=m;i++)
        {
            cur+=g[i-1][j],sum[i-j]+=g[i-1][j-1];///sum[i-j]=\sum_{k=1}^{j-1}g[i-k][j-k]
            poly now=g[i][j]-g[i+1][j]*v[m][m-i]*v[i+1][i-j+1]-(cur+sum[i-j])*v[m][i]*v[i-1][1];
            g[i+1][j+1]=i!=m?now*qpow(1ll*v[m][m-i]*v[i+1][j]%mod,mod-2):now;
        }
    }
    for(int j=3;j<=m;j++)
    {
        for(int i=3;i<=m;i++) a[j][i]=g[m+1][j][i];
        a[j][m+1]=mod-g[m+1][j][m+1];
    }
    gauss();
    int res=g[n][k][m+1];
    for(int i=3;i<=m;i++) res=(res+1ll*a[i][m+1]*qpow(a[i][i],mod-2)%mod*g[n][k][i])%mod;
    printf("%d\n",(res+mod)%mod);
    return 0;
}

posted on 2023-06-06 18:49  peiwenjun  阅读(10)  评论(0)    收藏  举报

导航