题意: 一个 NN 个点的树,树上任意两个节点的距离都是 11 ,一开始从 11 号节点出发,经过所有边后回到 11 号节点,你可以新建 KK 条边 (1K2)(1\leq K\leq 2),使得总的访问路径的长度最小,且满足新建的边正好只经过一次,输出这个最小的长度。

很明显,树上 NN 个点,那么就有 N1N-1 条边,要想每个节点都访问完,根据搜索可知,总长度应为 2(n1)2(n-1),每条边访问一次,回溯一次。

如果 K=1K=1,那么我们直接找到树的直径,在直径的两端点连一条边,这样能减掉最多的距离,此时最小长度为 2(n1)L+12(n-1)-L+1LL 即为直径长度。

如果 K=2K=2,我们先在直径的两端点连一条边,肯定会构成环,如果连第 22 条边时所构成的环与第一个环没有相交,那么正好满足题目条件,如果相交,我们会发现相交的那条边至少都要访问 22 次才能把所有边访问完。对于这种情况,我们可以先求出直径,长度为 L1L_1,把直径上的所有边权值改为 1-1,然后再求一次直径,直径长度为 L2L_2,所以答案就是:

2(n1)(L11)(L21)=2nL1L2 2(n-1)-(L_1-1)-(L_2-1)=2n-L_1-L_2

这时如果 L2L_2 中包含 L1L_1 中的边,由于权值改为 1-1,减去 1-1 正好把重叠的部分加了回来,满足条件,时间复杂度为 O(n)O(n)

需要注意的是,由于第一次找直径完后需要改变边权,所以需要记录边,故第一次最好使用搜索找直径,而第二次由于边权改为了 1-1,包含负权,所以必须用树形 DP 找直径,然后这里有个小技巧,在保存直径的路径时,直接选择保存边的编号,这时可以直接通过异或查询反边,从而得到上一个点,这样可免去了回溯。

#include<cstdio>
#include<cstring>
#include<cstdlib>
#include<cmath>
#include<algorithm>
using namespace std;
int n,k;
int f[101010],d[101010];
int dis[101010],vis[101010];
int ans,rt;
int head[201010],cnt;
struct node
{
	int next;
	int to;
	int num;
}e[201010];
void add(int from,int to,int num)
{
	e[++cnt].next=head[from];
	e[cnt].to=to;
	e[cnt].num=num;
	head[from]=cnt;
}
void dfs(int u)
{
	vis[u]=1;
	for(int i=head[u];i;i=e[i].next)
	{
		int v=e[i].to;
		if(vis[v])
			continue;
		d[v]=i;//注意,d数组存储路径存储的是编号,这样即可方便找到边。
		dis[v]=dis[u]+e[i].num;
		if(dis[v]>ans)
		{
			ans=dis[v];
			rt=v;
		}
		dfs(v);
	}
	return ;
}
void dp(int u,int fa)
{
	for(int i=head[u];i;i=e[i].next)
	{
		int v=e[i].to;
		if(v==fa)
			continue;
		dp(v,u);
		ans=max(ans,f[u]+f[v]+e[i].num);
		f[u]=max(f[u],f[v]+e[i].num);
	}
}
int main()
{
	scanf("%d%d",&n,&k);
	cnt=1;
	for(int i=1;i<n;i++)
	{
		int a,b;
		scanf("%d%d",&a,&b);
		add(a,b,1);
		add(b,a,1);
	}
	if(k==1)
	{
		dp(1,0);
		printf("%d",2*n-ans-1);
		return 0;
	}
	else
	{
		dfs(1);
		int p=rt;
		ans=0;
		rt=0;
		memset(dis,0,sizeof(dis));
		memset(d,0,sizeof(d));
		memset(vis,0,sizeof(vis));
		dfs(p);
		int l1=ans;
		for(;d[rt];rt=e[d[rt]^1].to)//直接用边,找到对应边的反边的to,即上一个
			e[d[rt]].num=e[d[rt]^1].num=-1;
		ans=0;
		dp(1,0);
		int l2=ans;
		printf("%d",2*n-l1-l2);
		return 0;
	}
	return 0;
}