CC++ & Algorithm

点分治 —— 分而治之的树上统计

极难2
语言版本:C++
概述:把一棵大树从重心切开,分别处理每个小部分,再合并结果,适用于统计路径等树上问题。

点分治:像切西瓜一样解决树上路径问题

你有没有遇到过这样的问题:在一棵很大的树上,想数一数有多少条路径的长度不超过某个数?比如,学校里的社团关系网,每条边是两个人之间的距离,你想知道有多少对同学的距离不超过10米。直接一个一个数,如果树有上万个结点,那就要花很长时间。今天我们要学的点分治,就是一种又快又聪明的办法——它把大树从中间切开,分成小块,分别处理,再合并结果。就像切西瓜:先找到中间一刀切下去,把西瓜切成几块,每块再继续切,最后加起来就是整个西瓜的份量。


什么是“分而治之”?

分而治之(Divide and Conquer)是一种解决问题的思路:把一个大问题分解成几个相同的小问题,分别解决,再把结果合并。点分治就是专门针对这种结构的分治法。它特别适合处理树上路径统计一类的问题,比如:

  • 统计距离不超过 k 的点对数量
  • 统计路径长度等于某个值的路径条数
  • 统计路径上权值乘积、和等满足条件的路径

它的核心思想是:先找到树的重心,把所有经过重心的路径统计出来,然后去掉重心,对剩下的每一棵子树递归进行同样的操作。为什么选择重心?因为重心能把树分成大小尽量均匀的几块,这样每次递归的规模都大约减少一半,总时间复杂度能控制在 O(n log n) 左右,非常高效。


关键概念一:树的重心

树的重心,简单理解就是:如果把树从某个结点处“切开”,分出的几块(子树)中,最大的一块的大小最小。这个结点就是重心。换句话说,重心是树上最“平衡”的点。

生活例子:想象你有一把树枝,想把它扎成几捆,你希望每捆差不多粗。你会先找到中间最粗的那根枝干,把它抽出来,这样剩下的几捆大小就差不多。那根中间的枝干就是树的重心。

如何找重心?

我们需要计算每个结点的子树大小,以及“去掉这个结点后剩下的最大子树大小”。具体做法:

  1. 任选一个结点作为根,算一遍所有子树的大小(用一次DFS)。
  2. 对于每个结点,考虑两部分:它的子结点下面的子树大小,以及它父方向的那一块的大小(等于总节点数减去当前结点自身的子树大小)。取这两部分的最大值。
  3. 所有结点中,这个最大值最小的那个结点就是重心。
// 计算子树大小,同时记录每个结点的最大子树大小(不包含父方向)
void get_size(int u, int f) {
    sz[u] = 1;             // 当前结点大小为1(自身)
    maxp[u] = 0;           // 初始化最大子树为0
    for (auto &e : G[u]) {
        int v = e.first;   // 邻居结点
        if (v == f || del[v]) continue; // 不能走回父结点或已删除
        get_size(v, u);
        sz[u] += sz[v];               // 累加子树大小
        maxp[u] = max(maxp[u], sz[v]); // 更新“子方向”的最大子树
    }
}

// 找重心:u当前结点,f父结点,total整棵树的总大小
int get_centroid(int u, int f, int total) {
    int res = u;                     // 先假设当前结点是重心
    maxp[u] = max(maxp[u], total - sz[u]); // 考虑父方向那一块的大小
    for (auto &e : G[u]) {
        int v = e.first;
        if (v == f || del[v]) continue;
        int tmp = get_centroid(v, u, total);
        if (maxp[tmp] < maxp[res]) res = tmp; // 保留最大子树最小的结点
    }
    return res;
}

关键概念二:统计经过重心的路径

找到重心后,我们要统计所有经过这个重心的路径。注意:这里的路径不一定是直的,可以是任意两个结点之间的路径。如果路径经过重心,那么它一定是由两条从重心出发的“半路径”拼接而成(包括重心自身)。

生活例子:你想统计学校里所有“经过操场”的两人组合。你可以先站在操场,记下每个同学到操场的距离,然后任意两个同学的距离就是他们各自到操场距离之和(如果踩在操场上的人不算)。我们要数出有多少对同学,他们的距离和 ≤ k。

具体做法:

  1. 从重心出发,对每个子树分别进行DFS,记录下该子树中所有结点到重心的距离,存到一个数组 sub 里。
  2. 在合并之前,先把这个子树内部的路径减去,因为后面我们会把所有子树合并再统计,如果直接合并,会把子树内部的不经过重心的路径也统计进来(比如两个结点都在同一个子树里,它们之间的路径不一定经过重心)。所以我们要先对每个子树单独算一遍“内部满足条件的对数”,从总答案中减去,最后再把所有距离混在一起算一遍,这样留下来的就是只经过重心的路径。
  3. 用一个双指针技巧,对排序后的距离数组统计满足 dist[i] + dist[j] ≤ k 的对数。

