[NOI 2026] 彩虹树

· · 题解

https://uoj.ac/problem/1107

考虑如何判定一个 x_i 数组是否可行。

有如下过程:

记录 z_i 表示 i 子树颜色和 i 祖先颜色的交集大小。

(x_1,z_1) (x_2,z_2) ... (x_k,z_k) 可以转移到 (X=\sum x_i-C,Z=\sum z_i-D),其中

然后要考虑加入子树的根 u

最后如果 z_1 = 0 则判定成功。

考虑贪心。整个过程可以贪心最小化 z_i。然后由于此时可以增大任何的 z,前半部分条件可以转化为:

前半部分的转移是最大化 D,后面是能减就减。

f_{i,x,z} 表示子树颜色个数为 x,最小化的 z_iz

合并若干子树之后我们要记录的信息有 (\sum x,\max x,\sum z,\max z)。然后枚举 X 减少的 C,根据这四个信息,贪心确定 D 是多少。

直接暴力是 O(n^6) 的。可以前缀和优化到 O(n^5)

然后是下一步优化,两位验题人对这步的难度反馈并不相同,不太确定难度是什么(

考虑优化,可以把一个 (\sum x,\max x,\sum z,\max z) 能贡献到的 (x,z) 位置画出来。发现一开始是 X,Z 两个一起减,后面变成只有 X 减。

也就是两段折线,一段是斜下,一段是水平。

发现考虑差分,就只需要 DP 出斜线的拐点是哪里,拆成了三个拐点的信息,需要的信息变成了 (\sum x,\sum z),(\sum x,\max z),(\max x,\max z)

对这三个分别 DP,复杂度瓶颈是 (\sum x,\sum z) 的二维树形背包,是 O(n^4) 的。用 NTT 优化瓶颈部分可以做到 O(n^3\log n)

(bonus:做到 O(n^3))

// what is matter? never mind. 
//#pragma GCC optimize("Ofast")
//#pragma GCC optimize("unroll-loops")
#include "rainbow.h"
//#pragma GCC target("sse,sse2,sse3,sse4,popcnt,abm,mmx,avx,avx2")
#include<bits/stdc++.h>
#define For(i,a,b) for(int i=(a);i<=(b);++i)
#define Rep(i,a,b) for(int i=(a);i>=(b);--i)
#define ll long long
using namespace std;
inline int read()
{
    char c=getchar();int x=0;bool f=0;
    for(;!isdigit(c);c=getchar())f^=!(c^45);
    for(;isdigit(c);c=getchar())x=(x<<1)+(x<<3)+(c^48);
    return f?-x:x;
}

#define fi first
#define se second
#define pb push_back
#define mkp make_pair
typedef pair<int,int>pii;
typedef vector<int>vi;

#define maxn 205
#define inf 0x3f3f3f3f

int n,siz[maxn];
vi e[maxn];

modint f[maxn][maxn][maxn];
modint dp1[maxn][maxn],dp2[maxn][maxn],dp3[maxn][maxn],tmp[maxn][maxn];
modint difH[maxn][maxn],difD[maxn][maxn],g[maxn][maxn];

int lx[maxn*maxn],lz[maxn*maxn];
modint lv[maxn*maxn];

inline void clear2(modint a[maxn][maxn],int n,int m){
    For(i,0,n)For(j,0,m)a[i][j].x=0;
}
inline void copy2(modint a[maxn][maxn],modint b[maxn][maxn],int n,int m){
    For(i,0,n)For(j,0,m)a[i][j]=b[i][j];
}

void dfs(int u)
{
    siz[u]=1;
    for(int v:e[u]){
        dfs(v);
        siz[u]+=siz[v];
    }

    int S=siz[u];

    clear2(dp1,S,S);
    clear2(dp2,S,S);
    clear2(dp3,S,S);

    dp1[1][1]=1;
    dp2[0][1]=1;
    dp3[1][1]=1;

    int cur=1;
    for(int v:e[u]){
        int sv=siz[v],ns=cur+sv,tot=0;

        For(x,1,sv)For(z,0,x)if(f[v][x][z].x){
            lx[++tot]=x;
            lz[tot]=z;
            lv[tot]=f[v][x][z];
        }

        clear2(tmp,ns,ns);
        For(sx,1,cur)For(sz,0,sx)if(dp1[sx][sz].x){
            modint w=dp1[sx][sz];
            For(i,1,tot)tmp[sx+lx[i]][sz+lz[i]]+=w*lv[i];
        }
        copy2(dp1,tmp,ns,ns);

        clear2(tmp,ns,ns);
        For(sd,0,cur)For(mz,1,cur)if(dp2[sd][mz].x){
            modint w=dp2[sd][mz];
            For(i,1,tot){
                int nz=mz>lz[i]?mz:lz[i];
                tmp[sd+lx[i]-lz[i]][nz]+=w*lv[i];
            }
        }
        copy2(dp2,tmp,ns,ns);

        clear2(tmp,ns,ns);
        For(mx,1,cur)For(mz,1,mx)if(dp3[mx][mz].x){
            modint w=dp3[mx][mz];
            For(i,1,tot){
                int nx=mx>lx[i]?mx:lx[i];
                int nz=mz>lz[i]?mz:lz[i];
                tmp[nx][nz]+=w*lv[i];
            }
        }
        copy2(dp3,tmp,ns,ns);

        cur=ns;
    }

    clear2(difH,S+1,S+1);
    clear2(difD,S+1,S+1);
    clear2(g,S,S);

    For(B,1,S)For(D,1,B)if(dp3[B][D].x){
        difH[B][D]+=dp3[B][D];
    }

    For(sd,0,S)For(D,1,S)if(dp2[sd][D].x){
        int t=sd+D;
        difH[t][D]-=dp2[sd][D];
        difD[sd][D]+=dp2[sd][D];
    }

    For(A,1,S)For(C,0,A)if(dp1[A][C].x){
        int d=A-C;
        difD[d][C+1]-=dp1[A][C];
    }

    For(z,0,S){
        modint now=0;
        For(x,1,S){
            now+=difH[x][z];
            if(now.x)g[x][z]+=now;
        }
    }

    For(d,0,S){
        modint now=0;
        for(int z=0;d+z<=S;++z){
            now+=difD[d][z];
            int x=d+z;
            if(x&&now.x)g[x][z]+=now;
        }
    }

    For(x,1,S)For(z,1,x)if(g[x][z].x){
        f[u][x][z]+=g[x][z];
        f[u][x][z-1]+=g[x][z];
    }
}

int rainbow(int cc,int nn,vector<int>F) {
    n=nn;
    For(i,2,n) e[F[i-1]+1].pb(i);
    dfs(1);
    modint ans=0;
    For(x,1,n)ans+=f[1][x][0];
    return ans.x;
}