O(n)代码求助
查看原帖
O(n)代码求助
701221
Chr0n1CleC楼主2023/7/20 16:25

貌似挺难调的。

错误情况:

1.遇到非链的树答案会错误。

2.代码答案 <= 实际答案


#include<cstdio>
#include<vector>
#include<algorithm>
using namespace std;

const int N = 1000009;

struct node
{
	int v, nxt;
	long long w;
}e[N << 1];

int head[N], cnt = 1;

void add(int u, int v, long long w)
{
	e[++ cnt].v = v, e[cnt].w = w, e[cnt].nxt = head[u], head[u] = cnt;
}

int s, t;

long long mx;

long long f[N], h[N], f1[N];

void dfs1(int u, int fa, long long dep)
{
	if (dep > mx)
		s = u, mx = dep;
	for (int i = head[u];i;i = e[i].nxt)
	{
		int v = e[i].v;
		long long w = e[i].w;
		if (v == fa)
			continue;
		dfs1(v, u, dep + w);
	}
}

int from[N], to[N];

void dfs2(int u, int fa)
{
	if (h[u] > mx)
		t = u, mx = h[u];
	for (int i = head[u];i;i = e[i].nxt)
	{
		int v = e[i].v;
		long long w = e[i].w;
		if (v == fa)
			continue;
		h[v] = h[u] + w;
		from[v] = i;
		dfs2(v, u);
		if (f[v] + w > f[u])
			f1[u] = f[u], f[u] = f[v] + w;
		else if (f[v] + w > f1[u])
			f1[u] = f[v] + w;	
	}
}

long long max(long long a, long long b)
{
	return a > b ? a : b;
}

long long min(long long a, long long b)
{
	return a < b ? a : b;
}

long long pre[N], lst[N]; 

int main()
{
	int n;
	scanf("%d", &n);
	for (int i = 1;i < n;++ i)
	{
		int u, v;
		long long w;
		scanf("%d%d%lld", &u, &v, &w);
		add(u, v, w), add(v, u, w);
	}
	dfs1(1, 0, 0);
	mx = 0;
	dfs2(s, 0);
	int cur = t;
	while (cur != s)
	{
		int v = e[from[cur] ^ 1].v;
		to[v] = from[cur];
		cur = v;
	}
	cur = s;
	int j = s;
	int count = 0;
	while (1)
	{
		long long val = h[cur] + f1[cur];
		while (j != cur)
		{
			int v = e[to[j]].v;
			if (max(h[v], val - h[v]) > max(h[j], val - h[j]))
				break;
			j = v;
		}
		++ count;
		if (j != s)
			pre[count] = max(h[j], val - h[j]);
		if (cur == t)
			break;
		cur = e[to[cur]].v;
	}
	int count1 = 0;
	cur = t, j = t;
	while (1)
	{
		long long val = f[cur] + f1[cur];
		while (j != cur)
		{
			int v = e[from[j] ^ 1].v;
			if (max(f[v], val - f[v]) > max(f[j], val - f[j]))
				break;
			j = v;
		}
		if (j != t)
			lst[count - count1] = max(f[j], val - f[j]);
		++ count1;
		if (cur == s)
			break;
		cur = e[from[cur] ^ 1].v;
	}
	cur = s;
	long long ans = f[s];
	int now = 1;
	while (now < count)
	{
		int v = e[to[cur]].v;
		ans = min(ans, max(pre[now] + lst[now + 1] + e[to[cur]].w, max(h[cur] + f1[cur], f[v] + f1[v])));
		++ now;
		cur = v;
	}
	printf("%lld\n", ans);
	
	return 0;
}

/*
5
1 2 1
2 3 3
2 4 4
1 5 5

7
1 2 2
1 3 3
2 5 5
2 6 6
3 4 4
3 7 7

*/
2023/7/20 16:25
加载中...