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;
}
浙公网安备 33010602011771号