18.并查集.md 9.8 KB

18. 并查集

概念

并查集(Union-Find / Disjoint Set,不相交集合)是一种树形数据结构,用于高效处理「不相交集合」的合并(Union)查询(Find)问题。它的核心思想:每个集合用一棵树表示,树的根结点作为整个集合的「代表元素」。判断两个元素是否属于同一集合,只需看它们的根是否相同。

三个基本操作:

  • 初始化(init):每个元素自成一个集合,即每个结点都是自己所在树的根(parent[i] = i)。
  • 查(Find):查找元素 x 所在集合的根。实现上沿 parent 指针不断向上走,直到 parent[x] == x 为止。
  • 并(Union):把两个元素所在的集合合并。做法是找到两个根 rxry,再把其中一个根挂到另一个根下面(parent[rx] = ry 或反过来)。

两个关键优化:

  • 路径压缩(Path Compression):在 find 的过程中,把沿途经过的结点直接挂到根上。这样下次再查这些结点时只需一步就能到根,树变得很「扁」。
  • 按秩合并(Union by Rank):合并时总是把较矮的树挂到较高的树上(记录每棵树的「秩」= 近似高度),避免树退化成一条长链。

为什么能近似 O(1):同时采用「路径压缩 + 按秩合并」后,单次操作的均摊时间复杂度为 O(α(n)),其中 α(n) 是反阿克曼函数。α(n) 增长极其缓慢——即使 n 达到可观测宇宙中的原子数量,α(n) 也不超过 4。因此在任何实际场景下,都可认为并查集的单次操作近似 O(1)

应用场景:判断连通分量、最小生成树算法(Kruskal)中判环、社交网络中两人是否在同一圈子、图像连通区域标记、等价类划分等。

核心操作 / 算法

  1. init(n)parent[i] = irank[i] = 0,初始化 n 个独立集合。
  2. find(x):递归(或迭代)沿 parent 链找根,同时做路径压缩
  3. union(x, y):先 find 两个根;若相同则已连通;否则按秩合并(矮树挂高树,秩相等时任意合并并让新根秩 +1)。
  4. 判连通:find(x) == find(y)
  5. 统计连通分量个数:维护一个 count,每次成功合并 count--(或在初始化后数根结点个数)。

复杂度分析

优化情况 find union(含 find) 说明
朴素实现(无优化) 最坏 O(n) 最坏 O(n) 树可能退化为单链,每次 find 都要走 O(n) 步
仅路径压缩 均摊 O(log n) 均摊 O(log n) 树变扁,但最坏仍是 O(n)
仅按秩合并 O(log n) O(log n) 树高被限制在 O(log n),树一定较平衡
路径压缩 + 按秩合并 均摊 O(α(n)) 均摊 O(α(n)) α(n) 为反阿克曼函数,实际中不超过 4,近似 O(1)

空间复杂度恒为 O(n):只需 parentrank 两个长度 n 的数组。

要点:两种优化缺一不可——路径压缩让树变扁、按秩合并防止退化,二者配合才得到近似 O(1) 的均摊界。

语言实现

以下四种实现完全等价,均包含「路径压缩 + 按秩合并」,并演示同一场景:给定 6 个结点与若干连接关系,判断元素是否连通、统计连通分量个数。

C

#include <stdio.h>

#define MAXN 1000

// 全局并查集数组:parent[x] 为 x 的父结点,rankArr[x] 为树的高度(秩)
int parent[MAXN];
int rankArr[MAXN];

// 初始化:每个元素自成一个集合
void init(int n) {
    for (int i = 0; i < n; i++) {
        parent[i] = i;
        rankArr[i] = 0;
    }
}

// 查找 x 所在集合的根,同时进行路径压缩
int find(int x) {
    if (parent[x] != x)
        parent[x] = find(parent[x]);   // 递归路径压缩:沿途结点直接挂到根上
    return parent[x];
}

// 合并两个元素所在的集合,按秩合并
void unionSet(int x, int y) {
    int rx = find(x);
    int ry = find(y);
    if (rx == ry) return;              // 已在同一集合,无需合并
    // 将较矮的树并入较高的树,秩相等时随意并入并让新根秩 +1
    if (rankArr[rx] < rankArr[ry]) {
        parent[rx] = ry;
    } else if (rankArr[rx] > rankArr[ry]) {
        parent[ry] = rx;
    } else {
        parent[ry] = rx;
        rankArr[rx]++;
    }
}

int main() {
    int n = 6;
    init(n);

    // 给定若干连接关系:0-1, 1-2, 3-4
    int links[][2] = {{0, 1}, {1, 2}, {3, 4}};
    int m = sizeof(links) / sizeof(links[0]);
    for (int i = 0; i < m; i++)
        unionSet(links[i][0], links[i][1]);

    // 判断连通性
    printf("0 与 2 是否连通: %s\n", find(0) == find(2) ? "是" : "否");
    printf("0 与 4 是否连通: %s\n", find(0) == find(4) ? "是" : "否");

    // 统计连通分量个数:统计根结点(parent[i] == i)的个数
    int count = 0;
    for (int i = 0; i < n; i++)
        if (find(i) == i) count++;
    printf("连通分量个数: %d\n", count);
    return 0;
}

