그래프가 주어졌을 때, 해당 그래프에서 최소 신장 트리(Minimum Spanning Tree, MST)를 찾아야 하는 경우가 자주 있습니다. MST란 가중치 그래프의 부분 집합으로, 모든 정점이 포함되어 있고 서로 연결되어 있으며 사이클(순환)이 존재하지 않는 트리를 의미합니다. '최소'라는 이름이 붙은 이유는 MST를 이루는 간선 가중치의 합이 그래프에서 가능한 모든 신장 트리 중 가장 작기 때문입니다.
이 글에서는 프림(Prim)의 MST 알고리즘을 사용하여 주어진 그래프에서 MST의 총 간선 가중치 합을 구하는 방법을 살펴보겠습니다.
문제 예시
정점의 개수 n = 4, 시작 정점 s = 3이고, 간선 정보가 다음과 같다고 가정해 보겠습니다.
- (1, 2, 5) — 정점 1과 정점 2를 잇는 가중치 5의 간선
- (1, 3, 5) — 정점 1과 정점 3을 잇는 가중치 5의 간선
- (2, 3, 7) — 정점 2와 정점 3을 잇는 가중치 7의 간선
- (1, 4, 4) — 정점 1과 정점 4를 잇는 가중치 4의 간선
이 경우 출력값은 14가 됩니다. 이 그래프의 MST는 다음 간선들로 구성됩니다.
- 정점 4 ↔ 정점 1 (가중치 4)
- 정점 1 ↔ 정점 2 (가중치 5)
- 정점 1 ↔ 정점 3 (가중치 5)
MST의 총 간선 가중치 합은 4 + 5 + 5 = 14입니다.
풀이 접근 방식
프림 알고리즘은 하나의 시작 정점에서 출발하여, 아직 트리에 포함되지 않은 정점 중 가장 낮은 가중치로 연결되는 정점을 반복적으로 추가하는 방식으로 동작합니다. 해결 과정은 다음과 같습니다.
- mst_find() 함수를 정의합니다. 인자로 그래프 G와 시작 정점 s를 받습니다.
- distance := 크기가 |G|인 리스트, 초기값은 양의 무한대(float("inf"))
- distance[s] := 0 으로 설정
- itr := 크기가 |G|인 방문 여부 리스트, 초기값은 False
- c := 0 (MST의 총 가중치 누적 변수)
- 무한 루프를 돌며 다음을 수행합니다.
- min_weight := 무한대, m_idx := -1 로 초기화
- 모든 정점 i를 순회하면서 itr[i]가 False인 정점 중 distance[i]가 가장 작은 정점을 찾습니다.
- m_idx가 -1이면 더 이상 방문할 정점이 없으므로 루프를 종료합니다.
- c에 min_weight를 더하고, itr[m_idx]를 True로 설정하여 해당 정점을 트리에 포함시킵니다.
- G[m_idx]에 연결된 모든 인접 정점 i에 대해 distance[i]를 min(distance[i], j)로 갱신합니다. 여기서 j는 간선의 가중치입니다.
- 루프가 끝나면 c를 반환합니다. 이것이 MST의 총 가중치입니다.
- solve() 함수에서 그래프 G를 생성하고, 각 간선 (u, v, w)에 대해 양방향으로 가중치를 저장합니다. 동일한 두 정점 사이에 여러 간선이 있다면 더 작은 가중치만 유지합니다.
- 마지막으로 mst_find(G, s)를 호출하여 결과를 반환합니다.
구현 예제
다음 파이썬 코드를 통해 위 접근 방식을 더 잘 이해할 수 있습니다.
def mst_find(G, s):
distance = [float("inf")] * len(G)
distance[s] = 0
itr = [False] * len(G)
c = 0
while True:
min_weight = float("inf")
m_idx = -1
for i in range(len(G)):
if itr[i] == False:
if distance[i] < min_weight:
min_weight = distance[i]
m_idx = i
if m_idx == -1:
break
c += min_weight
itr[m_idx] = True
for i, j in G[m_idx].items():
distance[i] = min(distance[i], j)
return c
def solve(n, edges, s):
G = {i: {} for i in range(n)}
for item in edges:
u = item[0]
v = item[1]
w = item[2]
u -= 1
v -= 1
try:
min_weight = min(G[u][v], w)
G[u][v] = min_weight
G[v][u] = min_weight
except KeyError:
G[u][v] = w
G[v][u] = w
return mst_find(G, s)
print(solve(4, [(1, 2, 5), (1, 3, 5), (2, 3, 7), (1, 4, 4)], 3))입력
4, [(1, 2, 5), (1, 3, 5), (2, 3, 7), (1, 4, 4)], 3
출력
14
마무리
위 코드는 프림 알고리즘의 기본 형태로, 시간 복잡도는 O(V²)입니다. 정점 개수가 많은 그래프의 경우 우선순위 큐(힙)를 활용하면 O(E log V)까지 성능을 개선할 수 있으니 참고하시기 바랍니다.