Union-Find (Disjoint Set Union)
Near-constant-time merging and connectivity queries. Path compression, union by size, counting components, detecting cycles, and weighted/offline variants.
The idea
Union-Find maintains a collection of disjoint sets and supports two operations:
find(x)— which set is x in? (returns the set's representative / root)union(a, b)— merge the sets containing a and b.
Each set is a tree of parent pointers; the root represents the set. With two optimizations, both operations run in O(α(n)) amortized — the inverse Ackermann function, which is ≤ 4 for any realistic n. Effectively constant.
The template
class DSU:
def __init__(self, n):
self.parent = list(range(n))
self.size = [1] * n
self.components = 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 # already connected
if self.size[ra] < self.size[rb]:
ra, rb = rb, ra
self.parent[rb] = ra # attach smaller under larger
self.size[ra] += self.size[rb]
self.components -= 1
return True
def connected(self, a, b):
return self.find(a) == self.find(b)
struct DSU {
vector<int> parent, sz;
int components;
DSU(int n) : parent(n), sz(n, 1), components(n) { iota(parent.begin(), parent.end(), 0); }
int find(int x) {
while (parent[x] != x) { parent[x] = parent[parent[x]]; x = parent[x]; }
return x;
}
bool unite(int a, int b) {
a = find(a); b = find(b);
if (a == b) return false;
if (sz[a] < sz[b]) swap(a, b);
parent[b] = a; sz[a] += sz[b];
components--;
return true;
}
};
Why the two optimizations matter
- Union by size/rank: always hang the smaller tree under the larger one. Tree height stays O(log n).
- Path compression: during
find, point nodes (closer) to the root, flattening the tree for future queries.
Without them, a chain of unions can create a linked list and find degrades to O(n).
Recursive full path compression (also common):
def find(self, x):
if self.parent[x] != x:
self.parent[x] = self.find(self.parent[x])
return self.parent[x]
Pattern 1: Counting components
def count_components(n, edges):
dsu = DSU(n)
for a, b in edges:
dsu.union(a, b)
return dsu.components
Pattern 2: Cycle detection in undirected graphs
def find_redundant_connection(edges):
dsu = DSU(len(edges) + 1)
for a, b in edges:
if not dsu.union(a, b):
return [a, b] # a and b were already connected → this edge closes a cycle
Valid tree ⇔ exactly n − 1 edges and no union fails.
Pattern 3: Mapping arbitrary keys to indices
When items are strings or coordinates, assign ids on the fly (or use a dict-based DSU):
class DictDSU:
def __init__(self):
self.parent = {}
def find(self, x):
self.parent.setdefault(x, x)
while self.parent[x] != x:
self.parent[x] = self.parent[self.parent[x]]
x = self.parent[x]
return x
def union(self, a, b):
self.parent[self.find(a)] = self.find(b)
Accounts merge
from collections import defaultdict
def accounts_merge(accounts):
dsu = DictDSU()
owner = {}
for acc in accounts:
name, first = acc[0], acc[1]
for email in acc[1:]:
dsu.union(email, first)
owner[email] = name
groups = defaultdict(list)
for email in owner:
groups[dsu.find(email)].append(email)
return [[owner[root]] + sorted(emails) for root, emails in groups.items()]
Grid as union-find
Map (r, c) to r * C + c. Union adjacent land cells. For rows/columns (Most Stones Removed), union row r with a column node c + 10001 so rows and columns don't collide.
Pattern 4: Online connectivity
Edges arrive one at a time and you must answer after each (Number of Islands II). BFS/DFS would redo work each time; DSU handles each addition in ~O(1).
def num_islands2(m, n, positions):
dsu = DSU(m * n)
land = set()
res, count = [], 0
for r, c in positions:
if (r, c) in land:
res.append(count); continue
land.add((r, c))
count += 1
for nr, nc in ((r + 1, c), (r - 1, c), (r, c + 1), (r, c - 1)):
if (nr, nc) in land and dsu.union(r * n + c, nr * n + nc):
count -= 1
res.append(count)
return res
Pattern 5: Offline queries sorted by threshold
“Is there a path from u to v using only edges with weight < limit?” for many queries: sort edges and queries by weight/limit, and add edges incrementally while answering queries in order. Each query is then a single find.
def distance_limited_paths_exist(n, edge_list, queries):
edge_list.sort(key=lambda e: e[2])
order = sorted(range(len(queries)), key=lambda i: queries[i][2])
dsu, res, j = DSU(n), [False] * len(queries), 0
for i in order:
u, v, limit = queries[i]
while j < len(edge_list) and edge_list[j][2] < limit:
dsu.union(edge_list[j][0], edge_list[j][1])
j += 1
res[i] = dsu.find(u) == dsu.find(v)
return res
Same idea: Swim in Rising Water (add cells by elevation), and Kruskal's MST (next topic).
Weighted union-find
Store with each node its relation to its parent (a ratio, an offset, a parity). find accumulates relations along the path while compressing.
class WeightedDSU: # weight[x] = value(x) / value(parent[x])
def __init__(self):
self.parent, self.weight = {}, {}
def find(self, x):
if x not in self.parent:
self.parent[x], self.weight[x] = x, 1.0
if self.parent[x] != x:
p = self.parent[x]
root = self.find(p)
self.weight[x] *= self.weight[p] # now relative to root
self.parent[x] = root
return self.parent[x]
def union(self, a, b, ratio): # a / b = ratio
ra, rb = self.find(a), self.find(b)
if ra != rb:
self.parent[ra] = rb
self.weight[ra] = ratio * self.weight[b] / self.weight[a]
def query(self, a, b):
if a not in self.parent or b not in self.parent or self.find(a) != self.find(b):
return -1.0
return self.weight[a] / self.weight[b]
Parity DSU (weight = 0/1 “same/different”) checks bipartiteness online.
Union-Find vs BFS/DFS
| Situation | Better choice |
|---|---|
| static graph, one-time component count | either; DFS is fine |
| edges added over time, queries interleaved | DSU |
| need actual paths or distances | BFS/DFS |
| edges removed over time | reverse time and add them (offline), or other structures |
| directed reachability | BFS/DFS (DSU is for undirected equivalence) |
Complexity
O(α(n)) amortized per operation with both optimizations; O(n) space.
Common mistakes
- Comparing
parent[a] == parent[b]instead offind(a) == find(b). - Forgetting to call
findbefore merging (merging non-roots corrupts sizes). - Counting components by counting distinct
parent[i]values without callingfindon each. - Using DSU for directed-graph problems where direction matters.
Practice problems
Ordered roughly from warm-up to hard. Try each one for 25–40 minutes before peeking at the key idea; if you needed the hint, mark it “review” and redo it in a few days.
Click the box to cycle: solved ✓ → needs review ↺ → not started. Links open LeetCode.