Union-Find

Union-Find (DSU)

Disjoint set union template with path compression — provinces, redundant edges, valid tree, accounts merge.

When to reach for Union-Find

Union-Find (Disjoint Set Union) maintains a forest of trees. find(x) returns the root; union(a,b) links roots. With path compression + union by rank, almost every op is amortized nearly O(1) (inverse Ackermann).

Analogy: merging friend circles

Each person starts alone. A friendship is a union: the smaller circle joins the larger one’s “leader.” Asking “same circle?” is find on both and compare leaders. Path compression is everyone pointing straight at the current mayor after you ask — next ask is faster.

Template: DSU with path compression + rank

Copy this into the interview first, then specialize.

class DSU:
    def __init__(self, n):
        self.p = list(range(n))
        self.rank = [0] * n
        self.components = n

    def find(self, x):
        while self.p[x] != x:
            self.p[x] = self.p[self.p[x]]  # path compress
            x = self.p[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.rank[ra] < self.rank[rb]:
            ra, rb = rb, ra
        self.p[rb] = ra
        if self.rank[ra] == self.rank[rb]:
            self.rank[ra] += 1
        self.components -= 1
        return True
class DSU {
    int[] p, rank;
    int components;
    DSU(int n) {
        p = new int[n]; rank = new int[n]; components = n;
        for (int i = 0; i < n; i++) p[i] = i;
    }
    int find(int x) {
        while (p[x] != x) { p[x] = p[p[x]]; x = p[x]; }
        return x;
    }
    boolean union(int a, int b) {
        int ra = find(a), rb = find(b);
        if (ra == rb) return false;
        if (rank[ra] < rank[rb]) { int t = ra; ra = rb; rb = t; }
        p[rb] = ra;
        if (rank[ra] == rank[rb]) rank[ra]++;
        components--;
        return true;
    }
}

Edges & gotchas

What to say in the first 60 seconds

Interviewers listen for the amortized bound. Say “nearly O(1) per op with path compression + rank — inverse Ackermann” once, then move on. Don’t derive Ackermann on the whiteboard.

Complexity you should quote

If you only remember one implementation detail under pressure: path compression on find. Rank is the second lever — together they keep trees flat.

Pattern bank (5 questions)

Q: Number of Provinces — n cities; isConnected[i][j] == 1 means an edge. Return # of provinces.

def find_circle_num(is_connected):
    n = len(is_connected)
    dsu = DSU(n)
    for i in range(n):
        for j in range(i + 1, n):
            if is_connected[i][j]:
                dsu.union(i, j)
    return dsu.components
int findCircleNum(int[][] isConnected) {
    int n = isConnected.length;
    DSU dsu = new DSU(n);
    for (int i = 0; i < n; i++)
        for (int j = i + 1; j < n; j++)
            if (isConnected[i][j] == 1) dsu.union(i, j);
    return dsu.components;
}

Q: Redundant Connection — undirected edges forming a tree + one extra edge; return the edge that creates the cycle (last such in input).

def find_redundant_connection(edges):
    dsu = DSU(len(edges) + 1)  # 1-indexed nodes
    for u, v in edges:
        if not dsu.union(u, v):
            return [u, v]
    return []
int[] findRedundantConnection(int[][] edges) {
    DSU dsu = new DSU(edges.length + 1);
    for (int[] e : edges)
        if (!dsu.union(e[0], e[1])) return e;
    return new int[]{};
}

Q: Graph Valid Tree — n nodes and edges; return true if they form a valid tree.

def valid_tree(n, edges):
    if len(edges) != n - 1:
        return False
    dsu = DSU(n)
    for u, v in edges:
        if not dsu.union(u, v):
            return False
    return dsu.components == 1
boolean validTree(int n, int[][] edges) {
    if (edges.length != n - 1) return false;
    DSU dsu = new DSU(n);
    for (int[] e : edges)
        if (!dsu.union(e[0], e[1])) return false;
    return dsu.components == 1;
}

Q: Accounts Merge — list of accounts [name, emails…]; merge accounts that share any email; return sorted emails per name.

from collections import defaultdict

def accounts_merge(accounts):
    n = len(accounts)
    dsu = DSU(n)
    email_to_id = {}
    for i, acc in enumerate(accounts):
        for email in acc[1:]:
            if email in email_to_id:
                dsu.union(i, email_to_id[email])
            else:
                email_to_id[email] = i
    groups = defaultdict(list)
    for email, i in email_to_id.items():
        groups[dsu.find(i)].append(email)
    return [[accounts[i][0], *sorted(emails)] for i, emails in groups.items()]
// Sketch: Map email→id, DSU.union shared emails, group by find(id), sort.
List<List<String>> accountsMerge(List<List<String>> accounts) {
    int n = accounts.size();
    DSU dsu = new DSU(n);
    Map<String, Integer> emailToId = new HashMap<>();
    for (int i = 0; i < n; i++) {
        for (int j = 1; j < accounts.get(i).size(); j++) {
            String email = accounts.get(i).get(j);
            if (emailToId.containsKey(email)) dsu.union(i, emailToId.get(email));
            else emailToId.put(email, i);
        }
    }
    Map<Integer, TreeSet<String>> groups = new HashMap<>();
    for (var e : emailToId.entrySet())
        groups.computeIfAbsent(dsu.find(e.getValue()), k -> new TreeSet<>()).add(e.getKey());
    List<List<String>> out = new ArrayList<>();
    for (var e : groups.entrySet()) {
        List<String> row = new ArrayList<>();
        row.add(accounts.get(e.getKey()).get(0));
        row.addAll(e.getValue());
        out.add(row);
    }
    return out;
}

Q: Earliest Moment When Everyone Become Friends — logs [timestamp, x, y]; return earliest time all n people are connected, else −1.

def earliest_acq(logs, n):
    logs = sorted(logs)  # by timestamp
    dsu = DSU(n)
    for t, a, b in logs:
        dsu.union(a, b)
        if dsu.components == 1:
            return t
    return -1
int earliestAcq(int[][] logs, int n) {
    Arrays.sort(logs, Comparator.comparingInt(a -> a[0]));
    DSU dsu = new DSU(n);
    for (int[] log : logs) {
        dsu.union(log[1], log[2]);
        if (dsu.components == 1) return log[0];
    }
    return -1;
}

Pattern drill: 45 minutes

← Lattice