Graph Algorithms contents
Minimum Spanning Tree: Kruskal's Algorithm
Connect all vertices at the smallest total cost by sorting edges and joining components with a disjoint set union, in O(m log m).
Read first: Disjoint Set Union (Union-Find)
Given a connected, weighted, undirected graph, a spanning tree is a subset of edges that connects all vertices without a cycle (so it has exactly edges). A minimum spanning tree (MST) is a spanning tree of the smallest possible total weight. Typical use: connect all cities with the cheapest road network, cable a set of offices, cluster points.
Kruskal's algorithm
A greedy algorithm:
- Sort all edges by weight, from smallest to largest.
- Go through them in that order; add an edge if it connects two vertices that are not yet connected (it won't form a cycle), otherwise skip it.
The correctness rests on the cut property: for any partition of the vertices into two groups, the lightest edge crossing the partition belongs to some MST. When Kruskal considers an edge that joins two different current components, that edge is the lightest one leaving one of them, so it is safe to take.
To test "are and already connected?" quickly, use a disjoint set union.
def kruskal(n, edges):
"""edges: list of (weight, u, v). Return (total_weight, list_of_chosen_edges)."""
parent = list(range(n))
size = [1] * n
def find(x):
while parent[x] != x:
parent[x] = parent[parent[x]]
x = parent[x]
return x
total = 0
chosen = []
for w, u, v in sorted(edges):
ru, rv = find(u), find(v)
if ru == rv:
continue # would close a cycle
if size[ru] < size[rv]:
ru, rv = rv, ru
parent[rv] = ru
size[ru] += size[rv]
total += w
chosen.append((u, v, w))
if len(chosen) == n - 1:
break
return total, chosen
edges = [(7, 0, 1), (5, 0, 3), (8, 1, 2), (9, 1, 3), (7, 1, 4), (5, 2, 4), (15, 3, 4), (6, 3, 5), (8, 4, 5), (9, 4, 6), (11, 5, 6)]
total, chosen = kruskal(7, edges)
assert total == 39 and len(chosen) == 6Complexity. Sorting costs ; the DSU operations cost . The sort dominates.
If the graph is disconnected, the algorithm returns a minimum spanning forest: fewer than edges. Check len(chosen) to detect that.
total, chosen = kruskal(4, [(1, 0, 1), (2, 2, 3)])
assert len(chosen) == 2 and total == 3 # a forest with 2 components: 4 - 2 edgesChecking against a brute force
For small graphs enumerate all subsets of edges, keep those that form a spanning tree, take the cheapest:
from itertools import combinations
import random
def is_spanning_tree(n, chosen):
parent = list(range(n))
def find(x):
while parent[x] != x:
parent[x] = parent[parent[x]]
x = parent[x]
return x
for _, u, v in chosen:
ru, rv = find(u), find(v)
if ru == rv:
return False
parent[ru] = rv
return True
def mst_brute(n, edges):
best = None
for subset in combinations(edges, n - 1):
if is_spanning_tree(n, subset):
w = sum(e[0] for e in subset)
best = w if best is None else min(best, w)
return best
random.seed(2)
for _ in range(200):
n = random.randint(2, 6)
# ensure connectivity with a random path, then add extra edges
perm = list(range(n))
random.shuffle(perm)
es = [(random.randint(1, 9), perm[i], perm[i + 1]) for i in range(n - 1)]
es += [(random.randint(1, 9), random.randrange(n), random.randrange(n)) for _ in range(random.randint(0, 5))]
es = [(w, u, v) for w, u, v in es if u != v]
assert kruskal(n, es)[0] == mst_brute(n, es)Properties and applications
- If all edge weights are distinct, the MST is unique. Otherwise there may be several with the same total weight.
- The MST minimizes the maximum edge on the path between any two vertices (the bottleneck path).
- Single-linkage clustering: stop Kruskal when components remain to obtain clusters.
- Maximum spanning tree: sort edges in decreasing order.
- Second-best MST: try replacing each non-tree edge into the tree (needs LCA with path maxima).
def k_clusters(n, edges, k):
"""Split n points into k groups by stopping Kruskal early; return the group id of each vertex."""
parent = list(range(n))
def find(x):
while parent[x] != x:
parent[x] = parent[parent[x]]
x = parent[x]
return x
components = n
for w, u, v in sorted(edges):
if components == k:
break
ru, rv = find(u), find(v)
if ru != rv:
parent[ru] = rv
components -= 1
roots = sorted({find(x) for x in range(n)})
return [roots.index(find(x)) for x in range(n)]
groups = k_clusters(6, [(1, 0, 1), (1, 1, 2), (2, 3, 4), (10, 2, 3), (1, 4, 5)], 2)
assert groups[0] == groups[1] == groups[2] and groups[3] == groups[4] == groups[5] and groups[0] != groups[3]For dense graphs, Prim's algorithm can be better.
Practice problems
- SPOJ - Koicost
- SPOJ - MaryBMW
- Codechef - Fullmetal Alchemist
- Codeforces - Edges in MST
- UVA 12176 - Bring Your Own Horse
- UVA 10600 - ACM Contest and Blackout
- UVA 10724 - Road Construction
- Hackerrank - Roads in HackerLand
- UVA 11710 - Expensive subway
- Codechef - Chefland and Electricity
- UVA 10307 - Killing Aliens in Borg Maze
- Codeforces - Flea
- Codeforces - Igon in Museum
- Codeforces - Hongcow Builds a Nation
- UVA - 908 - Re-connecting Computer Sites
- UVA 1208 - Oreon
- UVA 1235 - Anti Brute Force Lock
- UVA 10034 - Freckles
- UVA 11228 - Transportation system
- UVA 11631 - Dark roads
- UVA 11733 - Airports
- UVA 11747 - Heavy Cycle Edges
- SPOJ - Blinet
- SPOJ - Help the Old King
- Codeforces - Hierarchy
- SPOJ - Modems
- CSES - Road Reparation
- CSES - Road Construction