CC++ & Algorithm

树链剖分(HLD)简介

极难3
语言版本:通用
概述:把树切成几条“重链”,让树上路径问题变成区间问题,用线段树等数据结构快速处理。

树链剖分(HLD)——把大树拆成“区间”,让路径查询变简单

从生活中的例子引入

想象你有一棵大树(物理的树),你想知道从树梢的一片叶子到另一片叶子,沿途所有树枝的粗细之和。如果每次都要从叶子爬到分叉点、再爬上去,实在太慢。有没有办法给树枝标号,让一条路径上的编号连续,这样就能用“区间和”来快速计算?

其实,在计算机里,我们经常要处理“树上路径”的问题。比如:

  • 一个班级的座位用树形结构表示(班主任是根,每个同学有一个直接上级),每次要统计某个组长到另一个组长的所有同学的总成绩。
  • 网络路由中,数据包从一台电脑到另一台电脑,沿途经过的路由器有哪些?怎么快速算总延迟?
  • 家族族谱中,从爷爷到孙子这一条线的人数总和。

如果每次都在树上慢慢爬,最坏情况下要走 O(n) 步。但树链剖分(HLD, Heavy-Light Decomposition)就像给树装上了“快速通道”——它把一棵树剖分成若干条链,每条链上的节点在DFS序中是连续的。这样,树上任意两点之间的路径可以被拆分成 O(log n) 条连续的链,每条链对应一段连续区间,我们就可以用线段树、树状数组等数据结构去维护和查询。原来需要 O(n) 的操作,现在只需要 O(log² n) 甚至 O(log n)!

树链剖分的核心思想是:把树“压扁”成一条直线,让路径变成几个连续的片段

为什么要用树链剖分?

直接在一棵树上做路径查询,比如求两个节点之间所有节点权值之和,最直接的方法是从一个节点出发,暴力走到另一个节点,沿途累加。但这样每次查询要 O(n) 时间,如果查询次数很多(比如上千次),就会非常慢。而树链剖分结合线段树,可以把每次查询优化到 O(log² n),即使树很大(比如10万个节点)也能飞快完成。

除了路径求和,树链剖分还能做路径修改(比如给路径上所有节点加一个值)、求路径上的最大值/最小值、甚至处理子树操作(因为子树内的节点在DFS序中也是连续的)。它像一把“瑞士军刀”,是处理树上问题的利器。

数据结构原理和核心思想

1. 基本概念

我们给树上的每个节点和边起一些特殊的名字:

  • 重儿子 (heavy child):一个节点的所有儿子中,子树大小(包括自己)最大的那个。如果有多个大小相等的,可以任选一个(比如选第一个)。
  • 轻儿子 (light child):除了重儿子以外的其他儿子。
  • 重边 (heavy edge):连接父节点和重儿子的边。
  • 轻边 (light edge):连接父节点和轻儿子的边。
  • 重链 (heavy chain):由重边连接而成的链。每条重链的起点是某个轻儿子或者根节点(它们自己作为链顶)。

举个例子:假设你管理一群同学,每个同学下面有某些下属。你想知道谁手下的人最多,就选那个手下人最多的作为“重儿子”。这样,从你开始,一路沿着重儿子往下,就形成了一条“重链”。其他轻儿子则各自开始自己的小链。

重要性质:从任意节点到根节点的路径上,经过的轻边数量不超过 O(log n)。为什么呢?因为每经过一条轻边,子树的大小至少会减半(因为轻儿子的大小小于父节点大小的一半?准确说是轻儿子的子树大小 ≤ 父节点子树大小的一半,证明:如果轻儿子大小超过一半,它就会成为重儿子)。所以经过轻边的次数最多是 log₂n 级别。这意味着路径最多被分成 O(log n) 条重链片段。

2. 剖分流程

通常用两次DFS(深度优先遍历)完成:

第一次DFS:从根节点开始,计算每个节点的:

  • fa[u]:父节点。
  • dep[u]:深度(根深度为1或者0)。
  • sz[u]:子树大小(包括自己)。
  • son[u]:重儿子(子树大小最大的儿子,如果没有儿子则为0)。

第二次DFS:再次从根开始,按照“先重儿子,后轻儿子”的顺序遍历,分配DFS序(dfn[u])和链顶标记 top[u]。这样做的目的是:保证每条重链上的节点在DFS序中是连续的,并且重链内部的顺序是沿着从上到下的。

具体做法:

  • 进入一个节点时,给它分配一个递增的编号 dfn[u]
  • 如果它有重儿子,先递归重儿子,并且让重儿子的链顶等于当前节点的链顶(同一链)。
  • 然后遍历所有轻儿子(非重儿子),每个轻儿子作为新链的起点(链顶是自己),递归。

