문제 정보
- 문제 번호 : 13325
- 문제 이름 : 이진 트리
- 문제 링크 : https://www.acmicpc.net/problem/13325
- 정답 코드 : https://github.com/dalmengs/algorithm-solutions/blob/main/13325/main.py
- 난이도 : Gold 3
- 체감 난이도 : Gold 2
- 알고리즘 : Tree, DP
문제 설명
양수 가중치가 부여돼 있는 높이가 k인 포화 이진트리가 주어진다.
이 트리 간선의 가중치를 자유롭게 증가시켜 루트에서 모든 리프까지의 거리를 같게 만들면서, 가중치의 총합이 최소가 되도록 하는 것이다.

왼쪽 트리에서 가중치를 증가시켜 오른쪽 트리 처럼 만들어야 하고, 이 경우 루트에서 모든 리프까지의 거리가 5 이고, 가중치들의 총합은 15 이다. (이 경우가 최소이다.)
문제 해결
현재 정점을 now라고 할 때, now의 자식 정점을 각각 c_i이라고 하자. c_i를 루트로 하는 서브 트리는 c_i부터 모든 리프 노드까지의 거리가 같아야 한다.
왜냐하면 c_i보다 상위에 있는 간선의 가중치가 증가하면 이 트리의 리프 노드까지의 거리가 모두 같이 증가하기 때문이다.
따라서 now의 자식들은 재귀적으로 이미 리프까지의 거리가 모두 같은 상태여야 하고, now의 직계 자식들을 비교하여 거리가 다른 경우 가장 큰 값에 맞춰주면 된다.
가중치 업데이트 전 / 후 간선의 값을 모두 저장하기 위해, 그래프를 만들 때 [src, before_dest, after_dest] 형태로 저장했다.
나는 DFS를 두 번 돌려 문제를 해결했다.
첫 번째 DFS에서는 리프까지의 거리의 거리를 받아 가중치를 직접 업데이트하고, 두 번째 DFS에서는 바뀐 가중치로 트리를 다시 한 번 순회하여 모든 간선 가중치 합을 구했다.
정답 코드
import sys
sys.setrecursionlimit(123456)
k = int(input())
n = 2 ** (k + 1) - 1
w = list(map(int, input().split()))
g = { i: [] for i in range(1, n + 1) }
for i in range(len(w)):
src = (i + 2) // 2
dest = i + 2
g[src].append([dest, w[i], w[i]])
vis = [0 for i in range(n + 1)]
def dfs(now):
vis[now] = 1
ret = 0
ws = []
for nxt in g[now]:
node_idx = nxt[0]
node_weight = nxt[1]
if vis[node_idx]: continue
res = dfs(node_idx)
ws.append(res[0] + node_weight)
ret += res[1]
if len(ws) == 0:
return [0, 0]
m = max(ws)
q = 0
idx = -1
for nxt in g[now]:
idx += 1
q += m - ws[idx]
g[now][idx][2] += m - ws[idx]
return [m, ret + q]
dfs(1)
vis = [0 for i in range(n + 1)]
def ans(now):
vis[now] = 1
ret = 0
for nxt in g[now]:
node_idx = nxt[0]
node_weight = nxt[2]
if vis[node_idx]: continue
ret += (ans(node_idx) + node_weight)
return ret
print(ans(1))
마무리
직계 자식을 루트로 하는 서브 트리를 Sub Problem으로 인식할 수 있다면 자연스럽게 그 최댓값에 맞추어 가중치를 업데이트해야 한다는 아이디어는 쉽게 떠올릴 수 있다.
문제에 딱히 함정도 없고 써야 하는 알고리즘도 확실히 보여서 연습하기 좋은 문제라고 생각한다.

댓글 남기기