CF294E Shaass the Great

CF294E Shaass the Great

题意:

有一棵 \(n\) 个点的树,从 \(n-1\) 条边中去掉一条边,再构建一条相同长度的边重新构成一棵树(去掉的边和构造的边可以相同),问新的树中任意两点之间距离的总和最小是多少。

\(n \leq 5000\)

分析:

设去掉的是长度为 \(w\) 的边 \(u\rightarrow v\),并在点 \(x,y\) 间连一条长度为 \(w\) 的边,则 \(ans=sz_u\cdot sz_v\cdot w+sum_u+sum_v+\min\limits_{i\in u所在树的节点集合}\{sum_i\}\cdot sz_v+\min\limits_{i\in v所在树的节点集合}\{sum_v\}\cdot sz_u\),其中 \(sz_i\) 表示去掉一条边后以 \(i\) 为根的子树大小, \(sum_i\) 表示相同情况下以 \(i\) 为根时,其它树上节点到它的距离总和。

对于 \(sz_u\) 和 \(sz_v\),我们可以分别做一次 dfs 来求;至于 \(sum_i\) 则可以通过 换根dp 来求。对于每棵树的最小值,可以在换根时自底向上更新。复杂度为 \(O(n)\)。

所以我们可以每次枚举一条边,每次在不经过这条边做 换根dp ,复杂度为 \(O(n^2)\),可以接受。

对于 换根dp 时的求值,设 \(d_x\) 为以 \(x\) 为根的子树中 \(x\) 与子树内节点的距离总和,\(f_x\) 为以 \(x\) 为根的树中 \(x\) 与其它节点的距离总和。假设 \(f_x\) 已经求出,且 \(y\in son_x\),边 \(x\rightarrow y\) 的值为 \(z\),则 \(f_y=d_y+(f_x-d_y)+[(sz_{root}-sz_y)-sz_y]\cdot z=f_x+(sz_{root}-2sz_y)z\)。

具体实现见代码。

Code:

#include <bits/stdc++.h>
#define int long long
using namespace std;

const int N = 5010;

int n, ans = 2e18, cnt, d[N], f[N], U[N], V[N], W[N], sz[N], fa[N], mn[N], head[N];

struct xcj{
	int to, nxt, value;
} e[N << 1];

inline int read(){
	int s = 0, w = 1;
	char ch = getchar();
	for (; ch < '0' || ch > '9'; w *= ch == '-' ? -1 : 1, ch = getchar());
	for (; ch >= '0' && ch <= '9'; s = s * 10 + ch - '0', ch = getchar());
	return s * w;
}

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

void dfs(int x){
	sz[x] = 1;
	d[x] = 0;
	for (int i = head[x]; i; i = e[i].nxt){
		int y = e[i].to, z = e[i].value;
		if (y == fa[x]) continue;
		fa[y] = x;
		dfs(y);
		sz[x] += sz[y];
		d[x] += d[y] + sz[y] * z;
	}
}

void dp(int x, int root){
	mn[x] = f[x];
	for (int i = head[x]; i; i = e[i].nxt){
		int y = e[i].to, z = e[i].value;
		if (y == fa[x]) continue;
		f[y] = f[x] + (sz[root] - 2 * sz[y]) * z;
		dp(y, root);
		mn[x] = min(mn[x], mn[y]);
	}
}

signed main(){
	n = read();
	for (int i = 1; i < n; ++i){
		U[i] = read();
		V[i] = read();
		W[i] = read();
		add(U[i], V[i], W[i]);
		add(V[i], U[i], W[i]);
	}
	for (int i = 1, u, v, w, res; i < n; ++i){
		memset(mn, 0x3f, sizeof(mn));
		u = U[i];
		v = V[i];
		w = W[i];
		fa[u] = v;
		fa[v] = u;
		dfs(u);
		f[u] = d[u];
		dp(u, u);
		dfs(v);
		f[v] = d[v];
		dp(v, v);
		res = 0;
		for (int j = 1; j <= n; ++j) res += f[j];
		ans = min(ans, res / 2 + sz[u] * sz[v] * w + mn[u] * sz[v] + mn[v] * sz[u]);
	}
	printf("%lld\n", ans);
	return 0;
}
posted @ 2022-08-22 20:41  leoair  阅读(21)  评论(0)    收藏  举报