3. 查询/修改路径

对于节点 u 和 v,要处理它们路径上的信息(比如求和)。过程如下:

  1. top[u] != top[v] 时,说明它们不在同一条重链上。这时,我们需要让“链顶深度较大的那个节点”向上跳,跳到其链顶的父亲,同时处理从该节点到链顶的这一段区间(对应 [dfn[top[u]], dfn[u]])。
  2. 重复步骤1,直到 top[u] == top[v],即它们在同一条重链上。
  3. 最后,处理这两个节点之间的区间(注意,深度小的节点编号小,所以区间是 [dfn[u], dfn[v]] 或 [dfn[v], dfn[u]])。

由于每次跳过一个链顶,而链顶的深度至少减少1(实际上跳跃会跨越多个节点),总共跳的次数为 O(log n)。每次区间操作如果是线段树,复杂度 O(log n),总复杂度 O(log² n)。

ASCII图示

假设一棵树(括号里是子树大小):

       1(7)
      / \
     2(3) 3(3)
    / \    \
   4(1)5(1) 6(2)
            /
           7(1)
  • 子树大小:节点2有3个(2、4、5),节点3有3个(3、6、7)。它们相等,我们任选节点2作为重儿子。所以重边:1-2, 2-4? 节点4没有儿子,所以重儿子是0;节点5也是0;节点3的重儿子是节点6(子树2>0),节点6的重儿子是节点7。所以重链有:
    • 重链1:1-2-4(因为2是1的重儿子,4是2的重儿子,形成连续链)
    • 重链2:3-6-7(3是根的重儿子?不对,根是1,2才是重儿子,3是轻儿子,所以3作为链顶,6是3的重儿子,7是6的重儿子)
    • 重链3:5(轻儿子自己成链)
    • 重链4:1? 根也是链顶,但已经包含在链1中。

DFS序(先重后轻): 我们从根1开始,先走重儿子2:

  • dfn[1]=1
  • 到2:dfn[2]=2
  • 先走2的重儿子4:dfn[4]=3
  • 4没有子节点,返回
  • 然后走2的轻儿子5:dfn[5]=4
  • 返回2,回到1
  • 然后走1的轻儿子3:dfn[3]=5
  • 走3的重儿子6:dfn[6]=6
  • 走6的重儿子7:dfn[7]=7
  • 结束。

链顶:

  • top[1]=1(根)
  • top[2]=1(因为2是1的重儿子,链顶继承)
  • top[4]=1(链顶继承)
  • top[5]=5(轻儿子,自己为链顶)
  • top[3]=3(轻儿子自己为链顶)
  • top[6]=3(3的重儿子)
  • top[7]=3

现在查询路径4到7:

  • u=4, v=7。top[4]=1, top[7]=3,不相等。
  • 比较链顶深度:dep[1]=1, dep[3]=2? 假设根深度为1,那么dep[1]=1, dep[2]=2, dep[4]=3, dep[5]=3, dep[3]=2, dep[6]=3, dep[7]=4。链顶深度:top[4]=1深度1,top[7]=3深度2。深度大的链顶是节点3(因为dep[3]=2 > dep[1]=1),所以应该让v向上跳?注意代码中比较的是 dep[top[u]] 和 dep[top[v]],谁小谁在上面?我们需要让深度大的链顶先跳。实际上当top不同时,我们取链顶深度较大的节点跳,因为它的链顶更靠近下方。这里dep[top[7]]=dep[3]=2,dep[top[4]]=dep[1]=1,所以dep[top[7]] > dep[top[4]],说明v的链顶更深,让v跳。处理区间 [dfn[top[7]]=dfn[3]=5, dfn[7]=7],然后v = fa[top[7]] = fa[3]=1。
  • 现在u=4, v=1。top[4]=1, top[1]=1,相同。因为dep[4]=3 > dep[1]=1,所以u是深度大的,交换使得u在上方?实际上最后一步我们处理区间 [dfn[1], dfn[4]] = [1,3](因为深度小的节点编号小)。注意方向:路径是从4到1,但我们要求的是从u到v,所以最终区间 [dfn[u], dfn[v]] 其中u深度小。这里dep[1]=1 < dep[4]=3,所以u=1,v=4,区间[dfn[1]=1, dfn[4]=3]。
  • 这样路径被拆成两个连续区间:[5,7] 和 [1,3]。完美!

C++完整代码实现

