零、写在前面
树链剖分——重链剖分,这个主要掌握 Dfs 序,以及HLD求LCA
最近公共祖先,这个主要看如何将LCA转化为欧拉序列上的RMQ问题
0.1 LCA问题转化为RMQ问题
对一棵树进行 DFS,无论是第一次访问还是回溯,每次到达一个结点时都将编号记录下来,可以得到一个长度为 2n − 1 的序列,这个序列被称作这棵树的欧拉序列。
对于u, v,从u 走到 v 的过程一定会经历 LCA(u, v),并且 LCA(u, v) 是这段路径上深度最小的节点。
换言之,LCA(u, v) 是 u 到 v 欧拉序列上深度最小的节点。
我们将 LCA 转换为了区间最值问题,而区间最值问题是 RMQ 问题。
0.2 关键点
有一颗大小为 n 的树,每次选取树中的k个节点,记为集合 S,k 个节点中两两之间的LCA 构成的集合记为T,那么 |S ∪ T| <= 2k - 1
我们称 S ∪ T 为关键点。
证明:
假定选取的k个节点按dfs序排序后为:{u1, u2, …, ukl}
构造集合T’ = { lca(u1, u2), lca(u2, u3), …, lca(uk−1, uk) }
显然|T‘| <= k - 1
由于任意两个节点 ui, uj(i<j),它们的 lca(ui,uj) 一定也是区间 ui,ui+1,…,uj 中某一对相邻节点的 LCA 的祖先(由 0.1 易得),或者就是其中某一对的 LCA,所以T = T’。
所以|S ∪ T| = |S ∪ T‘| <= |S| + |T’| = 2k - 1
我们能否将关键点建立成一棵树呢?
一、虚树
1.1 虚树
将0.2 中的关键点,按照原树中的祖先关系连边,就可以得到一颗相较于原树更小的树。
通常情况下,我们会将根节点(一般是0)也加入,我们称这棵树为虚树(Virtual Tree)。
1.2 为什么要有虚树
一个经典的问题就是对于一颗n个节点的树,我们做q次查询,每次给定k个点,问最少删除多少个点可以使得k个点两两不连通。
根据虚树的定义,我们发现删除虚树外的点对于查询点连通性是没有影响的。
也就是说,我们对于每次查询,可以在更小的虚树上处理,从而达到更优的复杂度。
1.3 虚树的建立
虚树建立需要dfs序处理、某种快速lca查询方法、单调栈
由于重链剖分(HLD)本身就需要处理dfs序,所以我通常采用HLD求LCA来辅助
-
先预处理原树的 lca 和 dfs 序,然后用栈维护关键点,构造虚树,
-
对 k个查询点按 dfs 序排序,先加入根节点,再按顺序加入查询点。
-
栈维护从根向下的一条链上的关键点,按深度从小到大存储。
- 当加入 a[x]后,满足 stk[0] = root,stk.Peek() = a[x]。stk[i] 为 stk[i - 1]的后代
-
现在考虑加入查询点 a[i],设lca = LCA(s[top]。a[i]),分两类讨论
-
1)lca = stk.Peek(),即 a[i]是 stk.Peek()子树内的节点。直接把 a[i]入栈。
-
2)lca ≠ stk.Peek(),即 a[i]不是 stk.Peek()子树内的节点(如图所示)

