洛谷P4149 [IOI 2011] Race 题解 树上启发式合并(dsu on tree)
题目链接:https://www.luogu.com.cn/problem/P4149
题目大意
给定一棵 \(n\) 个节点的树,边有非负权值。求树上一条简单路径,使其长度恰好为 \(K\),并最小化路径上的边数。若不存在输出 \(-1\)。
算法:DSU on tree(树上启发式合并)
用 dis[u] 表示根节点 \(1\) 到 \(u\) 的路径长度,dep[u] 表示根到 \(u\) 的边数(深度)。
对于一条路径 \(x \to y\),设其 LCA 为 \(u\),则路径长度满足:
\[dis[x] + dis[y] - 2 \cdot dis[u] = K
\]
路径边数为:
\[(dep[x] - dep[u]) + (dep[y] - dep[u]) = dep[x] + dep[y] - 2 \cdot dep[u]
\]
我们使用 map<ll, int> mp 维护已合并子树中每个 dis 值对应的 最小深度 dep。
在 DSU on tree 过程中,对于当前节点 \(u\):
- 先递归处理轻儿子,不保留其贡献。
- 递归处理重儿子,保留其贡献。
- 遍历每个轻儿子子树,对于子树中的每个节点 \(v\):
- 查询是否存在已合并的节点 \(y\),使得
\[dis[y] = K + 2 \cdot dis[u] - dis[v]
\]
- 若存在,则路径边数为
\[(dep[y] - dep[u]) + (dep[v] - dep[u])
\]
更新全局最小值 ans。
- 查询完后,将轻儿子子树的所有节点合并到
mp中(仅保留相同dis的最小dep)。
- 处理以 \(u\) 为端点的路径:若
mp中存在dis[u] + K,则路径边数为dep[y] - dep[u],更新ans。 - 将 \(u\) 自身加入
mp。 - 若当前子树不需要保留(轻儿子),清空
mp。
剪枝:由于边权非负,dis 随深度递增。在查询时,若目标 disy <= dis[u],说明目标节点不可能在已合并的子树中(那些节点的 dis 都大于 dis[u]),可直接返回。
复杂度
每个节点最多被加入 mp \(O(\log n)\) 次,每次操作 map 为 \(O(\log n)\),总时间复杂度 \(O(n \log^2 n)\),空间 \(O(n)\)。
参考代码核心
void dfs_add(int u, int p, int x, int cnt, int id) {
if (id == 1) { // 查询答案
ll disy = K + 2 * dis[x] - dis[u];
if (disy <= dis[x]) return; // 剪枝
if (mp.find(disy) != mp.end()) {
int tmp = mp[disy] - dep[x] + cnt;
ans = (ans == -1 || ans > tmp) ? tmp : ans;
}
} else { // 合并到 mp
if (mp.find(dis[u]) == mp.end() || mp[dis[u]] > dep[u])
mp[dis[u]] = dep[u];
}
for (auto [v, w] : g[u])
if (v != p) dfs_add(v, u, x, cnt + 1, id);
}
在 dfs2 中按 DSU on tree 框架调用即可。
完整程序:
#include <bits/stdc++.h>
using namespace std;
using ll = long long;
const int maxn = 2e5 + 5;
struct Edge {
int v, w;
};
vector <Edge> g[maxn];
ll K, dis[maxn];
int n, sz[maxn], dep[maxn], son[maxn], ans = -1;
map<ll, int> mp; // key: 子树里所有结点的dis, val: 这个dis对应的边数
void dfs1(int u, int p, int _dep, ll _dis) {
sz[u] = 1;
dep[u] = _dep; // dep[u] 深度
dis[u] = _dis; // dis[u] 表示的是从根节点1到结点u的简单路径长度
for (auto [v, w] : g[u]) {
if (v != p) {
dfs1(v, u, _dep+1, _dis+w);
sz[u] += sz[v];
if (sz[v] > sz[ son[u] ])
son[u] = v;
}
}
}
void dfs_add(int u, int p, int x, int cnt, int id) {
if (id == 1) {
ll disy = K + 2 * dis[x] - dis[u];
if (disy <= dis[x])
return;
if (mp.find(disy) != mp.end()) {
int tmp = mp[disy] - dep[x] + cnt;
if (ans == -1 || ans > tmp)
ans = tmp;
}
}
else { // id == 2
if (mp.find(dis[u]) == mp.end() || mp[ dis[u] ] > dep[u])
mp[ dis[u] ] = dep[u];
}
for (auto [v, w] : g[u])
if (v != p)
dfs_add(v, u, x, cnt+1, id);
}
void dfs2(int u, int p, bool keep) {
for (auto [v, w] : g[u])
if (v != p && v != son[u])
dfs2(v, u, false);
if (son[u])
dfs2(son[u], u, true);
for (auto [v, w] : g[u]) {
if (v != p && v != son[u]) {
dfs_add(v, u, u, 1, 1); // 1: 更新答案
dfs_add(v, u, u, 1, 2); // 2: 将这棵子树合并到“前面的子树”
}
}
// 这里还得考虑以u为端点的路径
if (mp.find(dis[u] + K) != mp.end()) {
int cnt = mp[ dis[u] + K ] - dep[u];
if (ans == -1 || ans > cnt)
ans = cnt;
}
if (mp.find(dis[u]) == mp.end() || mp[ dis[u] ] > dep[u])
mp[ dis[u] ] = dep[u];
if (!keep) {
mp.clear();
}
}
int main() {
cin >> n >> K;
for (int i = 1, u, v, w; i < n; i++) {
cin >> u >> v >> w;
u++, v++;
g[u].push_back({v, w});
g[v].push_back({u, w});
}
dfs1(1, 0, 0, 0);
dfs2(1, 0, true);
cout << ans;
return 0;
}
浙公网安备 33010602011771号