C++

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

class UnionFind {
private:
    vector<int> parent;   // 父结点数组
    vector<int> rank_;    // 秩(树的近似高度)
    int count;            // 连通分量个数
public:
    // 初始化:每个元素自成一个集合
    UnionFind(int n) {
        parent.resize(n);
        rank_.assign(n, 0);
        count = n;
        for (int i = 0; i < n; i++) parent[i] = i;
    }

    // 查找根结点,带路径压缩
    int find(int x) {
        if (parent[x] != x)
            parent[x] = find(parent[x]);
        return parent[x];
    }

    // 按秩合并,成功合并时连通分量减一
    void unite(int x, int y) {
        int rx = find(x), ry = find(y);
        if (rx == ry) return;
        if (rank_[rx] < rank_[ry]) {
            parent[rx] = ry;
        } else if (rank_[rx] > rank_[ry]) {
            parent[ry] = rx;
        } else {
            parent[ry] = rx;
            rank_[rx]++;
        }
        count--;
    }

    bool connected(int x, int y) { return find(x) == find(y); }

    int components() { return count; }
};

int main() {
    int n = 6;
    UnionFind uf(n);

    // 给定若干连接关系:0-1, 1-2, 3-4
    int links[][2] = {{0, 1}, {1, 2}, {3, 4}};
    for (auto& e : links) uf.unite(e[0], e[1]);

    cout << "0 与 2 是否连通: " << (uf.connected(0, 2) ? "是" : "否") << endl;
    cout << "0 与 4 是否连通: " << (uf.connected(0, 4) ? "是" : "否") << endl;
    cout << "连通分量个数: " << uf.components() << endl;
    return 0;
}

Java

public class UnionFind {
    private int[] parent;   // 父结点数组
    private int[] rank_;    // 秩(树的近似高度)
    private int count;      // 连通分量个数

    // 初始化:每个元素自成一个集合
    public UnionFind(int n) {
        parent = new int[n];
        rank_ = new int[n];
        count = n;
        for (int i = 0; i < n; i++) parent[i] = i;
    }

    // 查找根结点,带路径压缩
    public int find(int x) {
        if (parent[x] != x) parent[x] = find(parent[x]);
        return parent[x];
    }

    // 按秩合并,成功合并时连通分量减一
    public void unite(int x, int y) {
        int rx = find(x), ry = find(y);
        if (rx == ry) return;
        if (rank_[rx] < rank_[ry]) {
            parent[rx] = ry;
        } else if (rank_[rx] > rank_[ry]) {
            parent[ry] = rx;
        } else {
            parent[ry] = rx;
            rank_[rx]++;
        }
        count--;
    }

    public boolean connected(int x, int y) { return find(x) == find(y); }

    public int components() { return count; }

    public static void main(String[] args) {
        int n = 6;
        UnionFind uf = new UnionFind(n);

        // 给定若干连接关系:0-1, 1-2, 3-4
        int[][] links = {{0, 1}, {1, 2}, {3, 4}};
        for (int[] e : links) uf.unite(e[0], e[1]);

        System.out.println("0 与 2 是否连通: " + (uf.connected(0, 2) ? "是" : "否"));
        System.out.println("0 与 4 是否连通: " + (uf.connected(0, 4) ? "是" : "否"));
        System.out.println("连通分量个数: " + uf.components());
    }
}

Python

class UnionFind:
    """并查集:路径压缩 + 按秩合并"""

    def __init__(self, n):
        self.parent = list(range(n))   # 每个元素的父结点,初始指向自己
        self.rank_ = [0] * n           # 秩(树的近似高度)
        self.count = n                 # 连通分量个数

    def find(self, x):
        """查找根结点,带路径压缩"""
        if self.parent[x] != x:
            self.parent[x] = self.find(self.parent[x])
        return self.parent[x]

    def unite(self, x, y):
        """按秩合并,成功合并时连通分量减一"""
        rx, ry = self.find(x), self.find(y)
        if rx == ry:
            return
        if self.rank_[rx] < self.rank_[ry]:
            self.parent[rx] = ry
        elif self.rank_[rx] > self.rank_[ry]:
            self.parent[ry] = rx
        else:
            self.parent[ry] = rx
            self.rank_[rx] += 1
        self.count -= 1

    def connected(self, x, y):
        return self.find(x) == self.find(y)

    def components(self):
        return self.count


if __name__ == "__main__":
    n = 6
    uf = UnionFind(n)

    # 给定若干连接关系:0-1, 1-2, 3-4
    links = [(0, 1), (1, 2), (3, 4)]
    for a, b in links:
        uf.unite(a, b)

    print("0 与 2 是否连通:", "是" if uf.connected(0, 2) else "否")
    print("0 与 4 是否连通:", "是" if uf.connected(0, 4) else "否")
    print("连通分量个数:", uf.components())