문제 소개
n개의 도시가 n−1개의 도로로 연결되어 있어, 어떤 도시에서든 다른 모든 도시로 이동할 수 있다고 가정해 봅시다. 즉, 도시들은 트리(tree) 구조를 이룹니다. 매일 우편 시스템은 k통의 편지를 처리하며, 각 편지의 목적지는 서로 다른 k개의 도시 중 하나입니다. 우체부는 매일 모든 편지를 수신지에 배달해야 하며, 우리는 이때 이동해야 하는 최소 총 거리를 구해야 합니다. 우체부는 편의상 어떤 도시에서든 출발할 수 있습니다.
예를 들어 입력이 아래 그림과 같고, 배달해야 할 도시(delv)가 1, 2, 4라고 가정하면 출력은 4가 됩니다.

우체부는 1, 2, 4번 도시 중 어디서든 배달을 시작할 수 있습니다. 1번 도시에서 출발하면 경로는 1 → 2 → 4가 되고, 4번 도시에서 출발하면 그 반대인 4 → 2 → 1이 됩니다. 두 경우 모두 총 비용은 1 + 3 = 4입니다. 반면 2번 도시에서 출발하면 왕복해야 하는 구간이 늘어나 비용이 더 커지게 됩니다.
풀이 접근 방식
이 문제의 핵심 아이디어는 다음과 같습니다.
- 배달 대상 도시들을 연결하는 데 필요한 간선만 실제로 사용되며, 이 간선들의 가중치 합을 SUM이라고 합시다. 우체부는 대부분의 간선을 지나쳤다가 되돌아와야 하므로 기본 비용은 2 × SUM입니다.
- 다만 출발 도시부터 마지막 배달 도시까지 이어지는 가장 긴 경로 하나는 한 번만 지나도 됩니다. 따라서 배달 도시들이 속한 트리의 최장 경로 길이를 MAX라고 하면 정답은
2 × SUM − MAX가 됩니다.
알고리즘 단계
이 문제를 해결하기 위해 다음 단계를 따릅니다.
depth_search()함수를 정의합니다. 이 함수는 현재 노드(node)와 부모 노드(p)를 인자로 받습니다.- d1 := −∞ (음의 무한대)
- d2 := −∞
- adj_list[node]의 각 쌍 (x, y)에 대해 다음을 수행합니다.
- x가 p와 같지 않다면:
- d1 := max(d1, depth_search(x, node) + y)
- d1 > d2이면 d1과 d2의 값을 서로 교환
- ti[node] := ti[node] + ti[x]
- 0 < ti[x] < k이면 SUM := SUM + y
- x가 p와 같지 않다면:
- d1 > 0이면 MAX := max(MAX, d1 + d2)
- d2 > 0이고 tj[node]가 0이 아니면 MAX := max(MAX, d2)
- tj[node]가 0이 아니면 d2 := max(0, d2)
- d2를 반환합니다.
- k := delv의 크기
- adj_list := 새로운 맵(딕셔너리)
- ti := 크기가 (nodes + 5)이고 0으로 초기화된 새 리스트
- tj := 크기가 (nodes + 5)이고 0으로 초기화된 새 리스트
- delv의 각 i에 대해 ti[i] := 1, tj[i] := 1로 설정
- roads의 각 항목(item)에 대해 다음을 수행합니다.
- x := item[0], y := item[1], c := item[2]
- x가 adj_list에 없으면 adj_list[x] := [] 생성
- y가 adj_list에 없으면 adj_list[y] := [] 생성
- adj_list[x]의 끝에 (y, c) 추가
- adj_list[y]의 끝에 (x, c) 추가
- SUM := 0, MAX := 0으로 초기화
- depth_search(1, 1) 호출
- SUM * 2 − MAX를 반환합니다.
예제 코드
아래 파이썬 구현을 통해 더 자세히 이해해 보겠습니다.
import sys
from math import inf as INF
sys.setrecursionlimit(10**5 + 5)
def depth_search(node, p):
global SUM, MAX
d1 = -INF
d2 = -INF
for x, y in adj_list[node]:
if x != p:
d1 = max(d1, depth_search(x, node) + y)
if d1 > d2:
d1, d2 = d2, d1
ti[node] += ti[x]
if 0 < ti[x] < k:
SUM += y
if d1 > 0: MAX = max(MAX, d1 + d2)
if d2 > 0 and tj[node]: MAX = max(MAX, d2)
if tj[node]: d2 = max(0, d2)
return d2
def solve(nodes, delv, roads):
global k, ti, tj, adj_list, SUM, MAX
k = len(delv)
adj_list = {}
ti = [0] * (nodes + 5)
tj = [0] * (nodes + 5)
for i in delv:
ti[i] = tj[i] = 1
for item in roads:
x, y, c = map(int, item)
if x not in adj_list: adj_list[x] = []
if y not in adj_list: adj_list[y] = []
adj_list[x].append([y, c])
adj_list[y].append([x, c])
SUM = 0
MAX = 0
depth_search(1,1)
return SUM * 2 - MAX
print(solve(5, [1, 2, 4], [(1,2,1),(2,3,2),(2,4,3),(1,5,1)]))입력
5, [1, 2, 4], [(1,2,1),(2,3,2),(2,4,3),(1,5,1)]
출력
4