下面的代码实现树链剖分,并用线段树维护区间和,支持路径求和与单点修改。变量名使用简短英文,每行添加中文注释。

#include <iostream>
#include <vector>
using namespace std;

const int MAXN = 100005; // 最大节点数

vector<int> adj[MAXN]; // 邻接表存树
int n, q;              // 节点数,查询次数

// 树剖数据结构
int fa[MAXN];    // 父节点
int dep[MAXN];   // 深度
int sz[MAXN];    // 子树大小
int son[MAXN];   // 重儿子(0表示无)
int top[MAXN];   // 链顶
int dfn[MAXN];   // dfs序(节点在数组中的下标)
int rnk[MAXN];   // 下标对应节点编号(反向映射)
int cnt = 0;     // dfs计数器

// 线段树
int tree[4*MAXN]; // 线段树数组
int arr[MAXN];    // 节点初始权值(下标从1开始)

// 第一次DFS:计算fa, dep, sz, son
void dfs1(int u, int f) {
    fa[u] = f;                 // 记录父节点
    dep[u] = dep[f] + 1;       // 深度 = 父节点深度+1
    sz[u] = 1;                 // 自身算1个节点
    int maxsz = -1;            // 记录最大子树大小
    for (int v : adj[u]) {     // 遍历所有儿子
        if (v == f) continue;  // 跳过父节点
        dfs1(v, u);            // 递归处理儿子
        sz[u] += sz[v];        // 累加子树大小
        if (sz[v] > maxsz) {   // 更新最大子树和重儿子
            maxsz = sz[v];
            son[u] = v;
        }
    }
}

// 第二次DFS:分配dfn和top,优先重儿子
void dfs2(int u, int t) {
    top[u] = t;                // 设置链顶
    dfn[u] = ++cnt;            // 分配dfs序编号(从1开始)
    rnk[cnt] = u;              // 记录编号对应的节点
    if (son[u]) {              // 如果有重儿子
        dfs2(son[u], t);       // 优先递归重儿子,继承链顶t
    }
    for (int v : adj[u]) {     // 遍历所有儿子
        if (v == fa[u] || v == son[u]) continue; // 跳过父节点和重儿子
        dfs2(v, v);            // 轻儿子自己作为链顶
    }
}

// 线段树构建(依据arr和rnk)
void build(int p, int l, int r) {
    if (l == r) {                              // 叶子节点
        tree[p] = arr[rnk[l]];                 // 该位置对应节点的权值
        return;
    }
    int mid = (l + r) / 2;
    build(p*2, l, mid);                        // 左子树
    build(p*2+1, mid+1, r);                    // 右子树
    tree[p] = tree[p*2] + tree[p*2+1];        // 区间和
}

// 单点修改:将pos位置的值改为val
void update(int p, int l, int r, int pos, int val) {
    if (l == r) {
        tree[p] = val;
        return;
    }
    int mid = (l + r) / 2;
    if (pos <= mid) update(p*2, l, mid, pos, val);
    else update(p*2+1, mid+1, r, pos, val);
    tree[p] = tree[p*2] + tree[p*2+1];
}

// 区间查询:求[ql, qr]的和
int query(int p, int l, int r, int ql, int qr) {
    if (ql <= l && r <= qr) return tree[p]; // 完全包含
    int mid = (l + r) / 2, res = 0;
    if (ql <= mid) res += query(p*2, l, mid, ql, qr);
    if (qr > mid) res += query(p*2+1, mid+1, r, ql, qr);
    return res;
}

// 路径查询:求节点u到v路径上所有节点权值之和
int path_query(int u, int v) {
    int res = 0;
    while (top[u] != top[v]) {          // 不在同一条重链上
        if (dep[top[u]] < dep[top[v]]) swap(u, v); // 让深度大的链顶节点在上方
        res += query(1, 1, n, dfn[top[u]], dfn[u]); // 处理从当前节点到链顶的区间
        u = fa[top[u]];                // 跳到链顶的父亲
    }
    // 现在在同一链上
    if (dep[u] > dep[v]) swap(u, v);   // 使u是深度小的节点(靠近根)
    res += query(1, 1, n, dfn[u], dfn[v]); // 处理最后一个区间
    return res;
}

