문제 소개
어떤 나라가 N개의 노드와 N-1개의 간선으로 이루어진 트리(tree) 구조로 표현되어 있다고 가정해 보겠습니다. 각 노드는 하나의 마을을 나타내고, 각 간선은 마을 사이를 잇는 도로를 의미합니다. 크기가 N-1인 두 리스트 source와 dest가 주어지며, i번째 도로는 source[i]와 dest[i]를 연결하고 모든 도로는 양방향입니다. 또한 크기가 N인 population 리스트에는 population[i] 값으로 i번째 마을의 인구가 담겨 있습니다.
우리는 몇 개의 마을을 도시로 승격시키려고 합니다. 이때 다음 조건들을 반드시 지켜야 합니다.
- 임의의 두 도시가 서로 인접해서는 안 됩니다.
- 마을에 인접한 모든 노드는 반드시 도시여야 합니다. 즉, 모든 도로는 항상 마을과 도시를 연결해야 합니다.
목표는 위 조건을 만족하면서 승격된 모든 도시의 인구 합을 최대화하는 것입니다.
예를 들어 입력이 다음과 같다고 해 보겠습니다.
source = [2, 2, 1, 1] dest = [1, 3, 4, 0] population = [6, 8, 4, 3, 5]
이 경우 출력은 15입니다. 0번, 2번, 4번 마을을 도시로 승격하면 인구가 6 + 4 + 5 = 15가 되기 때문입니다.
풀이 접근 방법
핵심 아이디어는 트리가 이분 그래프(bipartite graph)라는 점입니다. '모든 도로가 마을과 도시를 연결해야 한다'는 조건은 곧 노드를 두 그룹으로 나누어 번갈아 칠하는 2-색칠(2-coloring)과 같습니다. 트리에서는 이런 색칠이 정확히 두 가지뿐이므로, DFS로 한쪽 그룹의 인구 합을 구한 뒤 전체 인구에서 그 값을 뺀 결과와 비교하여 더 큰 쪽을 답으로 삼으면 됩니다.
다음 순서로 진행합니다.
- source와 dest를 이용해 그래프의 인접 리스트(adj)를 만듭니다.
- dfs(x, choose) 함수를 정의합니다.
- x를 이미 방문했다면 0을 반환합니다.
- x를 방문 처리합니다.
- ans := 0 으로 초기화합니다.
- choose가 참이면 ans에 population[x]를 더합니다.
- adj[x]의 모든 이웃 노드에 대해 ans에 dfs(neighbor, not choose)의 결과를 더합니다.
- ans를 반환합니다.
- 메인 로직에서 x := dfs(0, True)를 계산하고, max(x, sum(population) - x)를 반환합니다.
예제 코드
아래 구현을 통해 더 잘 이해해 보겠습니다.
from collections import defaultdict
class Solution:
def solve(self, source, dest, population):
adj = defaultdict(list)
for a, b in zip(source, dest):
adj[a].append(b)
adj[b].append(a)
seen = set()
def dfs(x, choose):
if x in seen:
return 0
seen.add(x)
ans = 0
if choose:
ans += population[x]
for neighbor in adj[x]:
ans += dfs(neighbor, not choose)
return ans
x = dfs(0, True)
return max(x, sum(population) - x)
ob = Solution()
source = [2, 2, 1, 1]
dest = [1, 3, 4, 0]
population = [6, 8, 4, 3, 5]
print(ob.solve(source, dest, population))
입력
[2, 2, 1, 1], [1, 3, 4, 0], [6, 8, 4, 3, 5]
출력
15
복잡도 분석
DFS가 모든 노드를 정확히 한 번씩 방문하므로 시간 복잡도는 O(N)입니다. 인접 리스트와 방문 집합을 저장해야 하므로 공간 복잡도 역시 O(N)입니다.