洛谷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\):

  1. 先递归处理轻儿子,不保留其贡献。
  2. 递归处理重儿子,保留其贡献。
  3. 遍历每个轻儿子子树,对于子树中的每个节点 \(v\):
    • 查询是否存在已合并的节点 \(y\),使得

\[dis[y] = K + 2 \cdot dis[u] - dis[v] \]

  • 若存在,则路径边数为

\[(dep[y] - dep[u]) + (dep[v] - dep[u]) \]

更新全局最小值 ans。

  • 查询完后,将轻儿子子树的所有节点合并到 mp 中(仅保留相同 dis 的最小 dep)。
  1. 处理以 \(u\) 为端点的路径:若 mp 中存在 dis[u] + K,则路径边数为 dep[y] - dep[u],更新 ans。
  2. 将 \(u\) 自身加入 mp。
  3. 若当前子树不需要保留(轻儿子),清空 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;
}
posted @ 2026-09-12 16:19  quanjun  阅读(17)  评论(0)    收藏  举报