前置知识: 算法与数据结构

并查集

2 minIntermediate2026/6/14

并查集(Union-Find)数据结构:路径压缩与按秩合并优化、连通性判断、Kruskal 最小生成树应用。

1. 并查集基础

1.1 问题背景

并查集(Disjoint Set Union, DSU)解决动态连通性问题:维护若干不相交集合,支持合并和查询操作。

初始: {0} {1} {2} {3} {4}  ← 每个元素独立
union(1,2): {0} {1,2} {3} {4}
union(3,4): {0} {1,2} {3,4}
union(1,3): {0} {1,2,3,4}
find(2) == find(4)? → true  ← 连通
find(0) == find(2)? → false ← 不连通

1.2 基本操作

操作含义时间复杂度
find(x)查找 x 的根节点O(α(n))O(\alpha(n))
union(x, y)合并 x 和 y 所在集合O(α(n))O(\alpha(n))
connected(x, y)判断 x 和 y 是否连通O(α(n))O(\alpha(n))

α(n)\alpha(n) 是反阿克曼函数,增长极慢,实际可视为常数。

2. 实现与优化

2.1 朴素实现

class UnionFind:
    def __init__(self, n):
        self.parent = list(range(n))

    def find(self, x):
        while self.parent[x] != x:
            x = self.parent[x]
        return x

    def union(self, x, y):
        root_x = self.find(x)
        root_y = self.find(y)
        if root_x != root_y:
            self.parent[root_x] = root_y

最坏情况退化为链表,findO(n)O(n)

2.2 路径压缩

查找时将路径上所有节点直接指向根:

def find(self, x):
    if self.parent[x] != x:
        self.parent[x] = self.find(self.parent[x])  # 递归压缩
    return self.parent[x]

# 迭代版本(避免栈溢出)
def find(self, x):
    root = x
    while self.parent[root] != root:
        root = self.parent[root]
    while self.parent[x] != root:
        self.parent[x], x = root, self.parent[x]
    return root
压缩前:  0 → 1 → 2 → 3 → 4
压缩后:  0 → 4, 1 → 4, 2 → 4, 3 → 4

2.3 按秩合并

将较矮的树合并到较高的树下:

class UnionFind:
    def __init__(self, n):
        self.parent = list(range(n))
        self.rank = [0] * n  # 树的高度

    def find(self, x):
        if self.parent[x] != x:
            self.parent[x] = self.find(self.parent[x])
        return self.parent[x]

    def union(self, x, y):
        root_x = self.find(x)
        root_y = self.find(y)
        if root_x == root_y:
            return False
        if self.rank[root_x] < self.rank[root_y]:
            self.parent[root_x] = root_y
        elif self.rank[root_x] > self.rank[root_y]:
            self.parent[root_y] = root_x
        else:
            self.parent[root_x] = root_y
            self.rank[root_y] += 1
        return True

2.4 按大小合并

另一种优化,将较小的树合并到较大的树下:

class UnionFind:
    def __init__(self, n):
        self.parent = list(range(n))
        self.size = [1] * n  # 集合大小

    def union(self, x, y):
        root_x = self.find(x)
        root_y = self.find(y)
        if root_x == root_y:
            return False
        if self.size[root_x] < self.size[root_y]:
            root_x, root_y = root_y, root_x
        self.parent[root_y] = root_x
        self.size[root_x] += self.size[root_y]
        return True

2.5 复杂度分析

优化findunion
无优化O(n)O(n)O(n)O(n)
仅路径压缩O(logn)O(\log n) 均摊O(logn)O(\log n) 均摊
仅按秩合并O(logn)O(\log n)O(logn)O(\log n)
两者结合O(α(n))O(\alpha(n))O(α(n))O(\alpha(n))

α(n)4(对所有实际可能出现的 n\alpha(n) \leq 4 \quad \text{(对所有实际可能出现的 } n\text{)}

3. 应用场景

3.1 Kruskal 最小生成树

def kruskal(n, edges):
    edges.sort(key=lambda e: e[2])  # 按权重排序
    uf = UnionFind(n)
    mst = []
    total = 0
    for u, v, w in edges:
        if uf.union(u, v):  # 不在同一集合 → 加入MST
            mst.append((u, v, w))
            total += w
            if len(mst) == n - 1:
                break
    return mst, total

3.2 省份数量(LeetCode 547)

def findCircleNum(isConnected):
    n = len(isConnected)
    uf = UnionFind(n)
    for i in range(n):
        for j in range(i + 1, n):
            if isConnected[i][j]:
                uf.union(i, j)
    return len(set(uf.find(i) for i in range(n)))

3.3 冗余连接(LeetCode 684)

def findRedundantConnection(edges):
    uf = UnionFind(len(edges) + 1)
    for u, v in edges:
        if not uf.union(u, v):  # 已连通 → 冗余边
            return [u, v]

3.4 岛屿数量(LeetCode 200)

def numIslands(grid):
    if not grid: return 0
    m, n = len(grid), len(grid[0])
    uf = UnionFind(m * n + 1)  # +1 为水域虚拟节点
    water = m * n

    for i in range(m):
        for j in range(n):
            if grid[i][j] == '0':
                uf.union(i * n + j, water)
            else:
                for di, dj in [(0,1),(1,0)]:
                    ni, nj = i + di, j + dj
                    if ni < m and nj < n and grid[ni][nj] == '1':
                        uf.union(i * n + j, ni * n + nj)

    roots = set()
    for i in range(m * n):
        if grid[i // n][i % n] == '1':
            roots.add(uf.find(i))
    return len(roots)