[AGC021F] Trinity 简要题解
题面
题解
考虑以列为维度进行DP。
设\(f(i,j)\)表示\(i\)行每行都至少有一格为黑色的大小为\(i\times j\)的表格的方案数,那么最终答案就是\(\sum\binom{n}{i}f(i,m)\)。
考虑如何转移:初始表格为 \(0\times 0\) , \(f(0,0)=1\) ,每次添加一列,在其中添加若干行(不一定都在原来行的下面),新添加的行的 \(A_i\) 等于新添加的列数。
假设当前列数为 \(j\) ,分两种情况讨论:
- 若行数不变,假设当前行数为 \(i\) ,列数 \(j\rightarrow j+1\) 。所有行的 \(A_i\) 不变,只考虑 \(B_{j+1},C_{j+1}\) 的方案数即可。若不填黑色,则方案数为 \(1\) ;若填一格黑色,方案数为 \(i\) ;若填多格黑色,只需考虑最上面一格和最下面一格的位置,方案数为 \(\binom{i}{2}\) 。因此总贡献为 \(f_{i,j}(1+i+\binom{i}{2})\)
- 若行数由 \(k\) 变为 \(i\) , \(j\rightarrow j+1\) ,新行的 \(A_i=j+1\)。考虑怎么计算 \(B_{j+1},C_{j+1}\) 。发现 \(B,C\) 可能由新的行得到,也可能由原来的行得到,直接计算不太好计算。考虑 \(B,C\) 之外一格的位置,即 \(B-1,C+1\) ,那么问题变成了从 \(i+2\) 的位置中选取 \(i-k+2\) 个位置,其中 \(i+2\) 个位置对应 \([0,i+1]\) ,选取的 \(i-k\) 个位置为新加的行,选取的 \(2\) 个位置为 \(B-1,C+1\) ,容易发现,每一种选择方案都能对应一种黑色的安排方式。贡献为 \(\sum_{k<i}\binom{i+2}{i-k+2}f_{k,j}\)
得到
\[f_{i,j+1} = (1+i+\binom{i}{2})f_{i,j}+\sum_{k<i}\binom{i+2}{i-k+2}f_{k,j}
\]
即
\[f_{i,j} = (1+i+\binom{i}{2})f_{i,j-1}+\sum_{k<i}\binom{i+2}{i-k+2}f_{k,j-1}
\]
直接做复杂度 \(O(n^2m)\)
发现转移左边可以 \(O(n)\) 直接转移,考虑右边的部分:
\[\sum_{k<i}\frac{(i+2)!}{(i-k+2)!k!}f_{k,j-1}
\]
令 \(g_1=\frac{1}{(i+2)!}\) ,\(g2=\frac{1}{i!}f_{i,j-1}\) ,得到:
\[(i+2)!\sum_{k<i}g_1(i-k)g_2(k)
\]
\(NTT\) 做一下即可。
复杂度 \(O(nmlogn)\)
#include<bits/stdc++.h>
#define ll long long
#define uit unsigned int
//#define int long long
using namespace std;
const int N=8e3+10,NTTNUM=N*8;
const int M=998244353;
int n,m,lim,limbit,revid[NTTNUM],g1[NTTNUM],g2[NTTNUM],dp[N][N],pw[N],cir[N],ans=0;
int ftp(int b,int p,int mod){
int r=1;
while(p){
if(p&1) r=1ll*r*b%mod;
b=1ll*b*b%mod;
p>>=1;
}
return r;
}
int INV(int x,int mod){
return ftp(x,mod-2,mod);
}
void NTT(int len,int *arr,int sign){
for(int i=0;i<len;i++){
if(i<revid[i]) swap(arr[i],arr[revid[i]]);
}
for(int l=2;l<=len;l<<=1){
int Wn=ftp((sign==1)?3:INV(3,M),(M-1)/l,M);
for(int be=0;be<len;be+=l){
int w=1;
for(int i=0;i<(l>>1);i++,w=1ll*w*Wn%M){
int tmp1=arr[be+i],tmp2=1ll*w*arr[be+i+(l>>1)]%M;
arr[be+i]=(tmp1+tmp2)%M;
arr[be+i+(l>>1)]=(tmp1-tmp2+M)%M;
}
}
}
if(sign==-1){
int inv=INV(len,M);
for(int i=0;i<len;i++) arr[i]=1ll*arr[i]*inv%M;
}
}
int C(int x,int y){
if(x<0||y<0||x<y) return 0;
return 1ll*pw[x]*cir[y]%M*cir[x-y]%M;
}
void solve(int R){
memset(g1,0,sizeof(int)*(lim));
memset(g2,0,sizeof(int)*(lim));
for(int i=1;i<=n;i++) g1[i]=cir[i+2];// pay attention to this i,it is from [1,n],which is because k<i
for(int i=0;i<=n;i++) g2[i]=1ll*cir[i]*dp[i][R-1]%M;
for(int i=0;i<=n;i++) dp[i][R]=(dp[i][R]+1ll*(1ll+i+C(i,2))%M*dp[i][R-1]%M)%M;
NTT(lim,g1,1);NTT(lim,g2,1);
for(int i=0;i<lim;i++) g1[i]=1ll*g1[i]*g2[i]%M;
NTT(lim,g1,-1);
for(int i=1;i<=n;i++) dp[i][R]=(dp[i][R]+1ll*g1[i]*pw[i+2]%M)%M;
}
int main(){
// freopen(".in","r",stdin);
// freopen(".out","w",stdout);
pw[0]=1;for(int i=1;i<N;i++) pw[i]=1ll*pw[i-1]*i%M;
cir[N-1]=INV(pw[N-1],M);for(int i=N-2;i>=0;i--) cir[i]=1ll*cir[i+1]*(i+1ll)%M;
scanf("%d%d",&n,&m);
for(lim=1,limbit=0;lim<=(n<<1);lim<<=1,limbit++);
for(int i=1;i<lim;i++) revid[i]=(revid[i>>1]>>1)|((i&1)<<(limbit-1));
dp[0][0]=1;
for(int i=1;i<=m;i++) solve(i);
for(int i=0;i<=n;i++) ans=(ans+1ll*C(n,i)*dp[i][m]%M)%M;
printf("%d\n",ans);
return 0;
}

浙公网安备 33010602011771号