[백준 / BOJ][Python] 5257 - timeismoney
https://www.acmicpc.net/problem/5257
5257번: timeismoney
In the first line of output print two numbers: the total time (SumTime) and total money (SumMoney) used in the optimal solution (the one with minimal value V), separated by one space. The next N-1 lines describe the links to be constructed. Each line conta
www.acmicpc.net
문제 풀이
거의 2년? 2년 반? 정도 전부터 언젠간 풀고 말겠다는 생각을 갖고 담아두던 문제였다. 백준의 서비스 종료가 코앞으로 다가온 지금, 하루라도 더 미룬다면 평생 못 풀수도 있겠다는 생각이 들어 풀이를 시작했다. 우선, time의 합과 cost의 합이 가장 작은 두 스패닝 트리를 구하는 것에서 시작한다. 두 스패닝 트리의 time과 cost의 합을 각각 t1, c1, t2, c2라 하면, 좌표 평면 위에 (t1, c1), (t2, c2)의 두 점을 찍어보자. 이제 잠시 위 과정은 뒤로 하고, 어떤 경우의 스패닝 트리가 합의 곱이 최소화되는지 생각해보자. time의 합을 X, cost의 합을 Y라 하면 XY = k라는 식을 세울 수 있는데, 우리의 목표는 최소가 되는 k를 찾는 것이다. 그렇다면 이 k값은 X와 Y가 어떤 경우일 때 최소가 될까? xy = k 형태의 함수를 좌표평면에 그리면, 이는 학창 시절 배운 유리함수 형태의 쌍곡선 그래프를 그린다. 즉, k를 최소로 하는 (x, y)의 후보는 좌표 평면에서 쌍곡선 그래프에 가장 먼저 닿는, 원점에 가장 가까운 점이 됨을 알 수 있다. 그렇다면 이 점들을 어떻게 구할 수 있겠는가? 쌍곡선의 형태는 바로 볼록껍질을 순회하며 구할 수 있다. 앞서 구한 두 점이 각각 좌하단 볼록껍질의 가장 왼쪽 점과 가장 아래쪽 점이 되기 때문에, 좌하단 볼록껍질을 구성하는 점들 중 이 두 점이 가장 끝 점이 되고, 이 두 점 사이의 볼록껍질을 구성하는 점들이 바로 k를 최소로 하는 점의 후보임을 알 수 있다. 따라서 이후에는 분할 정복을 통해 두 점 사이의 새 점을 구하고, 이 점의 구성 간선과 time_sum, cost_sum을 저장한 뒤 해당 점을 기준으로 탐색을 이어간다. 분할 정복 과정에서 간선을 정렬할 때의 기준이 되는 식은 외적을 통한 삼각형의 넓이 공식(aka 신발끈 공식)을 통해 유도할 수 있다. 양 끝점으로 부터 원점 방향으로 가장 먼 점을 골라야하기 때문.
코드
import sys
input = sys.stdin.readline
def find(x, graph):
if graph[x] != x:
graph[x] = find(graph[x], graph)
return graph[x]
def union(x, y, graph):
x = find(x, graph)
y = find(y, graph)
if x < y:
graph[y] = x
else:
graph[x] = y
def solve(ap, bp):
t1, c1 = ap
t2, c2 = bp
edge.sort(key = lambda x : (c1 - c2) * x[2] + (t2 - t1) * x[3])
graph = [i for i in range(n)]
cnt = 1
t_sum = c_sum = 0
tmp = []
for a, b, t, c in edge:
if find(a, graph) != find(b, graph):
union(a, b, graph)
tmp.append((a, b))
t_sum += t
c_sum += c
cnt += 1
if cnt == n:
break
cp = (t_sum, c_sum)
if cp in {ap, bp}:
return
point.append(cp)
edges.append(tmp)
solve(ap, cp)
solve(cp, bp)
n, m = map(int, input().split())
edge = [tuple(map(int, input().split())) for _ in range(m)]
point = []
edges = []
for k in [2, 3]:
edge.sort(key = lambda x : x[k])
graph = [i for i in range(n)]
cnt = 1
t_sum = c_sum = 0
tmp = []
for a, b, t, c in edge:
if find(a, graph) != find(b, graph):
union(a, b, graph)
tmp.append((a, b))
t_sum += t
c_sum += c
cnt += 1
if cnt == n:
break
point.append((t_sum, c_sum))
edges.append(tmp)
solve(point[0], point[1])
mn = 10 ** 18
mn_xy = []
for i in range(len(point)):
a, b = point[i]
if a * b < mn:
mn = a * b
mn_xy = [a, b]
ans = edges[i]
print(*mn_xy)
print()
for a, b in ans:
print(a, b)