int main() {
    cin >> n >> q;
    for (int i = 1; i <= n; i++) cin >> arr[i]; // 读入每个节点的初始权值
    for (int i = 1; i < n; i++) {
        int u, v; cin >> u >> v;
        adj[u].push_back(v);
        adj[v].push_back(u);
    }
    dfs1(1, 0);      // 第一次DFS,根为1,父节点为0
    dfs2(1, 1);      // 第二次DFS,根1的链顶是1
    build(1, 1, n); // 构建线段树
    while (q--) {
        int op, x, y;
        cin >> op >> x >> y;
        if (op == 1) {
            // 单点修改:将节点x的值改为y
            update(1, 1, n, dfn[x], y);
        } else {
            // 路径查询:输出x到y的路径和
            cout << path_query(x, y) << endl;
        }
    }
    return 0;
}

Python完整代码实现

import sys
sys.setrecursionlimit(1000000) # 防止递归深度过大

def dfs1(u, f, adj, fa, dep, sz, son):
    fa[u] = f                # 父节点
    dep[u] = dep[f] + 1      # 深度
    sz[u] = 1                # 自身大小
    maxsz = -1
    for v in adj[u]:
        if v == f:
            continue
        dfs1(v, u, adj, fa, dep, sz, son)
        sz[u] += sz[v]
        if sz[v] > maxsz:
            maxsz = sz[v]
            son[u] = v

cnt = 0
def dfs2(u, t, adj, fa, son, top, dfn, rnk):
    global cnt
    top[u] = t               # 链顶
    cnt += 1
    dfn[u] = cnt             # dfs序
    rnk[cnt] = u             # 反向映射
    if son[u]:
        dfs2(son[u], t, adj, fa, son, top, dfn, rnk) # 先重儿子
    for v in adj[u]:
        if v != fa[u] and v != son[u]:
            dfs2(v, v, adj, fa, son, top, dfn, rnk) # 轻儿子自己为链顶

class SegmentTree:
    def __init__(self, n, arr, rnk):
        self.n = n
        self.tree = [0] * (4 * n)
        self.arr = arr
        self.rnk = rnk
        self.build(1, 1, n)
    def build(self, p, l, r):
        if l == r:
            self.tree[p] = self.arr[self.rnk[l]]
            return
        mid = (l + r) // 2
        self.build(p*2, l, mid)
        self.build(p*2+1, mid+1, r)
        self.tree[p] = self.tree[p*2] + self.tree[p*2+1]
    def update(self, p, l, r, pos, val):
        if l == r:
            self.tree[p] = val
            return
        mid = (l + r) // 2
        if pos <= mid:
            self.update(p*2, l, mid, pos, val)
        else:
            self.update(p*2+1, mid+1, r, pos, val)
        self.tree[p] = self.tree[p*2] + self.tree[p*2+1]
    def query(self, p, l, r, ql, qr):
        if ql <= l and r <= qr:
            return self.tree[p]
        mid = (l + r) // 2
        res = 0
        if ql <= mid:
            res += self.query(p*2, l, mid, ql, qr)
        if qr > mid:
            res += self.query(p*2+1, mid+1, r, ql, qr)
        return res

def path_query(u, v, fa, dep, top, dfn, seg):
    res = 0
    while top[u] != top[v]:
        if dep[top[u]] < dep[top[v]]:
            u, v = v, u
        res += seg.query(1, 1, seg.n, dfn[top[u]], dfn[u])
        u = fa[top[u]]
    if dep[u] > dep[v]:
        u, v = v, u
    res += seg.query(1, 1, seg.n, dfn[u], dfn[v])
    return res

def main():
    n, q = map(int, input().split())
    arr = [0] + list(map(int, input().split()))  # 节点权值,下标从1开始
    adj = [[] for _ in range(n+1)]
    for _ in range(n-1):
        u, v = map(int, input().split())
        adj[u].append(v)
        adj[v].append(u)
    fa = [0]*(n+1)
    dep = [0]*(n+1)
    sz = [0]*(n+1)
    son = [0]*(n+1)
    dfs1(1, 0, adj, fa, dep, sz, son)
    global cnt
    cnt = 0
    top = [0]*(n+1)
    dfn = [0]*(n+1)
    rnk = [0]*(n+1)
    dfs2(1, 1, adj, fa, son, top, dfn, rnk)
    seg = SegmentTree(n, arr, rnk)
    for _ in range(q):
        op, x, y = map(int, input().split())
        if op == 1:
            seg.update(1, 1, n, dfn[x], y)
        else:
            print(path_query(x, y, fa, dep, top, dfn, seg))

if __name__ == "__main__":
    main()

