学习笔记 · 2026-08-06

NOIP2025T3树的价值

是在参加南京大学办的计算理论之美活动期间和舍友随机跳题跳到的, 结果这一段时间一直在想到现在才想清楚.

简述题意

题意

给定一个有 $n(\le 8000)$ 个节点, 树高不超过 $m(\le 800)$ 的有根树 $T$ , 树根是 $1$ . 要给每个节点 $i$ 安排一个权值 $a_i$ , 令集合

$$S_i:=\{a_j:i\ 是\ j\ 的祖先\}$$

求所有安排权值的方案中

$$\sum_i{\rm mex}(S_i)$$

的最大值是多少.

一开始的想法

手玩几个样例发现贪心是不对的, 一个可能的情况是两个子树在合并的时候, 其中一个故意牺牲掉若干个节点(让这些节点的 mex 值不严格大于自己的子树)来使得根节点的 mex 值更大. 这就牵扯到每个牺牲掉的节点被贡献到哪里, 我们需要记录下子树里有多少个已经牺牲但是没被贡献掉的节点, 记$d_i:={\rm mex}(S_i)$ , 于是我们就可以有一个最初的 DP 做法.

DP1

令状态 $f_{i,j,k}$ 表示以 $i$ 为根的子树中, 在满足 $d_i=j$ 且还剩 $k$ 个牺牲但是没被贡献的节点的前提下, 这个子树对答案的贡献的最大值是多少.

这显然是一个树上背包问题, 复杂度是 $O(n^4)$ . 我们称被牺牲的节点为黑点, 没被牺牲的节点为白点, 准确来讲, 一个节点 $i$ 是白的当且仅当

$$a_i+1=d_i$$

之后我们很容易可以得到一个观察

观察1

存在至少一种最优方案满足, 若 $i$ 是 $j$ 的祖先, 则 $a_i>a_j$ .

在这一前提下, 每个黑点一定贡献到自己的最近白色祖先. 现在发现优化的瓶颈卡在背包上, 树上背包统计贡献的方法是在节点 $i$ 处把对答案的贡献 $\mathrm{mex}(S_i)$ 加入. 我们可以考虑重新分配贡献, 把 $\mathrm{mex}(S_i)$ 拆分到会贡献到它身上的 $i$ 的子树的节点里. 同时假设 $i$ 的儿子节点是 $v_1,...,v_k$ 并且有 $l$ 个黑点贡献到 $i$ (所以 $i$ 是白的) , 那么一定有

$$d_i=\max\{d_{v_j}:1\le j\le k\}+l+1$$

而如果 $i$ 是黑的, 那么此时 $l=0$ , 就有

$$d_i=\max\{d_{v_j}:1\le j\le k\}$$

于是我们可以钦定 $i$ 的 $d$ 是从谁那里继承过来的, 这里我们就可以获得第二个观察

观察2

在满足观察1 的前提下, 每个非叶子白点 $i$ , 若有 $l$ 个黑点贡献到 $i$ , 那么总存在 $i$ 的儿子 $j$ 使得 $d_i=d_j+l+1$ 且这个 $j$ 也是白的.

并且我们不用刻意去维持 $\max$ , 因为答案要求最大值就天然帮我们排除掉了从非最大值的儿子那里继承过来的情形. 在这种视角下整棵树被这种继承关系拆分称了若干条链, 于是每个白点会使得自己到链头的每个节点的 $d$ 增加 $1$ ; 而对于黑点而言, 假定黑点 $i$ 的最近白色祖先为 $x$ , 则黑点 $i$ 贡献到 $x$ 上会使得 $x$ 到 $x$ 所在的链头的每个节点的 $d$ 增加 $1$ . 于是对于节点 $i$ 我们只需要记录下: (1) $i$ 所在的链头的深度是谁(确定深度就可以确定祖先); (2) $i$ 的链头的最近白色祖先到它(白色祖先)的链头的距离. 这样我们就得到了第二个 DP 做法.

DP2

令 $f_{i,j,k,c}$ 表示钦定 $i$ 的颜色为 $c$ ( $0$ 是白色, $1$ 是黑色) , $i$ 的链头的深度为 $j(\le\mathrm{dep}[i])$ 且 $i$ 所在的链的链头的最近白色祖先到自己的链头的距离为 $k$ 的前提下, 以 $i$ 为根的子树的贡献最大值是多少.

转移

我们假设 $i$ 的儿子集合是 $son$ . 若 $i$ 的颜色为 $0$ 则显然存在某个白色儿子继承过来, 所以有

