题解:P17138 [KOI 2026 #1] 步道
lailai0916 · · 题解
题意简述
给定一棵点带拥挤度、边带长度的树。选择两个不同结点,使路径长度减去路径上的最大点权尽量大,并输出一组最优端点。
解题思路
最大点权不便直接在线段之间合并,可以把它改写成一个逐渐增大的阈值。
对任意阈值
原问题的最优满意度满足:
先证明从左到右的不等式。任取一条路径
所以
再证明反向不等式。阈值
一条连接
将所有边按
不必单独处理只加入结点却没有加入边的阈值。这样的变化只会增加孤立点,不会产生端点不同的路径,也不会改变
接下来考虑如何在并查集中维护直径。每个连通块记录直径长度及两个端点。设两个连通块通过边
- 完全位于第一个连通块;
- 完全位于第二个连通块;
- 经过新加入的边。
对一棵树中的任意结点
在两条旧直径与这条跨块路径中取最长者,便得到合并后连通块的直径。并查集合并的同时维护全局最长直径及端点。
树上距离通过倍增 LCA 计算。题目给出的父结点满足
排序、并查集合并和距离查询的总时间复杂度为
正确性证明
阈值等式的两侧已经分别得到上界和可达下界,所以最大化满意度等价于最大化
合并两个连通块时,不经过新边的路径完全属于某个旧连通块,最长者就是两条旧直径之一。任何经过新边的路径都由第一侧到
因此,每次合并后记录的都是新连通块的真实直径,全局记录也是当前森林的真实最长直径。对所有边阈值计算
参考代码
#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;
}