Grow a Spanning Tree with Prim
Checking your account…
Sign in to save your code and progress across devices. The lesson and problem statement remain public.
Loading the interactive Practice workspace.If it does not appear, the problem and learning material remain readable, but browser execution is unavailable.Reload Practice workspace
Problem
Implement prim_spanning_weight(node_count, edges). Edges are undirected [u,v,weight]. Return the MST weight or -1 if disconnected.
Starter code
def prim_spanning_weight(node_count, edges):
passTest cases
frontier-choice
{
"args": [
4,
[
[
0,
1,
4
],
[
0,
2,
1
],
[
2,
1,
2
],
[
1,
3,
1
],
[
2,
3,
5
]
]
]
}Expected: 4
disconnected
{
"args": [
3,
[
[
0,
1,
1
]
]
]
}Expected: -1
Wizard outline
- Step 1: Handle an already complete tree
Return zero for a graph with at most one node. Prim needs no growth step for an empty or one-node graph.
- Step 2: Take edges directly from node 0
Build undirected adjacency and total an initial star frontier. Prim grows from one visited seed through crossing edges.
- Step 3: Grow through the cheapest crossing edge
Use a min-heap to expand the visited component until all nodes are reached. Prim must continually compare every edge crossing out of the current tree.
Footguns and prerequisites
- Adding the weight of a stale edge double-counts a vertex.
- Ending with an empty frontier does not imply all vertices were reached.
- trees and graphs
Reviewed references
Recommended approach and implementation
Grow from node zero using the lightest edge crossing from visited to unvisited vertices.
Why it works: At each step the heap minimum crossing edge is safe by the cut property. Adding its new endpoint preserves a tree; reaching every vertex yields an MST.
def prim_spanning_weight(node_count, edges):
import heapq
if node_count <= 1:
return 0
graph = [[] for _ in range(node_count)]
for left, right, weight in edges:
graph[left].append((weight, right)); graph[right].append((weight, left))
visited = set()
frontier = [(0, 0)]
total = 0
while frontier and len(visited) < node_count:
weight, node = heapq.heappop(frontier)
if node in visited:
continue
visited.add(node); total += weight
for next_weight, neighbor in graph[node]:
if neighbor not in visited:
heapq.heappush(frontier, (next_weight, neighbor))
return total if len(visited) == node_count else -1