트리의 모든 노드가 1부터 n까지 번호가 매겨져 있고, 각 노드에는 정수 값이 저장되어 있다고 가정해 봅시다. 이때 트리에서 하나의 간선을 제거하면 트리는 두 개의 부분 트리로 나뉘는데, 우리의 목표는 이 두 부분 트리의 노드 값 합의 차이가 최소가 되도록 하는 것입니다. 즉, 간선을 하나씩 제거해 가며 두 부분 트리 값 합의 차이를 계산하고, 그중 최솟값을 찾아 반환해야 합니다.
트리는 간선(edge) 목록의 형태로 주어지며, 각 노드의 값들도 함께 제공됩니다.
문제 예시
예를 들어 다음과 같은 입력이 주어졌다고 해보겠습니다.
- n = 6
- edge_list = [[1, 2], [1, 3], [2, 4], [3, 5], [3, 6]]
- values = [15, 25, 15, 55, 15, 65]
이 경우 출력은 0이 됩니다.
각 간선을 제거했을 때의 결과를 살펴보면 다음과 같습니다.
- 간선 (1, 2) 제거 → 두 부분 트리의 합은 80, 110 → 차이는 30
- 간선 (1, 3) 제거 → 두 부분 트리의 합은 95, 95 → 차이는 0
- 간선 (2, 4) 제거 → 두 부분 트리의 합은 55, 135 → 차이는 80
- 간선 (3, 5) 제거 → 두 부분 트리의 합은 15, 175 → 차이는 160
- 간선 (3, 6) 제거 → 두 부분 트리의 합은 65, 125 → 차이는 60
따라서 최소 차이는 0입니다.
풀이 접근 방법
이 문제는 리프 노드부터 시작해 위쪽으로 올라가면서 각 서브트리의 누적 합을 계산하는 방식으로 해결할 수 있습니다. 전체 단계는 다음과 같습니다.
- 크기가 n인 인접 리스트(adj_list)를 생성하고, 모든 간선에 대해 양방향 연결 정보를 저장합니다.
- 크기가 n인 value_list를 0으로 초기화합니다.
- 차수(연결된 간선 수)가 1인 노드, 즉 리프 노드들을 not_visited 집합에 넣습니다.
- not_visited가 빌 때까지 다음 과정을 반복합니다.
- 각 노드 i에 대해 value_list[i]에 자신의 값을 더합니다.
- adj_list[i]가 비어 있지 않다면, 인접한 부모 노드에서 자신을 제거하고 부모의 value_list에 자신의 누적 합을 더합니다.
- 그다음 후보 노드들의 부모 중 차수가 1이 된 노드들을 새로운 not_visited로 설정합니다.
- 전체 값의 합에서 각 노드의 서브트리 합의 2배를 뺀 절댓값을 계산합니다. 즉, |sum(values) - 2 * value_list[i]|가 곧 두 부분 트리의 차이입니다.
- 모든 노드에 대해 이 값을 비교하여 최솟값을 반환합니다.
여기서 |sum(values) - 2 * value_list[i]| 공식이 성립하는 이유는, 전체 합을 S, 한쪽 부분 트리의 합을 A라고 할 때 다른 쪽의 합은 S - A이므로 두 값의 차이는 |S - 2A|와 같기 때문입니다.
구현 예제
다음은 위 알고리즘을 Python으로 구현한 코드입니다.
def solve(n, edge_list, values):
adj_list = [[] for i in range(n)]
for edge in edge_list:
u = edge[0]
v = edge[1]
adj_list[u-1].append(v-1)
adj_list[v-1].append(u-1)
value_list = [0] * n
not_visited = {i for i in range(n) if len(adj_list[i]) == 1}
while(len(not_visited)):
for i in not_visited:
value_list[i] += values[i]
if(len(adj_list[i])):
adj_list[adj_list[i][0]].remove(i)
value_list[adj_list[i][0]] += value_list[i]
not_visited = {adj_list[i][0] for i in not_visited if
len(adj_list[i]) and len(adj_list[adj_list[i][0]]) == 1}
return_val = abs(sum(values) - 2 * value_list[0])
for i in range(1, n):
decision_val = abs(sum(values) - 2 * value_list[i])
if decision_val < return_val:
return_val = decision_val
return return_val
print(solve(6, [[1, 2], [1, 3], [2, 4], [3, 5], [3, 6]], [10, 20, 10, 50, 10, 60]))입력
6, [[1, 2], [1, 3], [2, 4], [3, 5], [3, 6]], [10, 20, 10, 50, 10, 60]
출력
0
이처럼 리프 노드부터 bottom-up 방식으로 서브트리의 합을 누적하면, 모든 간선을 실제로 제거해 보지 않고도 각 간선 제거 시의 두 부분 트리 값 차이를 효율적으로 계산할 수 있습니다.