双指针统计的原理

假设有一个排好序的距离数组 vec,我们要数有多少对 (i, j)i < j 满足 vec[i] + vec[j] ≤ k。我们用两个指针:一个 i 从头开始,一个 j 从尾开始。对于每个 i,把 j 向左移动,直到 vec[i] + vec[j] ≤ k 不成立为止。这时,i+1j 这些元素都可以和 vec[i] 配对,数量就是 j - i。累加所有 i 的结果即可。

// 计算一个距离数组中有多少对 (i,j) 满足距离和 <= k
int calc(vector<int> &vec) {
    sort(vec.begin(), vec.end());      // 升序排序
    int res = 0;
    int j = (int)vec.size() - 1;
    for (int i = 0; i < j; i++) {
        while (i < j && vec[i] + vec[j] > k) j--; // 移动右指针
        res += j - i;                  // 从 i+1 到 j 都可以配对
    }
    return res;
}

注意:上面代码中,vec 里可能包含重心自身的距离0,这样就能统计到以一个端点为重心的情况。


关键概念三:递归处理子树

统计完经过重心的路径后,我们把重心标记为“已删除”,然后对它的每一个子结点(也就是每一棵子树)递归地进行同样的点分治过程。这样,每棵子树内部的问题又被分解得更小,直到子树大小为1。

void divide(int u) {
    get_size(u, -1);                        // 计算当前树的子树大小
    int cen = get_centroid(u, -1, sz[u]);   // 找到重心
    del[cen] = true;                        // 标记重心已删除

    // 统计经过重心的路径
    vector<int> total;
    total.push_back(0);                     // 重心自身距离为0
    for (auto &e : G[cen]) {
        int v = e.first, w = e.second;      // 邻居结点和边权
        if (del[v]) continue;
        vector<int> sub;
        get_dist(v, cen, w, sub);           // 收集该子树所有结点到重心的距离
        ans -= calc(sub);                   // 减去子树内部重复统计
        total.insert(total.end(), sub.begin(), sub.end()); // 合并到总数组
    }
    ans += calc(total);                     // 加上合并后的结果

    // 递归处理每一个子树
    for (auto &e : G[cen]) {
        int v = e.first;
        if (!del[v]) divide(v);
    }
}

这里 get_dist 函数负责从某结点出发,收集所有未被删除结点到该起点的距离:

void get_dist(int u, int f, int d, vector<int> &vec) {
    vec.push_back(d);                       // 记录距离
    for (auto &e : G[u]) {
        int v = e.first, w = e.second;
        if (v == f || del[v]) continue;
        get_dist(v, u, d + w, vec);         // 累加距离继续深入
    }
}

新手常犯的错误

  1. 忘记标记 del 数组:如果不把重心标记为已删除,递归时还会访问它,导致无限循环或错误统计。
  2. 递归前没有先计算子树大小divide 函数入口处要先调用 get_size,否则 get_centroid 里用的 sz 可能是旧值。
  3. calc 函数中双指针的边界:注意 while 条件中的 i < j,以及 j 的初始值。如果 vec 为空或只有一个元素,要保证不会越界。
  4. 答案的初始化ans 要初始化为0,并且在每次递归前要保证当前树的 total 正确。
  5. 减去子树内重复时,注意顺序:要先统计子树内部的,再合并。如果反过来,合并后再减,可能会减错。

完整可运行的示例

下面是一个完整的程序,读入 n, k 和一棵无向树(边有权值),输出距离 ≤ k 的点对数量。代码中增加了详细的输入输出和注释。

#include <bits/stdc++.h>
using namespace std;

const int MAXN = 20005;
int n, k, ans;
vector<pair<int, int>> G[MAXN]; // 邻接表,存储 (邻居, 边权)
bool del[MAXN];                 // 标记重心已被删除
int sz[MAXN], maxp[MAXN];       // 子树大小,最大子树大小

// 计算子树大小,同时更新最大子树(子方向)
void get_size(int u, int f) {
    sz[u] = 1;                 // 当前结点大小为1
    maxp[u] = 0;               // 初始化最大子树为0
    for (auto &e : G[u]) {
        int v = e.first;       // 邻居结点
        if (v == f || del[v]) continue; // 跳过父结点或已删除
        get_size(v, u);
        sz[u] += sz[v];        // 累加子树大小
        maxp[u] = max(maxp[u], sz[v]); // 更新最大子树
    }
}

