Data Structures · Lesson 15 of 16

Union-Find (Disjoint Set Union)

Learn union-find (disjoint set union) with path compression and union by rank to group items, count connected components and detect cycles in Python.

  • Intermediate
  • 16 min read
  • 4 objectives

Before this lessonLesson 14: Graphs

What you will learn

  • Explain what find and union do
  • Implement path compression and union by size
  • Count connected components and detect cycles
  • Use union-find in Kruskal's minimum spanning tree

Your Progress

0 of 16 lessons 0%

  • Lessons0 / 16
  • Completed0
  • Est. time left~ 4 hours

Create a free account to keep your progress on every device.

Tip: pressing Next marks this lesson complete automatically.

Some problems are all about grouping. Which users are connected through friendships? Which accounts belong to the same person because they share an email? Do these network cables form a loop? You could run a graph search each time, but when connections keep arriving one by one, that repeats a lot of work.

Union-find, also called a disjoint set union (DSU), answers two questions almost instantly: "are a and b in the same group?" and "merge the groups of a and b". With two small optimizations, each operation costs effectively O(1).

The idea: every group has a representative

Each element points to a parent. Following parents upward eventually reaches a root that points to itself; the root is the group's representative. Two elements are in the same group exactly when they have the same root.

  • find(x): follow parents from x to the root and return it.
  • union(a, b): find both roots; if they differ, make one root point to the other.
parent = list(range(6))    # everyone starts as their own group

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

def union(a, b):
    ra, rb = find(a), find(b)
    if ra != rb:
        parent[rb] = ra

union(0, 1)
union(1, 2)
union(3, 4)
print(parent)
print(find(2) == find(0), find(2) == find(3))
Output
[0, 0, 0, 3, 3, 5]
True False

The problem: tall trees

This naive version can build long chains. Union 0-1, then 1-2, then 2-3 in an unlucky direction, and find becomes O(n), just like a degenerate BST. Two fixes solve it.

Fix 1: union by size (or rank)

When merging, attach the smaller tree under the root of the larger one. An element's depth only grows when its tree at least doubles in size, so no tree gets taller than log2 n. Union by rank (an upper bound on height) works the same way.

Fix 2: path compression

During find, once you know the root, make every node you passed point directly to it. The next find on any of them takes one step. The tree flattens itself as you use it.

With both fixes, a sequence of m operations costs O(m α(n)), where α is the inverse Ackermann function. It grows so slowly that it is at most 4 for any input that could fit in the universe, so treat each operation as constant time.

class DSU:
    def __init__(self, n):
        self.parent = list(range(n))
        self.size = [1] * n
        self.groups = n

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

    def union(self, a, b):
        ra, rb = self.find(a), self.find(b)
        if ra == rb:
            return False                        # already together
        if self.size[ra] < self.size[rb]:       # union by size
            ra, rb = rb, ra
        self.parent[rb] = ra
        self.size[ra] += self.size[rb]
        self.groups -= 1
        return True

d = DSU(8)
for a, b in [(0, 1), (2, 3), (1, 3), (5, 6)]:
    d.union(a, b)
print("groups:", d.groups)
print("0 and 2 connected:", d.find(0) == d.find(2))
print("size of 0's group:", d.size[d.find(0)])
Output
groups: 4
0 and 2 connected: True
size of 0's group: 4

Counting connected components

Start with n groups and union each edge. Every successful union reduces the count by one. This solves "number of friend circles" or "number of islands" style questions in near-linear time, and unlike BFS it also handles edges that arrive over time.

class DSU:
    def __init__(self, n):
        self.parent, self.size, self.groups = list(range(n)), [1] * n, n
    def find(self, x):
        while self.parent[x] != x:
            self.parent[x] = self.parent[self.parent[x]]   # path halving
            x = self.parent[x]
        return x
    def union(self, a, b):
        ra, rb = self.find(a), self.find(b)
        if ra == rb: return False
        if self.size[ra] < self.size[rb]: ra, rb = rb, ra
        self.parent[rb] = ra; self.size[ra] += self.size[rb]; self.groups -= 1
        return True

users = ["ada", "lin", "sam", "kai", "zoe"]
idx = {u: i for i, u in enumerate(users)}
friendships = [("ada", "lin"), ("sam", "kai"), ("lin", "ada")]

d = DSU(len(users))
for a, b in friendships:
    d.union(idx[a], idx[b])
print("friend circles:", d.groups)
Output
friend circles: 3

The loop above uses path halving, a one-pass variant of compression that points every other node to its grandparent. It has the same guarantees and is popular because it is short.

Detecting a cycle

In an undirected graph, if an edge joins two nodes that are already in the same group, adding it closes a loop. union returning False is the signal.

parent = list(range(4))
def find(x):
    while parent[x] != x:
        parent[x] = parent[parent[x]]
        x = parent[x]
    return x

edges = [(0, 1), (1, 2), (2, 3), (3, 1)]
for a, b in edges:
    ra, rb = find(a), find(b)
    if ra == rb:
        print(f"edge {a}-{b} creates a cycle")
        break
    parent[rb] = ra
Output
edge 3-1 creates a cycle

Kruskal's minimum spanning tree

Given towns and the cost of laying cable between pairs, what is the cheapest way to connect every town? Kruskal's algorithm sorts edges by cost and takes each one unless it would create a cycle. Union-find makes the cycle check nearly free, so the total cost is dominated by sorting: O(E log E).

def kruskal(n, edges):
    parent = list(range(n))
    def find(x):
        while parent[x] != x:
            parent[x] = parent[parent[x]]
            x = parent[x]
        return x
    total, chosen = 0, []
    for cost, a, b in sorted(edges):
        ra, rb = find(a), find(b)
        if ra != rb:
            parent[rb] = ra
            total += cost
            chosen.append((a, b, cost))
    return total, chosen

edges = [(4, 0, 1), (1, 1, 2), (3, 0, 2), (2, 2, 3), (5, 1, 3)]
print(kruskal(4, edges))
Output
(6, [(1, 2, 1), (2, 3, 2), (0, 2, 3)])

Limits

Union-find only merges. It cannot efficiently split a group or remove an edge, and it does not give you a path between two nodes, only whether one exists. For those, use a graph with BFS or DFS.

Complexity

Version                      find            union           space
naive (no optimizations)     O(n) worst      O(n) worst      O(n)
union by size/rank only      O(log n)        O(log n)        O(n)
path compression only        O(log n) amortized              O(n)
both                         O(alpha(n)) amortized, about O(1)   O(n)

Components with BFS/DFS: O(V + E) once, but redo work when edges arrive.
Components with DSU:     O(E alpha(V)) and supports edges arriving online.

Recap

  • Union-find tracks disjoint groups; each group is a tree identified by its root.
  • find returns the root; union links two roots.
  • Union by size plus path compression makes operations effectively O(1).
  • Use it to count components, detect cycles in undirected graphs and build minimum spanning trees.
  • It only merges groups; it cannot split them or return paths.
# Write your solution here

Finished reading? Mark this lesson complete to track your progress.

Up next · Lesson 16Segment Trees and Fenwick TreesAnswer range sum and range minimum queries with updates in O(log n) using prefix sums, Fenwick trees (binary indexed trees) and segment trees in Python.