A plain union-find, or disjoint-set union (DSU), answers one question fast: are these two elements in the same group? With union by size and path compression, a long sequence of operations costs almost constant time each. The core structure is explained in union-find from first principles; this page assumes it and starts where it stops.
Real problems ask more than membership. They ask how two elements relate inside a group, what the group's total is, whether it contains a cycle, how to move one element elsewhere, which slot is free next, and what the groups looked like before an edge was removed. Each question is answered by a small, specific change to the same forest of parent pointers. This page builds five such variants, runs each one on a worked example, and states what each change costs. All code is Python, and every output shown was produced by running it.
Choosing a variant
Potential DSU: relations inside a group
Many constraint problems state relations between pairs: service auth starts 12 ms after api, account A holds 3 times the balance of B, node u and node v have opposite colours. Store, for every node, its value relative to its parent: diff[x] = value(x) - value(parent[x]). The value of a node relative to its root is the sum along its path. A union records a new relation; if the two nodes are already in the same group, the relation is either consistent with what is known or a contradiction.
class WeightedDSU:
# Each node stores diff[x] = value(x) - value(parent[x]).
def __init__(self, n):
self.parent = list(range(n)); self.size = [1] * n; self.diff = [0] * n
def find(self, x):
path = []
while self.parent[x] != x:
path.append(x); x = self.parent[x]
root = x
# Re-point from the node nearest the root outwards, so each parent's
# diff is already relative to the root when its child folds it in.
for node in reversed(path):
p = self.parent[node]
if p != root:
self.diff[node] += self.diff[p]
self.parent[node] = root
return root
def union(self, a, b, w):
# Record value(a) - value(b) = w. Returns False on contradiction.
ra, rb = self.find(a), self.find(b)
da, db = self.diff[a], self.diff[b]
if ra == rb:
return da - db == w
if self.size[ra] < self.size[rb]:
ra, rb, da, db, w = rb, ra, db, da, -w
self.parent[rb] = ra
self.diff[rb] = da - db - w
self.size[ra] += self.size[rb]
return True
def delta(self, a, b):
if self.find(a) != self.find(b):
return None
return self.diff[a] - self.diff[b]Six start-up offsets between five services, in milliseconds, arrive from different log sources. The last one disagrees with the others:
auth - api = 12: ok
db - auth = 30: ok
cache - api = 4: ok
db - cache = 38: ok
queue - db = 5: ok
queue - api = 40: CONTRADICTION
queue - api = 47
cache - auth = -8The fourth fact is checked, not stored: db - cache must be 42 minus 4, which it is. The sixth contradicts the derived 47, so the source that produced it is wrong. The order of the compression loop is the part people get wrong. Folding offsets from the outermost node inwards adds a parent's offset before that parent has been made relative to the root, and the result is silently incorrect only on paths longer than two. For parity problems such as bipartite checks, replace the sum with XOR; for ratios, use multiplication and division, with exact rationals rather than floats.
Aggregates at the root, and cycle detection
Any value that combines associatively can live at the root and be merged during union: size, sum, minimum, maximum, a bitmask of labels. Counting edges as well as nodes adds one more answer for free. A connected component with as many edges as nodes contains a cycle.
def add_edge(self, a, b):
ra, rb = self.find(a), self.find(b)
if ra == rb:
self.edges[ra] += 1 # this edge closes a cycle
return
if self.size[ra] < self.size[rb]:
ra, rb = rb, ra
self.parent[rb] = ra
self.size[ra] += self.size[rb]
self.edges[ra] += self.edges[rb] + 1
self.total[ra] += self.total[rb]
self.lo[ra] = min(self.lo[ra], self.lo[rb])
def has_cycle(self, x):
r = self.find(x)
return self.edges[r] >= self.size[r]Six nodes with values 5, 3, 8, 1, 7, 2 and edges 0-1, 1-2, 3-4, 2-0:
component of 0: size=3 sum=16 min=3 cycle=True
component of 3: size=2 sum=8 min=1 cycle=False
component of 5: size=1 sum=2 min=2 cycle=FalseTwo rules keep this correct. Read aggregates only at the root, after find; values stored at non-root nodes are stale by design. And only merge values whose combine operation does not need to be undone. A minimum cannot be decreased back when a member leaves, which is why the next two variants exist. The same idea at tree scale, merging per-subtree sets cheaply, is small-to-large merging.
Moving an element with a fresh node
A DSU cannot remove an element from its group, because other nodes may hang underneath it. The standard trick is to stop addressing elements by node. Keep a handle[x] that maps element x to its current node. To move x, decrement the old group's live count, allocate a brand-new node for x, and attach that node to the target group. The old node stays in the forest as a dead leaf, still holding up whatever hangs below it.
def move(self, x, y):
rx = self.find(self.handle[x])
self.count[rx] -= 1 # the old node stays as a dead leaf
v = self._new_node()
self.handle[x] = v
ry = self.find(self.handle[y])
self.parent[v] = ry
self.count[ry] += 1After unions {0,1,2} and {3,4}, moving 2 into 4's group prints same(0,2): False same(2,4): True count(set of 0): 2. Each move adds one node, so memory grows with moves rather than elements. Rebuild from the live handles when dead nodes exceed live ones. Aggregates that can be subtracted, such as count and sum, follow the element; minimum and maximum cannot.
Next-free-slot union-find
Here the groups are runs of used slots, and each points to the first free slot after it. nxt[i] is i itself while slot i is free. Claiming a slot points it at the next one, so a later find skips the whole run. A sentinel slot n means nothing is free.
class NextFree:
def __init__(self, n):
self.nxt = list(range(n + 1)) # slot n is a sentinel
def find(self, i):
root = i
while self.nxt[root] != root:
root = self.nxt[root]
while self.nxt[i] != root: # full path compression
self.nxt[i], i = root, self.nxt[i]
return root
def take(self, i):
s = self.find(i)
if s < len(self.nxt) - 1:
self.nxt[s] = s + 1
return sWith eight slots, the claims take(3), take(3), take(3), take(6), take(6), take(6), take(0) return [3, 4, 5, 6, 7, 8, 0]; the 8 is the sentinel, meaning slot 6 onwards was full. This pattern assigns jobs to the first free day at or after a deadline, allocates ports or IDs, and paints intervals offline: process paint operations from last to first and skip already-painted cells, so each cell is written once. Linking always goes rightwards, so union by size is unavailable, and path compression alone gives amortised O(log n) per operation; in practice runs collapse quickly.
Rollback and offline dynamic connectivity
Path compression rewrites pointers on reads, and those writes cannot be undone cheaply. A rollback DSU drops compression, keeps union by size so every find is O(log n), and pushes each successful union onto a stack. Undoing pops the stack and restores one pointer and one size.
That is enough to handle edges that are added and removed in any order, as long as all operations are known in advance. Each edge is alive for an interval of time. Insert each interval into a segment tree over the timeline, so it lands on O(log m) nodes. Then walk the tree depth-first: on entering a node, union its edges; at a leaf, answer that moment's query; on leaving, roll back to the mark taken on entry.
def walk(node, lo, hi):
mark = len(dsu.stack)
for u, v in tree[node]:
dsu.union(u, v)
if hi - lo == 1:
if ops[lo][0] == "count":
out[lo] = dsu.comps
else:
mid = (lo + hi) // 2
walk(2 * node, lo, mid); walk(2 * node + 1, mid, hi)
dsu.rollback_to(mark)On five nodes, the operations add 0-1, count, add 1-2, add 3-4, count, remove 1-2, add 2-3, count, remove 0-1, remove 3-4, count print component counts: [4, 2, 2, 4]. Each edge is unioned O(log m) times at O(log n) each, so the total is O(m log m log n). An online version with arbitrary deletions needs a much heavier structure; if queries can be batched, this one is far simpler.
Persistent and concurrent variants
Two further variants come up less often. A persistent DSU keeps every past version queryable, typically by storing the parent and size arrays in a persistent array or persistent segment tree and again giving up path compression; each operation then costs an extra logarithmic factor. Concurrent union-find, used in parallel connected-components and graph clustering, replaces pointer writes with compare-and-swap and makes linking order deterministic so two threads cannot create a cycle. Both are worth reaching for only after measuring that a single-threaded or offline structure is the bottleneck.
Failure modes
- Compression order in the potential DSU. Folding offsets outermost-first gives wrong deltas on paths of three or more. Test with a long chain built by unions in both directions.
- Compression in a rollback DSU. Adding path halving to a rollback DSU makes undo restore the wrong parents. Keep
findread-only. - Reading aggregates at a non-root. Stale values look plausible. Always find first.
- Recursion depth. Recursive find in Python or on the JVM overflows on long chains before compression flattens them. Use the iterative form shown here.
- Floats in relation weights. Ratio constraints with floats report contradictions from rounding. Use integers, rationals, or logarithms with a tolerance you have chosen deliberately.
- Unbounded growth in the movable DSU. Every move allocates. Track dead nodes and rebuild.
Trade-offs
| Variant | Extra state | Find cost | Gives up |
|---|---|---|---|
| Potential | one value per node | near O(1) amortised | nothing, but compression must fold values |
| Aggregate | values at each root | near O(1) amortised | removal of members |
| Movable | handle per element, dead nodes | near O(1) amortised | memory proportional to moves |
| Next-free | none | amortised O(log n) | union by size |
| Rollback | undo stack | O(log n) worst case | path compression |
| Offline connectivity | segment tree of edge intervals | O(log m log n) per edge | online answers |
If the problem is a minimum spanning tree, union-find is the cycle check inside Kruskal's algorithm; compare it with Prim on dense graphs. If you need to know which edges hold a component together, union-find has thrown that information away; use bridges and articulation points instead.
What to do next
- Write down the question your problem asks of a group: membership, relation, aggregate, movement, next free slot, or history. That picks the variant.
- Implement the potential DSU and test it against a brute-force solver on random small constraint sets, including contradictions.
- Replace any recursive find in production code with an iterative one.
- If you undo unions, remove path compression and assert after each rollback that sizes sum to n.
- For connectivity with deletions, check whether queries can be collected first. If so, use the segment tree over time before considering an online structure.
- Benchmark with adversarial input, such as long chains and repeated moves, not only random graphs.