// 找重心,total为当前树的总大小
int get_centroid(int u, int f, int total) {
    int res = u;               // 先假设当前结点是重心
    maxp[u] = max(maxp[u], total - sz[u]); // 考虑父方向的那一块大小
    for (auto &e : G[u]) {
        int v = e.first;
        if (v == f || del[v]) continue;
        int tmp = get_centroid(v, u, total);
        if (maxp[tmp] < maxp[res]) res = tmp; // 保留最大子树最小的结点
    }
    return res;
}

// 收集以u为根(不经过已删除结点)的所有节点到起点的距离
void get_dist(int u, int f, int d, vector<int> &vec) {
    vec.push_back(d);          // 记录距离
    for (auto &e : G[u]) {
        int v = e.first, w = e.second;
        if (v == f || del[v]) continue;
        get_dist(v, u, d + w, vec);
    }
}

// 计算有序数组vec中满足距离和<=k的对数(双指针)
int calc(vector<int> &vec) {
    sort(vec.begin(), vec.end());
    int res = 0;
    int j = (int)vec.size() - 1;
    for (int i = 0; i < j; i++) {
        while (i < j && vec[i] + vec[j] > k) j--;
        res += j - i;          // 此时j右边的都不能配对
    }
    return res;
}

// 点分治主函数
void divide(int u) {
    get_size(u, -1);                              // 先计算子树大小
    int cen = get_centroid(u, -1, sz[u]);         // 找重心
    del[cen] = true;                              // 标记重心已删除

    // 统计经过重心的路径
    vector<int> total;
    total.push_back(0);                           // 重心自身
    for (auto &e : G[cen]) {
        int v = e.first, w = e.second;           // 邻居和边权
        if (del[v]) continue;
        vector<int> sub;
        get_dist(v, cen, w, sub);                 // 收集该子树所有距离
        // 注意:先减去子树内部的统计,避免重复
        ans -= calc(sub);
        // 把子树距离合并到总数组
        total.insert(total.end(), sub.begin(), sub.end());
    }
    ans += calc(total);                          // 加上合并后经过重心的对数

    // 递归处理每个子树
    for (auto &e : G[cen]) {
        int v = e.first;
        if (!del[v]) divide(v);
    }
}

int main() {
    ios::sync_with_stdio(false);
    cin.tie(0);

    // 输入:第一行n和k,接下来n-1行每行u v w
    cin >> n >> k;
    for (int i = 0; i < n - 1; i++) {
        int u, v, w;
        cin >> u >> v >> w;
        G[u].push_back({v, w});
        G[v].push_back({u, w});
    }

    ans = 0;                  // 初始化答案
    divide(0);                // 从0号结点开始(结点编号从0到n-1)
    cout << ans << endl;     // 输出满足条件的点对数量

    return 0;
}

输入样例

5 4
0 1 1
0 2 2
1 3 3
1 4 2

输出4(解释:距离≤4的点对有 (0,1)、(0,2)、(1,3)、(1,4) 共4对)


相关指引

点分治是处理树上路径问题的利器,继续学习可以关注:

  • 动态点分治:当树边或点权会变化时,使用点分树(Centroid Decomposition Tree)来维护。
  • 其他分治技巧:比如边分治、树分块等,各有适用场景。
  • 常见变体:统计路径长度等于k、路径上权值乘积满足要求、路径上颜色种类等。只要把 calc 函数中的统计逻辑换一下就行。
  • 离线与在线:点分治通常是离线算法,但结合数据结构(如树状数组)可以处理在线查询。

掌握了点分治,你就拥有了一把处理树上复杂统计的万能钥匙。下次遇到“树上路径”问题,记得试试“切西瓜”的方法!

例题精讲

1单选题

点分治算法在处理树上路径统计问题时,每次选择树的重心作为分治中心,其主要目的是?

A保证递归深度不超过O(log n)
B减少空间复杂度
C避免重复计数
D方便代码实现
2判断题

点分治中,处理子树时,需要先用当前重心计算经过重心的路径,然后对每个子树递归,并且在递归前需要减去子树内部形成的路径,以避免重复计数。

3填空题
void getdis(int u, int fa, int d) {
    ___; // 请填写代码
    for(int i=head[u];i;i=e[i].nxt){
        int v=e[i].to;
        if(v==fa || vis[v]) continue;
        getdis(v,u,d+e[i].w);
    }
}
4填空题
int calc(int u, int d) {
    cnt = 0;
    getdis(u, 0, d);
    sort(dis+1, dis+cnt+1);
    int l=1, r=cnt, res=0;
    while(l<r){
        if(dis[l]+dis[r] <= k){
            ___; // 请填写代码
            l++;
        } else r--;
    }
    return res;
}
5单选题

在点分治的实现过程中,通常会先计算经过当前重心的路径,然后递归处理各个子树。在计算经过重心的路径时,如果直接将所有节点到重心的距离混合排序统计,会导致什么错误?

A增加时间复杂度
B统计到不经过重心的路径
C重复统计路径
D导致递归深度增加