Computer >> 컴퓨터 >  >> 프로그래밍 >> Python

Python으로 분리된 트리(숲)를 하나의 트리로 연결하는 프로그램

그래프가 인접 리스트(adjacency list) 형태로 주어져 있다고 가정해 보겠습니다. 이 그래프는 실제로 서로 연결되지 않은 여러 개의 트리, 즉 숲(forest)으로 구성되어 있습니다. 우리는 여기에 적절한 수의 간선을 추가해 전체를 하나의 트리로 만들어야 하며, 그때 임의의 두 노드 사이의 최장 경로 길이가 가능한 한 최소가 되도록 해야 합니다.

Python으로 분리된 트리(숲)를 하나의 트리로 연결하는 프로그램

예를 들어 위와 같은 입력이 주어지면 출력은 4가 됩니다.

간선 0 → 5를 추가하면 최장 경로는 3 → 1 → 0 → 5 → 7 또는 4 → 1 → 0 → 5 → 7이 될 수 있으며, 방향을 반대로 한 경로 역시 마찬가지입니다. 따라서 답은 거리 4입니다.

핵심 아이디어

각 트리의 지름(diameter), 즉 트리 내에서 가장 긴 경로의 길이를 먼저 구합니다. 트리의 '중심'에서 가장 먼 노드까지의 거리는 지름의 절반을 올림한 값, 즉 ceil(지름 / 2)이 됩니다. 여러 트리를 하나로 연결할 때는 반지름이 가장 큰 두 트리를 서로 연결하는 것이 유리하며, 이때 두 트리를 잇는 최장 경로는 r1 + r2 + 1(r1, r2는 각 트리의 반지름)이 됩니다. 최종 정답은 개별 트리 지름의 최댓값과 트리들을 병합했을 때의 최장 경로 중 더 큰 값입니다.

풀이 절차

이 문제는 다음 단계를 따라 해결할 수 있습니다.

  • seen := 새로운 집합(set)

  • dic := graph

  • treeDepth() 함수를 정의합니다. 이 함수는 노드를 인자로 받습니다.

    • ret := 0

    • dfs1() 함수를 정의합니다. 이 함수는 노드와 부모 노드를 인자로 받습니다.

      • 현재 노드를 seen 집합에 추가

      • best2 := 빈 최소 힙(min heap) 생성

      • dic[node]의 각 인접 노드 nxt에 대해 다음을 반복

        • nxt가 부모 노드와 같지 않다면 dfs1(nxt, node) + 1 값을 best2에 push

        • best2의 크기가 2보다 커지면 힙에서 pop하여 항상 가장 큰 두 값만 유지

      • best2가 비어 있으면 0을 반환

      • ret := ret과 best2 요소들의 합 중 최댓값 (현재 노드를 지나는 가장 긴 경로)

      • best2의 최댓값 반환

    • dfs1(node, None) 호출

    • ret 반환

  • 메인 메서드에서는 다음을 수행합니다.

    • ret := 0, opt := 새로운 리스트, sing := 0으로 초기화

    • 0부터 그래프 크기까지의 각 노드에 대해 다음을 반복

      • 노드가 이미 seen에 있다면 다음 반복으로 건너뜀

      • res := treeDepth(node) (해당 트리의 지름)

      • sing := sing과 res 중 최댓값

      • res / 2의 올림값을 opt 리스트 끝에 추가

    • opt의 크기가 1 이하이면 sing을 반환

    • mx := opt의 최댓값

    • opt에서 mx와 같은 첫 번째 원소를 찾아 1을 감소시킨 뒤 반복 종료

    • opt의 모든 원소에 1씩 더함 (트리를 연결할 때 필요한 간선 수만큼 경로가 늘어나는 효과)

    • high2 := opt에서 가장 큰 두 원소

    • sum(high2)sing 중 최댓값을 반환

아래 구현을 통해 더 자세히 이해해 보겠습니다.

예제

import heapq, math
class Solution:
    def solve(self, graph):
        seen = set()
        dic = graph
        def treeDepth(node):
            self.ret = 0
            def dfs1(node, parent):
                seen.add(node)
                best2 = []
                for nxt in dic[node]:
                    if nxt != parent:
                        heapq.heappush(best2, dfs1(nxt, node) + 1)
                        if len(best2) > 2:
                            heapq.heappop(best2)
                if not best2:
                    return 0
                self.ret = max(self.ret, sum(best2))
                return max(best2)
            dfs1(node, None)
            return self.ret
        ret = 0
        opt = []
        sing = 0
        for node in range(len(graph)):
            if node in seen:
                continue
            res = treeDepth(node)
            sing = max(sing, res)
            opt.append(int(math.ceil(res / 2)))
        if len(opt) <= 1:
            return sing
        mx = max(opt)
        for i in range(len(opt)):
            if opt[i] == mx:
                opt[i] -= 1
                break
        for i in range(len(opt)):
            opt[i] += 1
        high2 = heapq.nlargest(2, opt)
        return max(sum(high2), sing)
ob = Solution()
graph = [
    [1, 2],
    [0,3,4],
    [0],
    [1],
    [1],
    [6,7],
    [5],
    [5]
]
print(ob.solve(graph))

입력

graph = [
    [1, 2],
    [0,3,4],
    [0],
    [1],
    [1],
    [6,7],
    [5],
    [5]
]

출력

4