Loading

P10141 [USACO24JAN] Merging Cells P 题解

给定 \(n\) 个细胞,大小为 \(a_i\),每次均匀随机选择两个相邻细胞合并,大小为 \(a_i+a_{i+1}\),新编号等于原来两者中较大的编号,相等取编号较大的,求每个细胞 \(i\) 最后存活的概率模 \(10^9+7\)

首先考虑转化一下题意。因为合并是比较困难的。考虑把它转化成分裂。初始时区间 \([i,i]\) 就表示细胞 \(i\),合并就是把 \([l,k]\)\([k+1,r]\) 合并成 \([l,r]\),所以反向看 \([l,r]\) 就是一个细胞。考虑合并规则,设 \(s(l,r)\) 表示 \(l\)\(r\) 的和,那么如果 \(s(l,k)>s(k+1,r)\) 编号继承左边,否则继承右边。反过来,等式推一下,就有 \(2s(l,k)>s(l,r)\) 进入左区间,否则进入右区间。

根据这个很容易设计一个 dp 状态 \(dp_{l,r}\) 表示达到区间 \([l,r]\) 的概率,显然初始 \(dp_{1,n}=1\),而答案就是 \(dp_{i,i}\)。当前区间中有 \(k\in[l,r-1]\),所以每个点被选中的概率为 \(1\over r-l\),所以转移到子区间的概率为 $$\Delta=\frac{dp_{l,r}}{r-l}$$ 那么什么时候转移到左边,什么时候转移到右边呢?考虑找到最小的 \(k_0\) 使得 \(2s(l,k_0)>s(l,r)\)。由于前缀和单调递增,所以一定有 \(k\geq k_0\) 进入左区间,\(k<k_0\) 进入右区间。如果找不到则 \(k_0=r\),表示都走右区间。那转移方程就是 $$\begin{cases}dp_{l,k}\leftarrow dp_{l,k}+\Delta&k\in[k_0,r-1]\\ dp_{k+1,r}\leftarrow dp_{k+1,r}+\Delta&k\in[l,k_0]\end{cases}$$ 枚举所有区间以及 \(k\),复杂度 \(O(n^3)\),考虑怎么优化。

观察转移方程,固定左端点 \(l\),第一种转移相当于区间加,固定右端点也一样。这样可以使用差分优化,把转移复杂度降到单次 \(O(1)\)。具体而言,我们从大到小枚举区间长度。对于第一种情况,右端点是单调递减,为了方便起见反转端点,这样改为单调递增,在原序列上区间 \(k\in[k_0,r-1]\) 加就转为 \(rev\in{n-r+1,n-k_0+1}\),然后动态维护当前差分数组的前缀和。第二种更新同理。那么区间的真实 dp 值就等于左边的前缀和加上右边的前缀和。那么这样更新是 \(O(1)\),找 \(k_0\) 二分 \(O(\log n)\),复杂度 \(O(n^2\log n)\),已经可以通过。

但还可以继续优化。观察 \(k_0\) 的变化过程,由于固定 \(l\)\(r\) 递减,所以 \(s(l,r)\) 递减,而可能有更小的 \(k_0\) 满足,所以 \(k_0\) 单调不增。那么可以维护指针 \(p_l\) 表示 \(l\) 为左端点时的 \(k\),每次看情况往左边移动。这样均摊 \(O(1)\),总复杂度降至 \(O(n^2)\)

接下来是 \(O(n^2\log n)\) 的代码,\(O(n^2)\) 的代码作为拓展就不放了。

#include<bits/stdc++.h>
#define L(a,b,c,d) for(int a=b;a<=c;a+=d)
#define R(a,b,c,d) for(int a=b;a>=c;a-=d)

using namespace std;
typedef long long i64;
typedef __int128 i128;

const int N=5e3+5,M=1e9+7;

void solve();
int n;
vector<int> a,pre,inv,Lr,Rr,Lp,Rp,Ld[N],Rd[N],ans;

signed main(){
  int Test=1;
// scanf("%d",&Test);
  while(Test--) solve();
  return 0;
}

void solve(){
  scanf("%d",&n);
  a.assign(n+1,0),pre.assign(n+1,0),inv.assign(n+1,0);
  L(i,1,n,1){
    scanf("%d",&a[i]);
    pre[i]=pre[i-1]+a[i];
    Ld[i].assign(n+1,0),Rd[i].assign(n+1,0);
  }
  inv[1]=1;
  L(i,2,n,1) inv[i]=1ll*(M-M/i)*inv[M%i]%M;
  Lr.assign(n+1,0),Rr.assign(n+1,0),Lp.assign(n+1,0),Rp.assign(n+1,0),ans.assign(n+1,0);
  R(len,n,1,1){
    L(i,1,n-len+1,1){
      int j=i+len-1;
      int rev=n-j+1;
      while(Lp[i]<rev){
        Lp[i]++;
        int x=Ld[i][Lp[i]];
        if(x) Lr[i]=(Lr[i]+x)%M;
      }
      while(Rp[j]<i){
        Rp[j]++;
        int x=Rd[j][Rp[j]];
        if(x) Rr[j]=(Rr[j]+x)%M;
      }
      int x;
      if(len==n) x=1;
      else x=(Lr[i]+Rr[j])%M;
      if(len==1){
        ans[i]=x;
        continue;
      }
      if(!x) continue;
      int sum=pre[j]-pre[i-1],l=i-1,r=j,k=j;
      while(l+1<r){
        int mid=(l+r)/2;
        int lsum=pre[mid]-pre[i-1];
        if(2ll*lsum>sum){
          k=mid;
          r=mid;
        }
        else l=mid;
      }
      int delta=1ll*x*inv[len-1]%M;
      l=k,r=j-1;
      if(l<=r){
        int lrev=n-l+1,rrev=n-r+1;
        Ld[i][rrev]=(Ld[i][rrev]+delta)%M;
        Ld[i][lrev+1]=(Ld[i][lrev+1]-delta+M)%M;
      }
      l=i+1,r=k;
      if(l<=r){
        Rd[j][l]=(Rd[j][l]+delta)%M;
        Rd[j][r+1]=(Rd[j][r+1]-delta+M)%M;
      }
    }
  }
  L(i,1,n,1) printf("%d\n",ans[i]);
}
posted @ 2026-08-18 10:25  jess1ca1o0g3  阅读(3)  评论(0)    收藏  举报