常见错误与避坑指南

  1. 忘记处理链顶深度比较时的方向:在 path_query 中,当 top[u] != top[v] 时,一定要比较 dep[top[u]]dep[top[v]],让链顶深度较大的节点向上跳。如果反过来,会导致死循环或错误结果。记住:谁的链顶更深(离根更远),谁就先跳

  2. 线段树建树时用错数组:建树时需要用 rnk[l] 得到节点编号,再取 arr[rnk[l]]。如果用 dfn[l] 就错了,因为 dfn 是节点到位置,rnk 是位置到节点。

  3. 节点编号从1开始,忘记调整:很多树的问题中根是1,但有时根可能不是0,注意数组大小要开够,比如 MAXN 设为最大节点数+5。

  4. 重儿子初始化:如果节点没有儿子,则 son[u] 应该为0。在第二次DFS中需要判断 if (son[u]) 才递归重儿子,否则会出错。

  5. 递归深度过大:Python默认递归深度有限,一定要设置 sys.setrecursionlimit(1000000)。C++一般没问题,但某些编译器也有栈限制,可以改成迭代或使用全局数组。

  6. 路径查询时最后一步的区间顺序:当 top[u]==top[v] 时,我们要处理从深度小的节点到深度大的节点。如果直接使用 [dfn[u], dfn[v]],要确保 dep[u] <= dep[v]。所以先判断深度,必要时交换。

完整示例与测试

假设输入如下(对应我们之前图示的树,节点权值分别为1~7):

7 5
1 2 3 4 5 6 7
1 2
1 3
2 4
2 5
3 6
6 7
2 4 7
1 4 100
2 4 7
2 1 3
2 5 5

第一行:7个节点,5次操作
第二行:每个节点的初始权值
接下来6行:树的边
然后5行操作:

  • 2 4 7:查询节点4到7的路径和
  • 1 4 100:将节点4的权值修改为100
  • 2 4 7:再次查询
  • 2 1 3:查询节点1到3的路径和
  • 2 5 5:查询节点5到5(自身)的路径和

预期输出:

路径4->7初始:节点4(1)+2(2)+1(1)+3(3)+6(6)+7(7) = 20? 不,路径是4-2-1-3-6-7,节点值:4:4? 注意我们权值arr: 节点1=1,2=2,3=3,4=4,5=5,6=6,7=7。所以路径和=4+2+1+3+6+7=23。
修改后:节点4变100,路径和=100+2+1+3+6+7=119。
节点1->3:1+3=4。
节点5->5:5。

可以手动验证。

总结要点

  1. 树链剖分把树上问题转化为区间问题,特别适合路径查询和路径修改。
  2. 通过重链剖分,每条路径被分解成O(log n)个区间,配合线段树可在O(log^2 n)内完成。
  3. 核心是两次DFS:第一次求子树大小、重儿子;第二次分配dfs序、标记链顶。
  4. 注意:路径查询时要比较链顶深度,确保向上跳的方向正确。
  5. 树剖还可以用于子树操作(因为子树内dfn连续),功能强大。

树链剖分就像给树装上了一个“分段导航系统”,让任何路径都能转化为连续的区间,轻松处理!如果你想深入学习,可以继续探索以下内容:

  • LCA(最近公共祖先):树剖也可以高效求LCA,而且比倍增法更简洁。
  • 线段树区间修改:如果要做路径修改(比如路径加一个值),需要线段树支持区间更新和懒标记。
  • 树状数组替代线段树:如果只求和,树状数组更快。
  • 动态树(Link-Cut Tree):更高级的树上路径处理,支持边和点的动态变化。

例题精讲

1单选题

在树链剖分(HLD)中,对于一个节点u,定义其“重儿子”为?

Au的所有子节点中子树大小最大的那个(若相等则任选一个)
Bu的所有子节点中深度最大的那个
Cu的所有子节点中编号最大的那个
Du的所有子节点中与u相连的边权最大的那个
2判断题

树链剖分可以将树上任意两点之间的路径分解为O(log n)条重链(或区间),从而利用线段树等数据结构进行高效的路径查询与修改。

3填空题
下面是一段使用树链剖分求树上两点u, v的LCA的代码片段,请补全缺失的部分。(假设已经预处理了fa, dep, top, son等数组,其中top[u]表示u所在重链的顶端节点,son[u]表示u的重儿子)

int lca(int u, int v) {
    while (top[u] != top[v]) {
        if (dep[top[u]] < dep[top[v]]) swap(u, v);
        u = ___;
    }
    return dep[u] < dep[v] ? u : v;
}
4单选题

树链剖分预处理(包括第一遍DFS求size, dep, fa, son,第二遍DFS求top, id)的时间复杂度是?

AO(n)
BO(n log n)
CO(n^2)
DO(log n)
5判断题

在树链剖分中,对于任意一条轻边(即连接父节点和非重儿子的边),从该轻边出发向下走,子树大小至少减半。