$$f_{i,j,k,0}=j-{\rm dep}[i]+1+\max_{x\in son}\{f_{x,j,k,0}+\sum_{\substack{y\in son\\y\not= x}}\max(f_{y,{\rm dep}[y],{\rm dep}[i]-j+1,0},f_{y,{\rm dep}[y],{\rm dep}[i]-j+1,1})\}$$

若 $i$ 的颜色为 $1$ , 那么它所继承的儿子的节点既可以是 $0$ 也可以是 $1$ 所以有

$$\begin{aligned}f_{i,j,k,1}=k+\max_{x\in son}\{\max(f_{x,j,k,0},f_{x,j,k,1})+\sum_{\substack{y\in son\\y\not=x}}\max(f_{y,{\rm dep}[y],k,0},f_{y,{\rm dep}[y],k,1})\}\end{aligned}$$

于是我们就得到了一个 $O(nm^2)$ 的做法, 可以得到多达 76 分.

长剖优化

我们还要继续优化(现在状态是 $O(nm^2)$ 的, 我们没办法优化转移了), 首先我们发现转移的时候绝大多数情况下都只需要用到 $f_{y,\mathrm{dep}[y],\cdot,\cdot}$ , 也就是 $y$ 恰好是 $y$ 所在链的链头的情形(这个时候我们不去关心 $y$ 的颜色, 因为这个链已经到头了), 我们记 $A_{y,k}:=\max(f_{y,\mathrm{dep}[y],k,0},f_{y,\mathrm{dep}[y],k,1})$ . 之后我们考虑不在 $f$ 上进行转移而在 $A$ 上进行, 这么做的原因是我们发现当一个决策在 $c=0$ 的时候, 值 $k$ 实际上是没有用的, 而当 $c=1$ 之后指标 $j$ 不参与到转移中, 单纯就是拿来标记我这个链什么时候该结束了, 所以着相当于尽管我们有三个维度, 但是在转移的前一段只有前两个维度是有用的, 在后一段只有第一,三维是有用的, 因而直接记录三个维度实际上是浪费了一个维度, 我们把这个浪费的维度压掉, 于是可以考虑越过 $f$ 而直接在 $A$ 上转移.

对于 $A_{i,j}$ 而言, 我们转移相当于先找到一个从点 $i$ 延伸到叶子的链 $\langle u_0,u_1,u_2,...,u_r\rangle$ (其中 $u_0=i$ , 之后假定链的长度为 $r+1$ ). 于是 观察2 告诉我们, 每个链的黑点一定在上面, 白点一定在下面. 所以我们只需要枚举这个分界点 $k$ , 这种前提下这个以 $i$ 为根的子树的贡献值最大就是

$$[\sum_{s\le k} j+(\sum_{\substack{x\in son(u_s)\\x\not=u_{s+1}}}A_{x,j})]+[\sum_{s>k}\mathrm{dep}[u_s]-\mathrm{dep}[i]+1+(\sum_{\substack{x\in son(u_s)\\x\not=u_{s+1}}}A_{x,\mathrm{dep}[x]-\mathrm{dep}[i]})]$$

对于后面这一坨, 我们实际上只关心两个信息: (1) $u_{s+1}$ 是谁 (2) $\mathrm{dep}[i]$ 的值是多少 , 这是因为后面这一坨东西可以看作是从 $u_{s+1}$ 出发一直延伸到叶子的一条链, 需要用到的外部信息只有 $i$ 的深度, 于是我们把这一坨东西记为 $g_{u_{s+1},\mathrm{dep}[i]}$ . 尽管如此我们还暂时没有得到优化, 我们需要最后一个观察

观察3

对于任意的 $i$ 以及 $x>y$ 有 $A_{i,x}\ge A_{i,y}$ .

这个结论需要结合我们对于 $A_{i,j}$ 的计算在树上实施归纳. 于是我们不需要枚举 $k$ 了, 因为这个时候选取 $k=j$ 就是最优的, 当然我们还需要处理一下边角的情况. 于是对于 $A_{i,j}$ 而言, 我们选取的 $u_{s+1}$ 就一定是满足 $\mathrm{dep}[u_{s+1}]-\mathrm{dep}[i]=j$ , 这对 $u_{s+1}$ 的深度产生了限制, 并且假定我们在计算 $A_{i,j}$ 时对 $i$ 子树里每一个 $x$ 都维护从 $x$ 出发一直到 $i$ 的上面式子的值 $val_x$ , 那么每次合并相当于给相同子树里的 $val_x$ 同时加上同一个数, 这就可以拿长链剖分维护. 具体的, 令链 $chain_{i,j}[k]$ 表示以 $i$ 为根的子树中, 在固定参数 $j$ 的前提下, 距离 $i$ 为 $k$ 的 $x$ 的 $val_x$ 的最大值, 对于全体加上某个数我们可以额外对每个链维护一个懒标记, 我们计算 $A_{i,j}$ 只需要查询 $chain_{i,j}[j]$ 即可, 每次合并子树就是相当于继承重儿子之后暴力合并轻儿子, 这样做的时间复杂度就是 $O(nm)$ 了.