- 那么我们将lca下方的点都出栈,出栈的时候将stk[top - 1] 和 stk[top] 连边
- 最后一次出栈的时候要将lca和栈顶连边
- 最后把lca和a[i]入栈
-
遍历结束后,把栈内维护的链连边
-
建树代码(C#)
- In[] 为 dfs序数组
Array.Sort(a, (i, j) => In[i].CompareTo(In[j]));
List<int> stk = [0];
if (a[0] > 0) stk.Add(a[0]);
for (int i = 1; i < k; ++i) {
int t = Lca(stk[^1], a[i]);
while (stk.Count > 1 && dep[stk[^2]] >= dep[t]) {
adj[stk[^2]].Add(stk[^1]);
stk.RemoveAt(stk.Count - 1);
}
if (t != stk[^1]) {
adj[t].Add(stk[^1]);
stk.RemoveAt(stk.Count - 1);
stk.Add(t);
}
stk.Add(a[i]);
}
while (stk.Count > 1) {
adj[stk[^2]].Add(stk[^1]);
stk.RemoveAt(stk.Count - 1);
}
二、题目练习
2.1 D. Kingdom and its Cities
原题链接
思路分析
对于查询点赋予点权f = 1,其他点默认点权f = 0
先考虑不可行的情况:存在f[u] > 0 并且 f[parent[u]] > 0
对于每组查询,我们先考虑暴力做法。
对整颗树做树形dp(注意此时是不存在f[u] > 0 并且 f[parent[u]] > 0的非法情况):
- 自顶向下递归,当前节点u,对子树递归
- 如果f[u] > 0
- 递归完子树v,如果 f[v] > 0,我们就删去v
- 如果f[u] = 0
- 递归完子树v,累加f[v] 到 f[u]
- 遍历完所有子树,如果 f[u] > 0,则删去u
- 如果f[u] > 0
单次树形dp时间复杂度:O(N)
总体时间复杂度:O(qN)
考虑到删除掉的点量级是 O(K)的,而题目保证了 ΣK <= 1E5,因而想到虚树
我们每次对查询点建立虚树做树形dp即可。
时间复杂度:O(N + ΣKlogK)
AC代码(C#)
public void Solve() {
int n = br.ReadInt32();
List<int>[] adj = new List<int>[n];
for (int i = 0; i < n; ++i) adj[i] = new();
for (int i = 1; i < n; ++i) {
int u = br.ReadInt32() - 1, v = br.ReadInt32() - 1;
adj[u].Add(v);
adj[v].Add(u);
}
int[] In = new int[n], dep = new int[n], siz = new int[n], fa = new int[n], top = new int[n];
fa[0] = -1;
int cur = 0;
// HLD 预处理
Dfs1(0);
Dfs2(0);
foreach (var v in adj) {
v.Clear();
}
int[] f = new int[n];
int q = br.ReadInt32();
int ans = 0;
while (q-- > 0) {
int k = br.ReadInt32();
int[] a = new int[k];
for (int i = 0; i < k; ++i) {
a[i] = br.ReadInt32() - 1;
f[a[i]] = 1;
}
bool ok = true;
foreach (int x in a) {
if (f[x] > 0 && x != 0 && f[fa[x]] > 0) {
ok = false;
break;
}
}
if (!ok) {
bw.AppendLine(-1);
foreach (int x in a) {
f[x] = 0;
}
continue;
}
// 建立虚树
Array.Sort(a, (i, j) => In[i].CompareTo(In[j]));
List<int> stk = [0];
if (a[0] > 0) stk.Add(a[0]);
for (int i = 1; i < k; ++i) {
int t = Lca(stk[^1], a[i]);
while (stk.Count > 1 && dep[stk[^2]] >= dep[t]) {
adj[stk[^2]].Add(stk[^1]);
stk.RemoveAt(stk.Count - 1);
}
if (t != stk[^1]) {
adj[t].Add(stk[^1]);
stk.RemoveAt(stk.Count - 1);
stk.Add(t);
}
stk.Add(a[i]);
}
while (stk.Count > 1) {
adj[stk[^2]].Add(stk[^1]);
stk.RemoveAt(stk.Count - 1);
}
ans = 0;
Dfs(0);
f[0] = 0;
bw.AppendLine(ans);
}
void Dfs1 (int u) {
if (fa[u] != -1) {
adj[u].Remove(fa[u]);
}
siz[u] = 1;
for (int i = 0; i < adj[u].Count; ++i) {
int v = adj[u][i];
fa[v] = u;
dep[v] = dep[u] + 1;
Dfs1(v);
siz[u] += siz[v];
if (siz[v] > siz[adj[u][0]]) {
(adj[u][0], adj[u][i]) = (adj[u][i], adj[u][0]);
}
}
}
void Dfs2 (int u) {
In[u] = cur++;
foreach (int v in adj[u]) {
top[v] = v == adj[u][0] ? top[u] : v;
Dfs2(v);
}
}
int Lca(int u, int v) {
while (top[u] != top[v]) {
if (dep[top[u]] > dep[top[v]]) {
u = fa[top[u]];
} else {
v = fa[top[v]];
}
}
return dep[u] < dep[v] ? u : v;
}
void Dfs (int u) {
if (f[u] > 0) {
foreach (int v in adj[u]) {
Dfs(v);
if (f[v] > 0) {
++ans;
f[v] = 0;
}
}
} else {
foreach (int v in adj[u]) {
Dfs(v);
if (f[v] > 0) {
f[u] += f[v];
f[v] = 0;
}
}
if (f[u] > 1) {
++ans;
f[u] = 0;
}
}
adj[u].Clear();
}
}
2.2 3786. Total Sum of Interaction Cost in Tree Groups
原题链接
3786. Total Sum of Interaction Cost in Tree Groups
思路分析
考虑暴力做法:
假定由 U 组,那么我们进行 U 次 dfs,统计每条边在每组中的贡献
每条边的贡献即 该边两端的两个集合各自包含的组内点数之积
这样的做法时间复杂度是**O(nU)**的
考虑建立每组建立虚树,再同样的方法dfs求贡献,这样的复杂度就跟U无关了
时间复杂度:O(nlogn)
AC代码(C#)
using i64 = long;
public class Solution {
public long InteractionCosts(int n, int[][] edges, int[] group) {
List<int>[] adj = new List<int>[n];
for (int i = 0; i < n; ++i) adj[i] = new();
foreach (var e in edges) {
var (u, v) = (e[0], e[1]);
adj[u].Add(v);
adj[v].Add(u);
}
int[] In = new int[n], dep = new int[n], siz = new int[n], fa = new int[n], top = new int[n];
fa[0] = -1;
int cur = 0;
// HLD 预处理
Dfs1(0);
Dfs2(0);
for (int i = 0; i < n; ++i) {
adj[i].Clear();
siz[i] = 0;
}
bool[] f = new bool[n];
i64 ans = 0;
int tot = 0;
foreach (var g in Enumerable.Range(0, n).GroupBy(x => group[x])) {
var ver = g.OrderBy(x=>In[x]).ToArray();
List<int> stk = [0];
if (ver[0] > 0) {
stk.Add(ver[0]);
}
tot = ver.Length;
f[ver[0]] = true;
for (int i = 1; i < ver.Length; ++i) {
f[ver[i]] = true;
int lca = Lca(stk[^1], ver[i]);
while (stk.Count > 1 && dep[stk[^2]] >= dep[lca]) {
adj[stk[^2]].Add(stk[^1]);
stk.RemoveAt(stk.Count - 1);
}
if (lca != stk[^1]) {
adj[lca].Add(stk[^1]);
stk.RemoveAt(stk.Count - 1);
stk.Add(lca);
}
stk.Add(ver[i]);
}
while (stk.Count > 1) {
adj[stk[^2]].Add(stk[^1]);
stk.RemoveAt(stk.Count - 1);
}
Dfs(0);
siz[0] = 0;
}
return ans;
void Dfs1(int u) {
if (fa[u] != -1) {
adj[u].Remove(fa[u]);
}
siz[u] = 1;
for (int i = 0; i < adj[u].Count; ++i) {
int v = adj[u][i];
fa[v] = u;
dep[v] = dep[u] + 1;
Dfs1(v);
siz[u] += siz[v];
if (siz[v] > siz[adj[u][0]]) {
(adj[u][0], adj[u][i]) = (adj[u][i], adj[u][0]);
}
}
}
void Dfs2(int u) {
In[u] = cur++;
foreach (int v in adj[u]) {
top[v] = v == adj[u][0] ? top[u] : v;
Dfs2(v);
}
}
int Lca(int u, int v) {
while (top[u] != top[v]) {
if (dep[top[u]] > dep[top[v]]) {
u = fa[top[u]];
} else {
v = fa[top[v]];
}
}
return dep[u] < dep[v] ? u : v;
}
void Dfs(int u) {
siz[u] = f[u] ? 1 : 0;
foreach (var v in adj[u]) {
Dfs(v);
siz[u] += siz[v];
ans += 1L * (dep[v] - dep[u]) * siz[v] * (tot - siz[v]);
}
adj[u].Clear();
f[u] = false;
}
}
}
2.3 P2495 【模板】虚树 / [SDOI2011] 消耗战
原题链接
思路分析
注意到,一个查询点的祖先也是查询点的话,我们一定是删除祖先往上的某条边(因为删除下面的边不足以使得连通性断掉)
所以,对查询点建立虚树,但是不同的是,lca == stk.back() 的时候,我们continue掉
这样构造出来的虚树,查询点都是叶子
然后跑树形dp
剩下的查询点均在叶子上。对叶子节点直接返回权值。
对父节点累加子节点的权值和,与父节点的权值比较,取最小值返回
时间复杂度:O(nlogn)
AC代码(C++)
#include <bits/stdc++.h>
namespace views = std::views;
namespace ranges = std::ranges;
using i64 = long long;
using u64 = unsigned long long;
constexpr i64 inf = 1E18;
int main() {
std::ios::sync_with_stdio(false);
std::cin.tie(nullptr);
int n;
std::cin >> n;
std::vector<std::vector<std::array<int, 2>>> adj(n);
for (int i = 1; i < n; ++i) {
int u, v, w;
std::cin >> u >> v >> w;
--u; --v;
adj[u].push_back({v, w});
adj[v].push_back({u, w});
}
std::vector<int> in(n), dep(n), siz(n), fa(n), top(n);
std::vector<i64> min(n, inf);
int cur = 0;
fa[0] = -1;
[&](this auto && self, int u) -> void {
if (~fa[u]) {
adj[u].erase(ranges::find_if(adj[u], [&](auto & e) {
return e[0] == fa[u];
}));
}
siz[u] = 1;
for (auto &e : adj[u]) {
auto [v, w] = e;
dep[v] = dep[u] + 1;
fa[v] = u;
min[v] = std::min<i64>(min[u], w);
self(v);
siz[u] += siz[v];
if (siz[v] > siz[adj[u][0][0]]) {
std::swap(e, adj[u][0]);
}
}
}(0);
[&](this auto && self, int u) -> void {
in[u] = cur++;
for (auto &[v, w] : adj[u]) {
top[v] = v == adj[u][0][0] ? top[u] : v;
self(v);
}
}(0);
auto Lca = [&](int u, int v) -> int {
while (top[u] != top[v]) {
if (dep[top[u]] > dep[top[v]]) {
u = fa[top[u]];
} else {
v = fa[top[v]];
}
}
return dep[u] < dep[v] ? u : v;
};
std::vector<std::vector<int>> g(n);
int m;
std::cin >> m;
while (m-- > 0) {
int k;
std::cin >> k;
std::vector<int> a(k);
for (int i = 0; i < k; ++i) {
std::cin >> a[i];
--a[i];
}
ranges::sort(a, {}, [&](int v) {
return in[v];
});
std::vector<int> stk{0};
if (a[0] > 0) {
stk.push_back(a[0]);
}
for (int i = 1; i < k; ++i) {
int lca = Lca(a[i], stk.back());
if (lca == stk.back()) {
continue;
}
while (stk.size() > 1 && dep[stk[stk.size() - 2]] >= dep[lca]) {
g[stk[stk.size() - 2]].push_back(stk.back());
stk.pop_back();
}
if (lca != stk.back()) {
g[lca].push_back(stk.back());
stk.pop_back();
stk.push_back(lca);
}
stk.push_back(a[i]);
}
while (stk.size() > 1) {
g[stk[stk.size() - 2]].push_back(stk.back());
stk.pop_back();
}
std::cout << [&](this auto && self, int u) -> i64 {
if (g[u].empty()) {
return min[u];
}
i64 res = 0;
for (int v : g[u]) {
res += self(v);
}
g[u].clear();
return std::min(res, min[u]);
}(0) << '\n';
}
return 0;
}
2.4 F. Unique Occurrences
原题链接
思路分析
对于虚树问题,往往可以先考虑暴力做法:
建图后,我们考虑枚举颜色 col,假定以 0 为根
然后定义 dp[u, 0] 为 u 子树内,经过 u 不含 col 的路径数(注意本题单独一个节点u不算路径),dp[u, 1] 为 为 u 子树内,经过 u 含 col 的路径数(注意本题单独一个节点u不算路径)
那么枚举子节点v,如果 w(u, v) = col,那么:
ans += (dp[v][0] + 1) * dp[u][0]; // (u, v)为特殊边,那么v内只能挑dp[v, 0]
dp[u][1] += dp[v][0] + 1;
否则:
ans += (dp[v][0] + 1) * dp[u][1]; // v 内挑不带颜色,v外挑带颜色
ans += dp[v][1] * dp[u][0]; // v内挑带颜色,v外挑不带颜色
dp[u][0] += dp[v][0] + 1;
dp[u][1] += dp[v][1];
最后还要累加 u 作为路径端点的方案:
ans += dp[u][1];
那么如何优化呢?
我们按颜色将边分组,那么所有边涉及的端点集合就是我们的特殊点,建立虚树跑dp
值得注意的是,此时我们不能初始化 dp[u, 0] = 0,因为不包含在虚树中的点也会对 dp[u, 0] 贡献,那么如何初始化?
我们递归到 u时,初始化 dp[u, 0] = siz[u] - 1,然后再:
for (int v : adj[u]) dp[u][0] -= g.siz[v];
也就是将不在虚树内的点都给算上了,虚树上的点我们单独算。
时间复杂度:O(nlogn)
本题另解:
线段树分治:O(nlog^2n)
换根dp:O(n)
AC代码(cpp)
#include <bits/stdc++.h>
namespace ranges = std::ranges;
using i64 = long long;
struct HLD {
int n;
std::vector<int> siz, top, dep, parent, in, out, seq;
std::vector<std::vector<int>> adj;
int cur;
HLD() {}
HLD(int _n) {
init(_n);
}
void init(int _n) {
n = _n;
siz.resize(n);
top.resize(n);
dep.resize(n);
parent.resize(n);
in.resize(n);
out.resize(n);
seq.resize(n);
cur = 0;
adj.assign(n, {});
}
void addEdge(int u, int v) {
adj[u].push_back(v);
adj[v].push_back(u);
}
void work(int root = 0) {
top[root] = root;
dep[root] = 0;
parent[root] = -1;
dfs1(root);
dfs2(root);
}
void dfs1(int u) { // dep和sz
if (~parent[u]) //转有向图
adj[u].erase(std::find(adj[u].begin(), adj[u].end(), parent[u]));
siz[u] = 1;
for (int &v : adj[u]) {
parent[v] = u;
dep[v] = dep[u] + 1;
dfs1(v);
siz[u] += siz[v];
if (siz[v] > siz[adj[u][0]])
std::swap(v, adj[u][0]);
}
}
void dfs2(int u) { // dfn
in[u] = cur ++;
seq[in[u]] = u;
for (int v : adj[u]) {
top[v] = v == adj[u][0] ? top[u] : v;
dfs2(v);
}
out[u] = cur; // 欧拉序
}
int lca(int u, int v) const {
while (top[u] != top[v]) {
if (dep[top[u]] > dep[top[v]])
u = parent[top[u]];
else
v = parent[top[v]];
}
return dep[u] < dep[v] ? u : v;
}
int dist(int u, int v) const {
return dep[u] + dep[v] - dep[lca(u, v)] * 2;
}
int jump(int u, int k) { // 往上k层祖先
if (dep[u] < k) return -1;
int d = dep[u] - k;
while (dep[top[u]] > d)
u = parent[top[u]];
return seq[in[u] - dep[u] + d];
}
bool isAncester(int u, int v) const { // 括号序判断u 是否是 v 的祖先
return in[u] <= in[v] && in[v] < out[u];
}
int rootedParent(int u, int v) const { // 以v为p的u的祖先
std::swap(u, v);
if (u == v)
return u;
if (!isAncester(u, v))
return parent[u];
auto it = std::lower_bound(adj[u].begin(), adj[u].end(), v, [&](int x, int y) -> bool {
return in[x] < in[y];
}) ;
return * --it;
}
int rootedSize(int u, int v) {
if (u == v)
return n;
if (!isAncester(v, u))
return siz[v];
return n - siz[rootedParent(u, v)];
}
int rootedLca(int a, int b, int c) {
return lca(a, b) ^ lca(b, c) ^ lca(c, a);
}
};
int main() {
std::ios::sync_with_stdio(false);
std::cin.tie(nullptr);
auto get = [](i64 u, i64 v) -> i64 {
if (u > v) std::swap(u, v);
return u << 20 | v;
};
int n;
std::cin >> n;
std::vector<std::vector<int>> st(n);
std::vector<std::unordered_set<i64>> good(n);
HLD g(n);
for (int i = 1; i < n; ++i) {
int u, v, x;
std::cin >> u >> v >> x;
--u; --v; --x;
st[x].push_back(u);
st[x].push_back(v);
g.addEdge(u, v);
good[x].insert(get(u, v));
}
g.work();
std::vector<std::vector<int>> adj(n);
std::vector<std::array<i64, 2>> dp(n);
i64 ans = 0;
for (int val = 0; val < n; ++val) {
auto &ver = st[val];
if (ver.empty()) continue;
ranges::sort(ver);
ver.resize(std::unique(ver.begin(), ver.end()) - ver.begin());
ranges::sort(ver, {}, [&](int x) {
return g.in[x];
});
std::vector<int> stk{0};
if (ver[0] > 0) {
stk.push_back(ver[0]);
}
for (int i = 1; i < ver.size(); ++i) {
int lca = g.lca(ver[i], stk.back());
for (; stk.size() > 1 && g.dep[stk[stk.size() - 2]] >= g.dep[lca]; ) {
adj[stk[stk.size() - 2]].push_back(stk.back());
stk.pop_back();
}
if (stk.back() != lca) {
adj[lca].push_back(stk.back());
stk.pop_back();
stk.push_back(lca);
}
stk.push_back(ver[i]);
}
for (; stk.size() > 1; ) {
adj[stk[stk.size() - 2]].push_back(stk.back());
stk.pop_back();
}
stk.clear();
[&](this auto && self, int u) -> void {
dp[u][0] = g.siz[u] - 1;
dp[u][1] = 0;
for (int v : adj[u]) dp[u][0] -= g.siz[v];
for (int v : adj[u]) {
self(v);
if (good[val].contains(get(u, v))) {
ans += (dp[v][0] + 1) * dp[u][0];
dp[u][1] += dp[v][0] + 1;
} else {
ans += (dp[v][0] + 1) * dp[u][1];
ans += dp[v][1] * dp[u][0];
dp[u][0] += dp[v][0] + 1;
dp[u][1] += dp[v][1];
}
}
ans += dp[u][1];
adj[u].clear();
}(0);
}
std::cout << ans << '\n';
return 0;
}

说些什么吧!