题解:P17316 [KismetOI 2026 I] 孤独绽放的彼岸花
lailai0916 · · 题解
题意简述
给定一棵以
解题思路
令
把不降关系反向会更便于计数。令
若某个深度大于
先计算全部合法局面。对非根节点
根节点的值固定为
叶子满足
代码计算每个
再计算事件的补集。补集要求每个深度为
补集条件等价于这些点均满足
考虑深度不超过
若
若
这里仍然只有多项式乘法和前缀和。相同的次数归纳说明
根节点的
用补集计数除以总方案数,最终答案为:
若树的最大深度小于
每个多项式保存
参考代码
#include <bits/stdc++.h>
using namespace std;
using ll=long long;
const int N=3005;
const int mod=998244353;
int n,k,mx;
ll s,x,l;
int dep[N],siz[N];
int fac[N],ifac[N],pre[N],suf[N];
int f[N][N],g[N][N];
vector<int> e[N];
ll Pow(ll x,ll y)
{
x%=mod;
ll res=1;
while(y)
{
if(y&1)res=res*x%mod;
x=x*x%mod;
y>>=1;
}
return res;
}
int calc(int *a,ll x,int d)
{
x%=mod;
if(x<=d)return a[x];
pre[0]=1;
for(int i=0;i<=d;i++)pre[i+1]=(int)((ll)pre[i]*(x-i+mod)%mod);
suf[d+1]=1;
for(int i=d;i>=0;i--)suf[i]=(int)((ll)suf[i+1]*(x-i+mod)%mod);
ll ans=0;
for(int i=0;i<=d;i++)
{
ll res=(ll)a[i]*pre[i]%mod*suf[i+1]%mod*ifac[i]%mod*ifac[d-i]%mod;
if((d-i)&1)ans-=res;
else ans+=res;
}
return (int)((ans%mod+mod)%mod);
}
void dfs(int u)
{
siz[u]=1;
mx=max(mx,dep[u]);
for(int i=0;i<=n;i++)f[u][i]=1;
for(auto v:e[u])
{
dep[v]=dep[u]+1;
dfs(v);
siz[u]+=siz[v];
for(int i=0;i<=n;i++)f[u][i]=(int)((ll)f[u][i]*f[v][i]%mod);
}
if(u!=1)
{
for(int i=1;i<=n;i++)
{
f[u][i]+=f[u][i-1];
if(f[u][i]>=mod)f[u][i]-=mod;
}
}
}
void dfs2(int u)
{
if(dep[u]==k)
{
int val=calc(f[u],l-1,siz[u]);
fill(g[u],g[u]+n+1,val);
return;
}
for(int i=0;i<=n;i++)g[u][i]=1;
for(auto v:e[u])
{
dfs2(v);
for(int i=0;i<=n;i++)g[u][i]=(int)((ll)g[u][i]*g[v][i]%mod);
}
if(u!=1)
{
int val=calc(f[u],l-1,siz[u]);
for(int i=0;i<=n;i++)
{
g[u][i]+=i?g[u][i-1]:val;
if(g[u][i]>=mod)g[u][i]-=mod;
}
}
}
void solve()
{
cin>>n>>k>>s>>x;
for(int i=1;i<=n;i++)e[i].clear();
for(int i=2;i<=n;i++)
{
int p;
cin>>p;
e[p].push_back(i);
}
dep[1]=0;
mx=0;
dfs(1);
if(mx<k)
{
cout<<0<<'\n';
return;
}
if(x<=s)
{
cout<<1<<'\n';
return;
}
l=x-s;
dfs2(1);
int all=calc(f[1],x,n);
int bad=calc(g[1],s,n);
cout<<(ll)(all-bad+mod)%mod*Pow(all,mod-2)%mod<<'\n';
}
int main()
{
ios::sync_with_stdio(false);
cin.tie(nullptr);
fac[0]=1;
for(int i=1;i<N;i++)fac[i]=(int)((ll)fac[i-1]*i%mod);
ifac[N-1]=(int)Pow(fac[N-1],mod-2);
for(int i=N-1;i;i--)ifac[i-1]=(int)((ll)ifac[i]*i%mod);
int t;
cin>>t;
while(t--)solve();
return 0;
}