代码如下


#include <bits/stdc++.h>

// #define int long long
#define ar(u,n) array<u,n>
#define vec vector

using namespace std;

inline int read(){
    int x=0,f=1;char ch=getchar();
    for(;!isdigit(ch);ch=getchar())f^=ch=='-';
    for(;isdigit(ch);ch=getchar())x=x*10+(ch^48);
    return f?x:-x;
}
template<typename T>inline void chmin(T &x,T y){x>y?x=y:y;}
template<typename T>inline void chmax(T &x,T y){x<y?x=y:y;}

const int mo=998244353,inf=1e9;

mt19937 rnd(time(0));

inline void red(int &x){(x>=mo)&&(x-=mo);}
inline int qpow(int x,int t=mo-2){
    int ret=1;
    for(;t;t>>=1,x=x*x%mo)if(t&1)ret=ret*x%mo;
    return ret;
}

vec<vec<int>> edge,f,g;
vec<int> len,son,dep;

int n,m;

struct chain{
    int add;
    vector<int> val;
    int len(){
        return val.size();
    }
};
vec<vec<chain>> tr;

void dfs(int u,int fa){
    len[u]=dep[u]=dep[fa]+1;
    int &s=son[u];
    for(int v:edge[u]){
        dfs(v,u);
        if(len[v]>len[s]){
            s=v;
        }
    }
    if(s)len[u]=len[s];
    for(int i=1;i<=dep[u];++i){
        int sum=0,l=dep[u]-i+1;
        chmax(g[u][i],l);
        for(int v:edge[u]){
            sum+=f[v][l];
        }
        for(int v:edge[u]){
            chmax(g[u][i],g[v][i]+l+sum-f[v][l]);
        }
    }
    for(int i=1;i<=dep[u];++i){//
        int sum=0;
        for(int v:edge[u]){
            sum+=f[v][i];
        }
        tr[u][i].val.clear();
        tr[u][i].add=0;
        if(s){
            swap(tr[u][i],tr[s][i]);
            int &add=tr[u][i].add,usiz=tr[u][i].len();
            add+=sum-f[s][i];
            for(int v:edge[u])if(v!=s){
                int vadd=tr[v][i].add,siz=tr[v][i].len();
                vadd+=sum-f[v][i]-add;
                for(int j=1;j<=siz;++j){
                    chmax(tr[u][i].val[usiz-j],tr[v][i].val[siz-j]+vadd);
                }
                vec<int>().swap(tr[v][i].val);
            }
        }
        tr[u][i].val.push_back(g[u][dep[u]-i+1]-tr[u][i].add);
        if(tr[u][i].len()>=i){
            f[u][i]=tr[u][i].add+tr[u][i].val[tr[u][i].len()-i]+(i-1)*i;
        }else{
            f[u][i]=tr[u][i].add+tr[u][i].len()*i;
        }
        chmax(f[u][i],g[u][dep[u]]);
    }
}

void solve(){
    n=read(),m=read()+1;
    edge.assign(n+1,{});
    f.resize(n+1);
    g.resize(n+1);
    tr.resize(n+1);
    len.assign(n+1,0);
    dep.assign(n+1,0);
    son.assign(n+1,0);
    for(int i=1;i<=n;++i){
        f[i].assign(m+1,-inf);
        g[i].assign(m+1,-inf);
        tr[i].resize(m+1);
        if(i>1){
            int x=read();
            edge[x].push_back(i);
        }
    }
    dfs(1,0);
    printf("%lld\n",g[1][1]);
    // for(int i=1;i<=n;++i,puts(""))for(int j=1;j<=len[1];++j){
    //     printf("g[%lld][%lld]=",i,j);
    //     if(g[i][j]<0){
    //         printf("-inf, ");
    //     }else{
    //         printf("%lld, ",g[i][j]);
    //     }
    // }
    // for(int i=1;i<=n;++i,puts(""))for(int j=1;j<=len[1];++j){
    //     printf("f[%lld][%lld]=",i,j);
    //     if(f[i][j]<0){
    //         printf("-inf, ");
    //     }else{
    //         printf("%lld, ",f[i][j]);
    //     }
    // }
    return;
}

signed main(){
    for(int cas=read();cas--;){
        solve();

    }
    return 0;
}