题解:P17138 [KOI 2026 #1] 步道

· · 题解

题意简述

给定一棵点带拥挤度、边带长度的树。选择两个不同结点,使路径长度减去路径上的最大点权尽量大,并输出一组最优端点。

解题思路

最大点权不便直接在线段之间合并,可以把它改写成一个逐渐增大的阈值。

对任意阈值 X,只保留点权不超过 X 的结点及其之间的边,得到一片森林。记这片森林中端点不同的最长路径长度为 D_X。若森林还没有边,则暂不定义 D_X

原问题的最优满意度满足:

\operatorname{ans}=\max_X(D_X-X)

先证明从左到右的不等式。任取一条路径 P,设其长度为 S,最大点权为 M。整条路径都存在于阈值 M 对应的森林中,因此:

S\le D_M

所以 P 的满意度 S-M 不超过 D_M-M,原问题的最优值也不会超过右侧。

再证明反向不等式。阈值 X 下取得 D_X 的路径,其实际最大点权至多为 X。因此它的真实满意度至少为 D_X-X。右侧的每个候选值都能由一条合法路径达到或超过,等式成立。

一条连接 iC_i 的边何时出现,只取决于两个端点何时同时保留。定义它的出现阈值为:

B_i=\max(A_i,A_{C_i})

将所有边按 B_i 升序加入。处理完所有满足 B_i=X 的边后,并查集中的连通块恰好是阈值 X 森林的所有非平凡连通块。

不必单独处理只加入结点却没有加入边的阈值。这样的变化只会增加孤立点,不会产生端点不同的路径,也不会改变 D_X。在相邻两个边阈值之间,森林保持不变,而 D_X-XX 增大而减小。因此只在每组相同的 B_i 全部合并后统计答案即可。

接下来考虑如何在并查集中维护直径。每个连通块记录直径长度及两个端点。设两个连通块通过边 (x,y) 合并,边长为 w。新直径只有三类可能:

对一棵树中的任意结点 x,距离 x 最远的结点一定可以在某一对直径端点中找到。因此,只需比较第一个连通块的两个直径端点到 x 的距离,选出较远端点 u;再对第二个连通块和 y 得到 v。最长的跨块路径长度就是:

\operatorname{dist}(u,x)+w+\operatorname{dist}(y,v)

在两条旧直径与这条跨块路径中取最长者,便得到合并后连通块的直径。并查集合并的同时维护全局最长直径及端点。

树上距离通过倍增 LCA 计算。题目给出的父结点满足 C_i>i。以结点 N 为根后,从 N-11 递减枚举,就能直接计算每个结点的深度、根距离和倍增祖先,无须执行 DFS。

排序、并查集合并和距离查询的总时间复杂度为 O(N\log N),空间复杂度为 O(N\log N)

正确性证明

阈值等式的两侧已经分别得到上界和可达下界,所以最大化满意度等价于最大化 D_X-X。每条边恰在 X=B_i 时第一次出现,按 B_i 分组加入后,并查集维护的森林与对应阈值下的非平凡森林完全一致。没有边出现的中间阈值不会改变直径,只会使减去的阈值更大,因此不会遗漏最优答案。

合并两个连通块时,不经过新边的路径完全属于某个旧连通块,最长者就是两条旧直径之一。任何经过新边的路径都由第一侧到 x 的路径、新边和第二侧从 y 出发的路径组成。树的直径端点性质保证两侧的最远结点都能从各自的直径端点中选出,所以算法构造的跨块路径是所有跨块路径中最长的。

因此,每次合并后记录的都是新连通块的真实直径,全局记录也是当前森林的真实最长直径。对所有边阈值计算 D_X-X 并保留最大值,最终输出的两个端点必定构成满意度最大的合法路径。

参考代码

#include <bits/stdc++.h>
using namespace std;

using ll=long long;
using pli=pair<ll,int>;
const int N=300005;
const int K=20;
const ll inf=1LL<<62;
struct node
{
    int x,y;
    ll len;
};
int n,c[N],f[N],siz[N],dep[N];
int fa[K][N];
ll a[N],l[N],dis[N];
node d[N],mx;
pli p[N];
int find(int x)
{
    return f[x]==x?x:f[x]=find(f[x]);
}
int lca(int x,int y)
{
    if(dep[x]<dep[y])swap(x,y);
    for(int i=K-1;i>=0;i--)
    {
        if(dep[fa[i][x]]>=dep[y])x=fa[i][x];
    }
    if(x==y)return x;
    for(int i=K-1;i>=0;i--)
    {
        if(fa[i][x]!=fa[i][y])
        {
            x=fa[i][x];
            y=fa[i][y];
        }
    }
    return fa[0][x];
}
ll get_dis(int x,int y)
{
    int z=lca(x,y);
    return dis[x]+dis[y]-2*dis[z];
}
void check(node &x,node y)
{
    if(y.len>x.len)x=y;
}
pli get_far(int x,int y)
{
    pli i={get_dis(d[x].x,y),d[x].x};
    pli j={get_dis(d[x].y,y),d[x].y};
    return max(i,j);
}
void join(int x,int y,ll w)
{
    int u=find(x),v=find(y);
    if(u==v)return;
    pli p=get_far(u,x),q=get_far(v,y);
    node res=d[u];
    check(res,d[v]);
    check(res,{p.second,q.second,p.first+w+q.first});
    if(siz[u]<siz[v])swap(u,v);
    f[v]=u;
    siz[u]+=siz[v];
    d[u]=res;
    check(mx,res);
}
int main()
{
    ios::sync_with_stdio(false);
    cin.tie(nullptr);
    cin>>n;
    for(int i=1;i<=n;i++)cin>>a[i];
    for(int i=1;i<n;i++)cin>>c[i];
    for(int i=1;i<n;i++)cin>>l[i];
    dep[n]=1;
    for(int i=n-1;i>=1;i--)
    {
        fa[0][i]=c[i];
        dep[i]=dep[c[i]]+1;
        dis[i]=dis[c[i]]+l[i];
        for(int j=1;j<K;j++)fa[j][i]=fa[j-1][fa[j-1][i]];
    }
    for(int i=1;i<=n;i++)
    {
        f[i]=i;
        siz[i]=1;
        d[i]={i,i,0};
    }
    for(int i=1;i<n;i++)p[i]={max(a[i],a[c[i]]),i};
    sort(p+1,p+n);
    mx={0,0,-1};
    ll ans=-inf;
    node res={0,0,0};
    for(int i=1;i<n;)
    {
        int j=i;
        while(j<n&&p[j].first==p[i].first)
        {
            int x=p[j].second;
            join(x,c[x],l[x]);
            j++;
        }
        if(mx.len-p[i].first>ans)
        {
            ans=mx.len-p[i].first;
            res=mx;
        }
        i=j;
    }
    cout<<res.x<<' '<<res.y<<'\n';
    return 0;
}