题解:P16365 [OOI 2026] Tasks from Sasha

· · 题解

直接 dp 也太难搞了,来点容斥做法。

我们发现,如果每个 S_i,去掉非空的限制,且把 lca=x_i 改成 lcay_i 的子树中,那么问题是非常好做的。令 a_u 表示 uu 的祖先中有几个 Y 中元素,然后每个 u 就在 a_u 个集合或者不选中任选一个,答案是 \prod (a_u+1)

然后我们考虑把每个限制容斥成上面的形式,假设 x_i=u,其儿子集合为 son(u),那么就可以容斥成以下形式,其中 y_i=\varnothing 的意思是强制 S_i=\varnothing,也就相当于删掉这个任务:

\\=[y_i=u]-\sum\limits_{v\in son(u)}[y_i=v]+(|son(u)|-1)[y_i=\varnothing]

这样我们就可以设计 dp 了,令 f_{u,i} 表示考虑节点 u 的子树,如果节点 u 的子树的祖先中有 i 个点在 Y 中(不包括由于 x_i=u 产生的),那么节点 u 中所有的 Y\prod (a_u+1) 的和是多少,最终答案就是 f_{1,0}

于是结合上面容斥的式子,就有以下转移。

然后对 \prod\limits_{w\in son(u),w\neq v}f_{w,i} 这一项,预处理一下前后缀积。其他项直接算,就可以做到 O(nk)

然后我们注意到,如果 u\notin Xf_u 的式子非常简单,所以我们考虑对所有 X 中的点建虚树。如果某一个子树 u 中没有任何 X 中的点,我们称其为空子树,那么显然 f_{u,i}=(i+1)^{sz_u+1}。然后虚树上每条边跳过的一段就相当于给 f_{u} 的每个位置乘上 (i+1) 的若干次方,这些都是容易处理的。

但是这样有一个问题,如果某个虚树上的点下有非常多的空儿子,那么我们还是 O(deg\cdot k) 的处理的话复杂度还是会退化到 O(nk)。但是我们注意到如果两个空儿子的大小相同,那么我们容易把他们一起处理了,所以此时每个节点处理空子树的复杂度就不超过 O(\sqrt{e_i}),其中 e_i 表示这个节点的所有空儿子大小之和。又 \sum e_i=O(n),所以简单的不等式放一下可以发现 \sum\sqrt{e_i}=O(\sqrt{nk})

于是现在我们的复杂度就是非空子树 O(k^2) 个 dp 状态,每次 O(1)。空子树 O(n^{0.5}k^{1.5}) 的个 dp 状态,但是每次需要一个快速幂所以带一个 \log n,不太能过,需要优化。

注意到,每次快速幂的指数的和也是 O(n) 的,且底数都是 1\sim k+1,所以快速幂的次数其实不超过 O(n^{0.5}k),记忆化一下即可。

所以总的复杂度是 O(n+k^2+n^{0.5}k^{1.5}+n^{0.5}k\log n)

代码,写的有点丑。

#include<bits/stdc++.h>
#define ll long long
#define N 1000009
#define K 2009
#define pb push_back
#define p 998244353ll
using namespace std;
int n,k;
ll pw[K<<1][K];
int pww[N],tt;
ll QPOW(ll x,ll y){
    ll z=1;while(y){
        if(y&1)z=z*x%p;
        x=x*x%p;y>>=1;
    }return z;
}
ll qpow(int x,int y){
    if(pww[y])return pw[pww[y]][x];
    pww[y]=++tt;
    for(int i=0;i<=k+1;i++)pw[tt][i]=QPOW(i,y);
    return pw[tt][x];
}
bitset<N> is;
ll f[K<<1][K],pre[K<<1][K],suf[K<<1][K];
vector<int> g[N];
int h[N][2],tot;
void dfs(int u){
    int cnt=0;if(is[u])cnt=100;
    for(int v:g[u])dfs(v),cnt+=(h[v][0]!=0);
    if(cnt<=1){
        for(int v:g[u]){
            h[u][0]|=h[v][0];
            h[u][1]+=h[v][1];
        }
        h[u][1]++;
        return;
    }
    ++tot;h[u][0]=tot;h[u][1]=0;
    if(!is[u]){
        cnt=1;
        for(int i=0;i<=k;i++)f[tot][i]=1;
        for(int v:g[u]){
            cnt+=h[v][1];
            if(h[v][0]){
                for(int i=0;i<=k;i++)f[tot][i]=f[tot][i]*f[h[v][0]][i]%p;
            }
        }
        for(int i=0;i<=k;i++)f[tot][i]=f[tot][i]*qpow(i+1,cnt)%p;
        return;
    }
    vector<array<int,2>> T;
    map<int,int> mp;
    for(int v:g[u]){
        if(h[v][0])T.pb({h[v][0],h[v][1]});
        else mp[h[v][1]]++;
    }
    for(auto t:mp)T.pb({-t.second,t.first});
    int sz=T.size();
    if(!sz){
        for(int i=0;i<=k;i++)f[tot][i]=1;
        return;
    }
    for(int i=0;i<=k;i++){pre[0][i]=1;suf[sz-1][i]=1;}
    for(auto [x,y]:T){
        if(x<0)continue;
        for(int i=0;i<=k;i++)f[x][i]=f[x][i]*qpow(i+1,y)%p;
    }
    for(int i=sz-1;i;--i){
        auto [x,y]=T[i];
        if(x>0)for(int j=0;j<=k;j++)suf[i-1][j]=suf[i][j]*f[x][j]%p;
        else for(int j=0;j<=k;j++)suf[i-1][j]=suf[i][j]*qpow(j+1,-x*y)%p;
    }
    for(int i=0;i<sz;i++){
        auto [x,y]=T[i];
        if(x>0){
            for(int j=0;j<=k;j++){
                pre[i+1][j]=pre[i][j]*f[x][j]%p;
                f[tot][j]+=pre[i][j]*suf[i][j]%p*f[x][j+1]%p;
            }
        }
        else{
            for(int j=0;j<=k;j++){
                pre[i+1][j]=pre[i][j]*qpow(j+1,-x*y)%p;
                f[tot][j]+=pre[i][j]*suf[i][j]%p*qpow(j+2,y)%p*qpow(j+1,(-1-x)*y)%p*(-x)%p;
            }
        }
    }
    int tmp=g[u].size();
    for(int i=0;i<=k;i++){
        f[tot][i]=pre[sz][i+1]*(i+2)%p+(-f[tot][i]+(tmp-1)*pre[sz][i])%p*(i+1);
        f[tot][i]=(f[tot][i]%p+p)%p;
    }
}
int main(){
    ios::sync_with_stdio(0);cin.tie(0);cout.tie(0);
    cin>>n>>k;
    for(int i=1,x;i<=k;i++)cin>>x,is[x]=1;
    for(int i=2,x;i<=n;i++)cin>>x,g[x].pb(i);
    dfs(1);
    if(h[1][0])cout<<f[h[1][0]][0]<<"\n";
    else cout<<1<<"\n";
    return 0;
}