虚树
原理
虚树是一种特殊的树形数据结构,主要用于解决一类特殊的树上的动态规划或查询问题。
它的核心思想是:在一棵庞大的树中,只保留“关键点”及其两两之间的最近公共祖先,将这些点按照原树的祖先关系连接成一棵新的树。这棵树保留了原树中关键节点之间的相对结构,但节点数量被压缩到了 \(O(k)\) 级别(其中 \(k\) 是关键点数量),从而大幅降低时间复杂度。
为什么要用虚树?
在算法竞赛或复杂系统设计中,我们常遇到如下问题:
给出一棵包含 \(n\) 个节点的树(\(n\) 可能很大,如 \(10^5\))。
有 \(q\) 次询问,每次询问给出若干个关键点(总关键点数量 \(\sum k\) 与 \(n\) 同阶)。
每次询问需要只针对这些关键点(或者它们之间的关系)做树形 DP,比如求最短路、最小割、连通块代价等。
每次询问根据给定的 \(k\) 个关键点构建一棵大小为 \(O(k)\) 的虚树,然后在虚树上跑 DP。复杂度为 \(O(\sum k \log k)\) 或 \(O(\sum k)\)。
举个例子:蓝色的是询问点。红色点就会在虚树上。


图来自:https://www.cnblogs.com/zzqsblog/p/5560645.html 。
构建过程(单调栈法)
- 我们先做一次
DFS,求出每个点的dfn和dep,以及预处理使其具备LCA查询能力。 - 然后将给定的关键点数组 \(a\) 按
dfn排序。 - 遍历排序后的关键点,计算相邻关键点的
LCA,将这些LCA也加入数组。
我们来稍微证明一下为什么只加 相邻关键点的 LCA 就可以覆盖任意两点的 LCA:
对于任意两个节点 \(u\) 和 \(v\)(假设 \(dfn[u] < dfn[v]\)),它们要么是祖先-后代关系(\(u\) 是 \(v\) 的祖先),要么分属不同的子树。
设三个节点按 DFS 序排列为 \(a, b, c\),即 \(dfn[a] < dfn[b] < dfn[c]\)。
令 \(z = LCA(a, c)\)。我们需要证明:
因为 \(z\) 是 \(a\) 和 \(c\) 的 LCA,所以 \(z\) 的子树在 DFS 序上是一个连续的区间,并且这个区间同时包含了 \(a\) 和 \(c\)。
由于 \(dfn[a] < dfn[b] < dfn[c]\),
\(b\) 位于 \(a\) 和 \(c\) 之间,因此 \(b\) 一定也落在这个区间内。所以:
\(b\) 也在 \(z\) 的子树内。
在 \(z\) 的子树中,
\(a\) 和 \(c\) 必然位于 \(z\) 的不同孩子分支。
现在看 \(b\) 的位置,它只有三种可能:
情况 1:
\(b\) 在 \(a\) 所在的那个分支里
此时 \(b\) 和 \(a\) 在同一分支,它们的 LCA 位于这个分支内,记作 \(x = LCA(a, b)\)。
因为 \(x\) 在 \(a\) 的分支内,而 \(z\) 是这个分支的祖先,所以 \(z\) 是 \(x\) 的祖先。
同时,
\(b\) 在 \(a\) 分支,\(c\) 在另一个分支,所以 \(b\) 和 \(c\) 的 LCA 就是 \(z\),即 \(LCA(b, c) = z\)。
于是:
因为 \(z\) 是 \(x\) 的祖先,所以它们的 LCA 就是 \(z\)。成立。
情况 2:
\(b\) 在 \(c\) 所在的那个分支里
完全对称:
\(LCA(b, c)\) 是 \(z\) 的后代,而 \(LCA(a, b) = z\),所以结论同样成立。
情况 3:
\(b\) 在 \(z\) 的第三个分支里(既不在 \(a\) 分支,也不在 \(c\) 分支)
那么 \(a\) 和 \(b\) 在 \(z\) 的不同分支,所以 \(LCA(a, b) = z\)。
同理 \(b\) 和 \(c\) 也在 \(z\) 的不同分支,所以 \(LCA(b, c) = z\)。
于是 \(LCA(z, z) = z\),显然成立。
- 如果我们发现栈顶元素 \(top\) 与我们想要加入的元素 \(x\) 的
LCA为 \(top\),那说明 \(x\) 的父亲为 \(top\),与 \(top\) 在一条链上,我们直接把 \(x\) 加入栈。否则我们一直退栈(在退栈过程中顺便连一下边),一直推到 \(LCA(top,x)=top\)(也就是栈顶元素是 \(x\) 的父亲,他们在一条链上)或者栈里仅剩下一个元素为止,然后再入栈即可。最后别忘了把还没空的栈给清空。
for(int i=1;i<=m;i++) cin>>key[i];
sort(key+1,key+1+m,cmp);
int len=m;
for(int i=1;i<m;i++) key[++len]=lca(key[i],key[i+1]);
sort(key+1,key+1+len,cmp);
len=unique(key+1,key+1+len)-key-1;
tp=0;
st[++tp]=1;
for(int i=2;i<=len;i++) {
while(lca(key[st[tp]],key[i])!=key[st[tp]]) tp--;
to[st[tp]].push_back(i),st[++tp]=i;
}
本文来自博客园,作者:MistyPost,转载请注明原文链接。

浙公网安备 33010602011771号