2 条题解

  • 0
    @ 2026-7-4 22:33:24

    #include <cstdio>
    #include <cstring>
    const int M = 2005; 
    const int MOD = 998244353;
    #define int long long
    int read()
    {
    	int x=0,f=1;char c;
    	while((c=getchar())<'0' || c>'9') {if(c=='-') f=-1;}
    	while(c>='0' && c<='9') {x=(x<<3)+(x<<1)+(c^48);c=getchar();}
    	return x*f;
    }
    int n,k,x,y,p,q,dp[M][M],f[M],b[M],s[M],t[M],a[M],r[M];
    void add(int &x,int y) {x=(x+y)%MOD;}
    int qkpow(int a,int b)
    {
    	int r=1;
    	while(b>0)
    	{
    		if(b&1) r=r*a%MOD;
    		a=a*a%MOD;
    		b>>=1;
    	}
    	return r;
    }
    int work(int k)
    {
    	memset(a,0,sizeof a);
    	memset(b,0,sizeof b);
    	memset(f,0,sizeof f);
    	memset(s,0,sizeof s);
    	memset(t,0,sizeof t);
    	memset(r,0,sizeof r); 
    	memset(dp,0,sizeof dp);
    	for(int i=0;i<=k+1;i++) dp[0][i]=1;
    	for(int j=k;j>=1;j--)
    		for(int i=1;i*j<=k;i++)
    		{
    			int r=0;
    			for(int k=1;k<=i;k++)
    				add(r,dp[k-1][j+1]*dp[i-k][j]);
    			dp[i][j]=r*qkpow(p,j)%MOD*q%MOD;
    			add(dp[i][j],dp[i][j+1]); 
    		}
    	k++;
    	for(int i=0;i<k;i++) f[i+1]=dp[i][1]*q%MOD;
    	for(int i=1;i<=k;i++) b[i]=MOD-f[i];
    	b[0]=s[0]=a[1]=r[0]=1;
    	for(int i=1;i<k;i++) for(int j=0;j<i;j++)
    		add(s[i],s[j]*f[i-j]%MOD);
    	int m=n+1,ans=0;
    	while(m>0)
    	{
    		if(m&1)
    		{
    			for(int i=0;i<=k;i++) t[i]=r[i],r[i]=0;
    			for(int i=0;i<=k;i++)
    				for(int j=0;j<=k;j++)
    					add(r[i+j],t[i]*a[j]%MOD);
    			for(int i=k<<1;i>=k;i--)
    				for(int j=k;j>=0;j--)
    					add(r[i-j],MOD-r[i]*b[j]%MOD);
    		}
    		for(int i=0;i<=k;i++) t[i]=a[i],a[i]=0;
    		for(int i=0;i<=k;i++)
    			for(int j=0;j<=k;j++)
    				add(a[i+j],t[i]*t[j]%MOD);
    		for(int i=k<<1;i>=k;i--)
    			for(int j=k;j>=0;j--)
    				add(a[i-j],MOD-a[i]*b[j]%MOD);
    		m>>=1;
    	}
    	for(int i=0;i<k;i++) add(ans,s[i]*r[i]%MOD);
    	return ans*qkpow(q,MOD-2)%MOD;
    }
    signed main()
    {
    	n=read();k=read();x=read();y=read();
    	p=x*qkpow(y,MOD-2)%MOD;q=(1-p+MOD)%MOD;
    	printf("%lld\n",(work(k)-work(k-1)+MOD)%MOD);
    }
    
    
    • 0
      @ 2026-5-14 18:00:47

      尽量讲得详细点,很适合作为入门题。有错欢迎来直接 diss。

      差分,求面积不超过 kk 的概率,则用 fif_i 表示考虑了前 ii 列,其最大子矩形大小不超过 kk 的概率,进行 DP。

      我们只关心每一列最靠下的障碍在哪里,由于算的是概率,我们不用关心障碍上面的格子的状态。

      假设第 ii 列最下面一个障碍的纵坐标为 zz,则称这一列的高度 hi=z1h_i=z-1,这是一个直方图最大子矩形的问题,具体地,设 lenilen_i 表示高度不低于 ii 的连续段的最长长度,则 max(i×leni)\max(i\times len_i) 即为答案。这是显然的。

      可是仍不好做,我们设 dpi,jdp_{i,j} 表示一个长度为 ii 的序列,每个位置的高度都不低于 jj 且恰有一个地方高度为 jj 的概率,这个限制就很严格了。

      方便起见,我们钦定 dpi,jdp_{i,j} 还要满足其最大子矩形大小不大于 kk,有个必要条件是 ijkij\le k。这启发我们状态数是 O(klnk)O(k\ln k) 的。

      考虑 dpi,jdp_{i,j} 怎么转移,套路地枚举最靠左的高度为 jj 的位置 xx,分割为子问题:

      $$dp_{i,j}=(1-p)p^j\sum_{x=1}^i(\sum_{y>j}dp_{x-1,y})(\sum_{y\ge j}dp_{i-x,y})$$

      很好理解,我们钦定了 xx 是最靠左的高度为 jj 的位置,那么 <x<x 的位置必然高度都 >j>j,右边就没这么多限制。

      注意到可以使用后缀和优化,预处理 dpdp 数组容易做到 O(k2lnk)O(k^2\ln k)。不出意外需要一定程度地精细实现。

      初值为 dp1,j=(1p)pj[jk]dp_{1,j}=(1-p)p^j[j\le k],方便起见可以预处理 pp 的幂次。请注意 dp0,dp_{0,*} 的处理,可能会用到。

      思考怎么用 dpdp 去转移 ff,为了避免算重,同样对 ff 进行钦定:fif_i 为考虑了前 ii 列,其最大子矩形大小不超过 kkhi=0h_i=0 的概率,最后的答案就是 fn+11p\frac{f_{n+1}}{1-p},注意本题一定有 p<1p<1 所以不需要特判。

      常规地,枚举上一个连续段的长度:

      $$f_n=(1-p)(f_{n-1}+\sum_{x=1}^{k-1}f_{n-x-1}\sum_{y=1}^{\lfloor\frac{k}{x}\rfloor}dp_{x,y})$$

      gig_i 表示 $\sum\limits_{y=1}^{\lfloor\frac{k}{x}\rfloor}dp_{x,y}$,特别地,g0=1g_0=1

      fn=(1p)x=0kfnx1gxf_n=(1-p)\sum_{x=0}^{k}f_{n-x-1}g_x

      gg 右移一位并乘上 1p1-p 就是:fn=x=1k+1fnxgxf_n=\sum_{x=1}^{k+1}f_{n-x}g_x

      这还看不出来是常系数齐次线性递推的可以不做这个题了。

      注意到 F(x)G(x)+1=F(x)F(x)G(x)+1=F(x),则有 F(x)=11G(x)F(x)=\frac{1}{1-G(x)}

      使用 Bostan Mori 做到 O(k2logn)O(k^2\log n) 或者也可以使用多项式卷积做到 O(klogklogn)O(k\log k\log n)

      非常烦的是因为我们差分了,所以要做两次。

      #include<bits/stdc++.h>
      
      using namespace std;
      
      const int mod=998244353,N=1010;
      struct modint {
          int val;
          static int norm(const int& x) { return x < 0 ? x + mod : x; }
          modint inv() const {
              int a = val, b = mod, u = 1, v = 0, t;
              while (b > 0) t = a / b, swap(a -= t * b, b), swap(u -= t * v, v);
              return modint(u);
          }
          modint() : val(0) {}
          modint(const int& m) : val(norm(m)) {}
          modint(const long long& m) : val(norm(m % mod)) {}
          modint operator-() const { return modint(norm(-val)); }
          bool operator==(const modint& o) { return val == o.val; }
          bool operator<(const modint& o) { return val < o.val; }
          modint& operator+=(const modint& o) { return val = (1ll * val + o.val) % mod, *this; }
          modint& operator-=(const modint& o) { return val = norm(1ll * val - o.val), *this; }
          modint& operator*=(const modint& o) { return val = static_cast<int>(1ll * val * o.val % mod), *this; }
          modint operator-(const modint& o) const { return modint(*this) -= o; }
          modint operator+(const modint& o) const { return modint(*this) += o; }
          modint operator*(const modint& o) const { return modint(*this) *= o; }
          friend std::ostream& operator<<(std::ostream& os, const modint& a) { return os << a.val; }
      }P,p[N],dp[N][N],sum[N][N],f[N];
      int n,k,x,y;using Poly=vector<modint>;
      Poly g;
      
      Poly operator *(Poly &v,Poly &x){
      	Poly t;t.resize(v.size()+x.size()-1);int A=v.size(),B=x.size();
      	for(int i=0;i<A;++i) for(int j=0;j<B;++j) t[i+j]+=v[i]*x[j];return t;
      }
      Poly& operator *=(Poly &v,Poly &x){return v=v*x;}
      inline Poly nega(Poly v){for(int i=1;i<(int)v.size();i+=2) v[i]=-v[i];return v;}
      inline modint Bostan(Poly g,int k){
      	Poly f;f.resize(g.size());f[0]=1;
      	while(k){
      		Poly tmp=nega(g);
      		f*=tmp;g*=tmp;int up=g.size()/2,sz=0;
      		for(int i=0;i<=up;++i) g[i]=g[i<<1];
      		g.resize(up+1);up=f.size();
      		for(int i=0;i<up;++i) if((i&1)==(k&1)) f[sz++]=f[i];
      		f.resize(sz);k>>=1;
      	}
      	return f[0]*g[0].inv();
      }
      
      inline modint work(int k){
      	g.clear();g.resize(k+2);sum[0][k+1]=1;
      	for(int i=1;i<=k;++i) dp[1][i]=p[i]*(-P+1),sum[0][i]=1;
      	for(int i=k;i;--i) sum[1][i]=sum[1][i+1]+dp[1][i];
      	g[1]=(-P+1);g[2]=sum[1][1]*(-P+1);
      	for(int i=2;i<=k;++i){
      		for(int j=1;j<=k/i;++j){
      			modint s=0;
      			for(int x=1;x<=i;++x) s+=sum[x-1][j+1]*sum[i-x][j];
      			dp[i][j]=s*p[j]*(-P+1);
      		}
      		for(int j=k/i;j;--j) sum[i][j]=sum[i][j+1]+dp[i][j];
      		g[i+1]=sum[i][1]*(-P+1);
      	}
      	g[0]=1;for(int i=1;i<=k+1;++i) g[i]=-g[i];
      	return Bostan(g,n+1)*((-P+1).inv());
      }
      
      int main(){
      	ios::sync_with_stdio(0);cin.tie(0);cout.tie(0);
      	cin>>n>>k>>x>>y;P=modint(y).inv()*x;p[0]=1;
      	for(int i=1;i<=k+2;++i) p[i]=p[i-1]*P;
      	modint res1=work(k-1),res2=work(k);
      	cout<<res2-res1<<'\n';
      //	cout<<res1<<'\n'; 
      	return 0;
      }
      
      • 1

      信息

      ID
      6613
      时间
      3000ms
      内存
      512MiB
      难度
      10
      标签
      递交数
      4
      已通过
      2
      上传者