题解:P16163 [ICPC 2016 NAIPC] YATP
lailai0916 · · 题解
题意简述
给定一棵边带权的树,节点
分别求每个起点的最小路径费用,再将这些最小值相加。
解题思路
对每个节点
乘积项与两个端点有关,不能仅保留子树中点权最小的节点:点权更小的节点也可能距离很远。我们用点分治处理距离,再用直线下凸壳处理乘积。
在当前连通块中取重心
固定起点
这里可以将当前连通块的所有节点一起加入直线集合,不必排除与
同一分支的节点只会给出不小于真实费用的候选值,不可能将答案错误地压低。随后递归处理每个分支,真正经过更深层重心的路径会补上准确费用。
完整地说,对于固定的最优端点对
将当前连通块中的节点按
考虑三条斜率严格递减的直线
两个斜率差都是正数,将交点比较交叉相乘,得到删除条件:
用栈维护上述条件即可建出下凸壳。查询时反向遍历排序后的节点,查询横坐标
代码中 q 保存凸壳对应的节点编号,calc(u,v) 计算
每次先迭代遍历当前连通块,逆序累加子树大小,找到删除后所有部分大小都不超过一半的节点。然后从该重心重新迭代遍历,计算距离。树上的这两次遍历都不用递归,因此长链不会造成深度为
边权为正,long long。不需要用浮点数求交点。
点分治每层节点总数为
参考代码
#include <bits/stdc++.h>
using namespace std;
using ll=long long;
using pii=pair<int,int>;
const int N=200005;
vector<pii> G[N];
int p[N],a[N],q[N],fa[N],siz[N];
ll dis[N],ans[N];
bool vis[N];
bool cmp(int u,int v){return p[u]!=p[v]?p[u]>p[v]:dis[u]<dis[v];}
ll calc(int u,int v)
{
return 1LL*p[u]*p[v]+dis[v];
}
void solve(int x)
{
a[1]=x;
fa[x]=0;
int tot=1;
for(int i=1;i<=tot;i++)
{
int u=a[i];
siz[u]=1;
for(auto [v,w]:G[u])
{
if(v==fa[u]||vis[v])continue;
fa[v]=u;
tot++;
a[tot]=v;
}
}
int rt=x;
for(int i=tot;i>=1;i--)
{
int u=a[i],mx=tot-siz[u];
for(auto [v,w]:G[u])
{
if(fa[v]==u&&!vis[v])mx=max(mx,siz[v]);
}
if(mx<=tot/2)rt=u;
if(fa[u])siz[fa[u]]+=siz[u];
}
a[1]=rt;
fa[rt]=0;
dis[rt]=0;
int cnt=1;
for(int i=1;i<=cnt;i++)
{
int u=a[i];
for(auto [v,w]:G[u])
{
if(v==fa[u]||vis[v])continue;
fa[v]=u;
dis[v]=dis[u]+w;
cnt++;
a[cnt]=v;
}
}
sort(a+1,a+tot+1,cmp);
int tail=0;
for(int i=1;i<=tot;i++)
{
int u=a[i];
if(tail&&p[q[tail]]==p[u])continue;
while(tail>1)
{
int v=q[tail-1],w=q[tail];
if((dis[w]-dis[v])*(p[w]-p[u])<(dis[u]-dis[w])*(p[v]-p[w]))break;
tail--;
}
tail++;
q[tail]=u;
}
int head=1;
for(int i=tot;i>=1;i--)
{
int u=a[i];
while(head<tail&&calc(u,q[head+1])<=calc(u,q[head]))head++;
ans[u]=min(ans[u],dis[u]+calc(u,q[head]));
}
vis[rt]=1;
for(auto [v,w]:G[rt])if(!vis[v])solve(v);
}
int main()
{
ios::sync_with_stdio(false);
cin.tie(nullptr);
int n;
cin>>n;
for(int i=1;i<=n;i++)
{
cin>>p[i];
ans[i]=1LL*p[i]*p[i];
}
for(int i=1;i<n;i++)
{
int u,v,w;
cin>>u>>v>>w;
G[u].push_back({v,w});
G[v].push_back({u,w});
}
solve(1);
ll res=0;
for(int i=1;i<=n;i++)res+=ans[i];
cout<<res<<'\n';
return 0;
}