그래프에서 모든 정점을 연결하면서 전체 간선 가중치의 합이 최소가 되는 트리를 말합니다.
'스패닝 트리'는 모든 정점을 포함하는 부분 트리이며,
'최소 신장 트리(MST)'는 그 중에서 가중치의 합이 최소인 트리입니다.
Prim 알고리즘은 하나의 정점에서 출발하여 가장 저렴하게 연결되는 인접 정점을 하나씩 확장해 나가는 방식입니다.
처음에는 하나의 정점만 가지고 있다가, 그 정점과 연결된 가장 비용이 낮은 간선을 골라 새로운 정점을 포함시킵니다. 이후 그 덩어리에서 다시 가장 저렴한 간선을 골라 연결해 나갑니다.
덩어리를 조금씩 확장해 나가는 탐색형 알고리즘입니다.
다만, 이 방식은 그래프가 희소한 경우 우선순위 큐에 불필요한 후보 정점들이 많이 쌓이게 되어 시간 낭비가 생길 수 있습니다.
Kruskal 알고리즘은 간선 가중치를 기준으로 전체 그래프를 정렬한 뒤, 가장 가벼운 간선부터 하나씩 연결해 나가는 방식입니다. 연결 과정에서 사이클이 생기지 않도록 유니온 파인드(Disjoint Set) 구조를 사용합니다.
서로 다른 소규모 트리(또는 집합)를
가장 저렴한 간선으로 병합하여
전체를 하나의 거대한 트리로 만드는 방식입니다.
Prim처럼 한 지점에서 출발하지 않고, 처음부터 모든 트리가 독립적인 상태에서 시작하며 간선 중심으로 필요한 연결만 정확히 수행합니다.
두가지 방식 다 엄밀한 증명을 위해 cut proprety가 필요하며 이는 각자 찾아보거나 증명해 보시면 좋을 것 같습니다
이번 실험에서 사용한 그래프는 정점이 1000개, 간선이 3000개인 희소 그래프입니다. 이런 구조에서는 Prim의 큐에 불필요한 간선 후보들이 많이 쌓이게 되어 성능상 손해가 발생합니다.
반면 Kruskal 알고리즘은 애초에 간선을 정렬한 뒤, 정확히 필요한 간선만 선택해 연결하므로 더 효율적이었습니다.
[Kruskal-반복] MST: 191823, Time: 0.003144초
[Kruskal-재귀] MST: 191823, Time: 0.002988초
[Prim ] MST: 191823, Time: 0.005745초
| 항목 | Prim (덩어리 확장) | Kruskal (트리 병합) |
|---|---|---|
| 출발 방식 | 한 정점에서 시작 | 모든 간선 기준 |
| 확장 흐름 | 좌측에서 우측으로 덩어리 확장 | 각지 마을을 병합 후 하나의 트리로 |
| 핵심 개념 | 우선순위 큐 + 방문 체크 | 간선 정렬 + 유니온 파인드 |
| 희소 그래프에서 | 느릴 수 있음 | 빠름 |
| 밀집 그래프에서 | 효율적일 수 있음 | 오히려 느릴 수 있음 |
| 인상적인 점 | 탐색 중심 | 병합 중심 |
Prim 알고리즘은 덩어리를 하나씩 확장해 나가는 방식이며,
Kruskal 알고리즘은 작은 트리들을 가장 저렴한 간선으로 병합해 나가는 방식입니다.
이번 실험을 통해,
그래프 구조에 따라 어떤 알고리즘이 더 적합할지 판단하는 감각을 얻을 수 있었습니다.
특히 희소 그래프에서는 Kruskal이 훨씬 유리하다는 것을 직접 체감했습니다.
# ====================== Union-Find + kruskal =====================
def find(x):
root = x
while parent[root] != root:
root = parent[root]
while x != root:
next_x = parent[x]
parent[x] = root
x = next_x
return root
def find_recursion(x):
if parent[x] != x:
parent[x] = find_recursion(parent[x])
return parent[x]
def union(x, y, use_recursion=False):
if use_recursion:
parent_x = find_recursion(x)
parent_y = find_recursion(y)
else:
parent_x = find(x)
parent_y = find(y)
if parent_x != parent_y:
parent[parent_y] = parent_x
def kruskal(que, use_recursion=False):
total = 0
while que:
dist, x, y = heapq.heappop(que)
if use_recursion:
if find_recursion(x) == find_recursion(y):
continue
total += dist
union(x, y, use_recursion=True)
else:
if find(x) == find(y):
continue
total += dist
union(x, y)
return total
# ====================== prim =====================
def prim(n,m,graph):
visited=[False]*(n+1)
total_dist=0
que=[(0,1)]#dist,start_node
while que:
dist, node=heapq.heappop(que)
if visited[node]==True:
continue
visited[node]=True
total_dist+=dist
for next_node,cost in graph[node]:
heapq.heappush(que,(cost,next_node))
return total_dist