题解:P9962 [THUPC 2024 初赛] 一棵树

· · 题解

闵和好题。

首先介绍一下闵和(闵可夫斯基和)

两个凸包的闵和可以看做是把一个凸包的顶点不断换成另一个凸包顶点后最外层组成的凸多边形,这里不详细展开,主要分析下凸壳时的情况。

闵和主要是用来优化 \left(\min,+\right) 卷积的一个方法,有重要性质:

两个下凸壳的 \left(\min,+\right) 卷积依然是一个下凸壳。

假设要把下凸壳 f,g\left(\min,+\right) 卷积,即要计算

h_i=\min_{j+k=i}\set {f_j+g_k}

那么此时

h 的差分数组就是 f,g 差分数组的归并。

有这个重要性质,就可以考虑用平衡树等数据结构快速维护凸壳的差分数组。

回到此题,显然可以使用一个 \mathrm{dp},定义 f_{u,i} 表示考虑 u 为根的子树内,染了 i 个黑点时的最小代价,类比树形背包,得到转移:

f_{u,i}\leftarrow \min_{j+k=i}\set{f_{u,j}+f_{v,k}}

然后要加上代价,考虑 \mathrm{fa}_u\to u 这条边,上面有 K-i 个,下面有 i 个,于是差为 |K-i-i|=|K-2i|,于是

f_{u,i}\leftarrow f_{u,i}+|K-2i|

初值 f_{u,0}=f_{u,1}=0,最终目标:f_{1,k}

直接做背包时间和空间都不能接受,但是观察式子:

\min_{j+k=i}\set{f_{u,j}+f_{v,k}}

于是考虑使用刚刚提到的闵和优化一下,很明显,此时 f_uf_uf_v 进行 (\min,+) 卷积得到的,如果我们维护 f_u,f_v 的差分数组,即维护第 i 个位置为 f^{\prime}_if^{\prime}_{i-1} 的差分,那么如果支持快速进行归并,即在值域上进行合并,那么就可以得出 f_u 的差分数组。而对于最后的 |K-2i|,这是绝对值函数,分类讨论一下

m=\lfloor\frac{K}{2}\rfloor

于是我们相当于要用一个数据结构维护差分数组,其支持值域上合并,区间加上一个数。很明显使用平衡树即可。其值域上合并的复杂度为 O(\log^2 n),区间加可以打 \mathrm{tag},于是就可以在 O(n\log^2 n) 的时间内求解出 f 的差分数组。

那么求答案也是简单的。考虑到我们已经求出了差分数组,只要知道 f_{1,0},那么代入差分数组推下去就可以得到 f_{1,k}f_{1,0} 在原来的 \mathrm{dp} 中很明显值为 (n-1)k,于是从前往后加上差分数组就做完了。

代码:

#include <iostream>
#include <cstring>
#include <algorithm>
#include <random>
#include <ctime>

using namespace std;

const int N=500010;
typedef long long ll;
int n,k;
int h[N],e[N<<1],ne[N<<1],idx;

void adde(int u,int v) { e[idx]=v,ne[idx]=h[u],h[u]=idx++; }

mt19937 rnd(time(0));

namespace FHQ
{
    const int M=N*2;
    ll val[M],add[M];
    int pri[M],ch[M][2],sz[M],root[N],idx;

    inline int new_node(ll v)
    {
        int u=++idx;
        val[u]=v;
        pri[u]=rnd();
        sz[u]=1;
        return u;
    }

    inline void pushup(int u) { sz[u]=sz[ch[u][0]]+sz[ch[u][1]]+1; }

    inline void Add(int u,ll d)
    {
        add[u]+=d;
        val[u]+=d;
    }

    inline void pushdown(int u)
    {
        if (add[u])
        {
            Add(ch[u][0],add[u]);
            Add(ch[u][1],add[u]);
            add[u]=0;
        }
    }

