题解:P16163 [ICPC 2016 NAIPC] YATP

· · 题解

题意简述

给定一棵边带权的树,节点 u 有点权 p_u。从 uv 的简单路径费用为两点距离加 p_up_v,允许 u=v

分别求每个起点的最小路径费用,再将这些最小值相加。

解题思路

对每个节点 u,要求的是:

\min_v\left(\operatorname{dist}(u,v)+p_up_v\right)

乘积项与两个端点有关,不能仅保留子树中点权最小的节点:点权更小的节点也可能距离很远。我们用点分治处理距离,再用直线下凸壳处理乘积。

在当前连通块中取重心 c,记 d_u=\operatorname{dist}(u,c)。若 u,v 位于删除 c 后的不同连通块中,或者其中一个点就是 c,它们的路径经过 c,费用可以写成:

d_u+d_v+p_up_v

固定起点 u 后,d_u 是常数。每个候选终点 v 提供一条直线 y=p_vx+d_v,在 x=p_u 处查询最小值,再加 d_u 即可。

这里可以将当前连通块的所有节点一起加入直线集合,不必排除与 u 同一分支的节点。原因是树上距离满足三角不等式:

d_u+d_v\ge\operatorname{dist}(u,v)

同一分支的节点只会给出不小于真实费用的候选值,不可能将答案错误地压低。随后递归处理每个分支,真正经过更深层重心的路径会补上准确费用。

完整地说,对于固定的最优端点对 (u,v),在点分树上第一次把它们分开的重心处,两者之间的路径一定经过该重心,因此这一次得到的候选值等于真实最优值;其他层的候选值则都不小于真实最优值。对所有层取最小值,恰好得到答案。u=v 的候选费用为 p_u^2,可以直接用于初始化。

将当前连通块中的节点按 p_v 递减排序,相同点权再按 d_v 递增排序。相同斜率只保留截距最小的一条直线。

考虑三条斜率严格递减的直线 A,B,C,分别写成 y=k_Ax+b_A 等形式。如果 A,B 的交点横坐标不小于 B,C 的交点横坐标,那么 B 不可能单独成为最优直线,应当删除。因为交点两侧分别有 AC 不劣于它。

两个斜率差都是正数,将交点比较交叉相乘,得到删除条件:

(b_B-b_A)(k_B-k_C)\ge(b_C-b_B)(k_A-k_B)

用栈维护上述条件即可建出下凸壳。查询时反向遍历排序后的节点,查询横坐标 p_u 递增。最优直线在凸壳上的位置也单调向后,因此只需一个指针:下一条直线在当前横坐标处不劣时,就向后移动。

代码中 q 保存凸壳对应的节点编号,calc(u,v) 计算 p_up_v+d_v。凸壳的建立和全部查询都只需线性时间,当前层的主要开销是排序。

每次先迭代遍历当前连通块,逆序累加子树大小,找到删除后所有部分大小都不超过一半的节点。然后从该重心重新迭代遍历,计算距离。树上的这两次遍历都不用递归,因此长链不会造成深度为 n 的调用栈;只有点分治过程递归,深度为 O(\log n)

边权为正,d_v<2\times10^{11},斜率不超过 10^6。凸壳交叉相乘的绝对值小于 2\times10^{17},最终答案不超过 n\times10^{12}\le2\times10^{17},均可用 long long。不需要用浮点数求交点。

点分治每层节点总数为 O(n),排序至多花费 O(n\log n),共有 O(\log n) 层。总时间复杂度为 O(n\log^2 n),空间复杂度为 O(n)

参考代码

#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;
}