题解 P3478 【[POI2008]STA-Station】

· · 题解

经典的树形dp,运用到了换根法,我详细讲一下状态转移方程的推导.

首先考虑暴力,跑n遍堆优化dijkstra,时间复杂度O(n^2*log(n)) 能得50分.

正解是树形dp.

此题用到了换根法,先明确几个东西.

n为总节点个数

size[x]表示x的子树的节点个数(包括自己)

f[x]表示以x为根节点时的其他节点到根节点距离和

对于任意一个非根节点x,他的父亲节点为y,状态转移方程如下:

f[x] = f[y] + (n-size[x]) - size[x];

如何理解这个状态转移方程呢?

1.对于不在x的子树中的节点,他们到x的距离等于他们到y的距离 +(n-size[x])

这个应该不难理解,每个不在x的子树中的节点,到x的距离比到y的距离多了1,一共

有(n-size[x])个节点不在x的子树中.

2对于在x的子树中的节点,他们到x的距离等于他们到y的距离-size[x]

这个也很好理解,每个在x的子树中的节点,到x的距离比到y的距离少了1,一共有

size[x]个节点在x的子树中.

合并,可得状态转移方程.

有了状态转移方程,代码就好写了.

1.首先选任意一个点当做根节点(我一般选1)递归求出每个子树的节点个数(包括

自己)顺便计算出其他节点到根节点的距离和.

2.再来一遍dfs进行状态转移.

不过还有一个细节要注意,在我的代码中,根节点的到根节点距离不能设为0,因为这

是个无向图,我在判断有没有往回跑的时候用的是判断v有没有被赋过值,如果设根节

点到根节点距离为0,那么会导致误判.

所以我设根节点到根节点的距离为1,其他节点到根节点的距离相应也多了1,最后减

去节点个数就行了.

别忘开long long!!!


#include<iostream>
#include<cstdio>
#include<algorithm>
#include<cstring>
using namespace std;
long long size[1000010],dis[1000010],f[1000010];//f数组表示以i为根节点

的距离和,dis数组只是计算每一个点到一开始假设的根的距离,之后换根的时候就没

用了
int head[1000010];
int idx,n;
long long num;
struct node{
    int nxt,to;
}edge[2000010];
void add(int u,int v)
{
    edge[++idx].nxt=head[u];
    edge[idx].to=v;
    head[u]=idx;
}
void dfs1(int now)//递归求出每个子树的节点个数顺便计算出其他节点到根节点的距离和 
{
    size[now]=1;
    for(int i=head[now];i;i=edge[i].nxt)
    {
        int v=edge[i].to;
        if(dis[v]) continue;
        dis[v]=dis[now]+1;
        dfs1(v);
        size[now]+=size[v];
    }
}
void dfs2(int now,int fath)
{
    if(now!=1)//如果是根节点的话不能进行状态转移
    f[now]=f[fath]-size[now]+n-size[now];
    for(int i=head[now];i;i=edge[i].nxt)
    {
        int v=edge[i].to;
        if(v==fath) continue;
        dfs2(v,now);
    }
}
int main()
{
    scanf("%d",&n);
    for(int i=1;i<=n-1;i++)
    {
        int u,v;
        scanf("%d%d",&u,&v);
        add(u,v);
        add(v,u);
    }
    dis[1]=1;
    dfs1(1);
    for(int i=1;i<=n;i++)
    {
        num+=d[i];
    }
    num-=n;
    f[1]=num;
    dfs2(1,0);
    long long ans=0,temp;
    for(int i=1;i<=n;i++)
    {
        if(f[i]>ans)
        {
            ans=f[i];
            temp=i;
        }
    }
    printf("%lld",temp);
    return 0;
}