    void split(int rt,ll k,int &rt1,int &rt2)
    {
        if (!rt)
        {
            rt1=rt2=0;
            return;
        }
        pushdown(rt);
        if (val[rt]<=k)
        {
            rt1=rt;
            split(ch[rt][1],k,ch[rt][1],rt2);
        }
        else
        {
            rt2=rt;
            split(ch[rt][0],k,rt1,ch[rt][0]);
        }
        pushup(rt);
    }

    void cntsplit(int rt,int k,int &rt1,int &rt2)
    {
        if (!rt)
        {
            rt1=rt2=0;
            return;
        }
        pushdown(rt);
        if (sz[ch[rt][0]]+1<=k)
        {
            rt1=rt;
            cntsplit(ch[rt][1],k-sz[ch[rt][0]]-1,ch[rt][1],rt2);
        }
        else
        {
            rt2=rt;
            cntsplit(ch[rt][0],k,rt1,ch[rt][0]);
        }
        pushup(rt);
    }

    int merge(int rt1,int rt2)
    {
        if (!rt1 || !rt2) return rt1^rt2;
        if (pri[rt1]<pri[rt2])
        {
            pushdown(rt1);
            ch[rt1][1]=merge(ch[rt1][1],rt2);
            pushup(rt1);
            return rt1;
        }
        else
        {
            pushdown(rt2);
            ch[rt2][0]=merge(rt1,ch[rt2][0]);
            pushup(rt2);
            return rt2;
        }
    }

    int Merge(int rt1,int rt2)
    {
        if (!rt1 || !rt2) return rt1^rt2;
        if (pri[rt1]>pri[rt2]) swap(rt1,rt2);
        pushdown(rt1);
        int x,y;
        split(rt2,val[rt1],x,y);
        ch[rt1][0]=Merge(ch[rt1][0],x);
        ch[rt1][1]=Merge(ch[rt1][1],y);
        pushup(rt1);
        return rt1;
    }

    ll qrys(int u)
    {
        if (!u) return 0;
        pushdown(u);
        return val[u]+qrys(ch[u][0])+qrys(ch[u][1]);
    }

    void adj(int u)
    {
        pushdown(u);
        if (ch[u][0]) adj(ch[u][0]);
        if (ch[u][1]) adj(ch[u][1]);
        pushup(u);
    }

    void print(int u)
    {
        pushdown(u);
        if (ch[u][0]) print(ch[u][0]);
        cout << val[u] << " ";
        if (ch[u][1]) print(ch[u][1]);
    }

    void prt(int rt) { print(rt);cout << "\n"; }
}
using namespace FHQ;

inline void Mrg(int u,int v)
{
    root[u]=Merge(root[u],root[v]);
    if (sz[root[u]]>k)
    {
        int x,y;
        cntsplit(root[u],k,x,y);
        root[u]=x;
    }
}

void dfs(int u,int fa)
{
    root[u]=merge(root[u],new_node(0));
    root[u]=merge(root[u],new_node(1e18));

    for (int i=h[u];i!=-1;i=ne[i])
    {
        int v=e[i];
        if (v==fa) continue;
        dfs(v,u);

        Mrg(u,v);
    }

    // |k-2x|

    if (fa!=-1)
    {
        int mid=k/2;
        if (k&1)
        {
            int x,y,z;
            cntsplit(root[u],mid,x,y);
            Add(x,-2);
            if (sz[y]>1)
            {
                cntsplit(y,1,y,z);
                Add(z,2);
                root[u]=merge(merge(x,y),z);
            }
            else root[u]=merge(x,y);
        }
        else
        {
            int x,y;
            cntsplit(root[u],mid,x,y);
            Add(x,-2);
            Add(y,2);
            root[u]=merge(x,y);
        }
    }
}

int main()
{
    ios::sync_with_stdio(false);
    cin.tie(0),cout.tie(0);

    cin >> n >> k;
    memset(h,-1,sizeof(h));
    for (int i=1;i<n;i++)
    {
        int u,v;
        cin >> u >> v;
        adde(u,v);adde(v,u);
    }

    dfs(1,-1);

    ll res=1ll*(n-1)*k;
    int x,y;
    cntsplit(root[1],k,x,y);
    res+=qrys(x);

    cout << res << "\n";

    return 0;
}