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

파이썬으로 이진 트리에서 가장 큰 BST의 노드 합 구하기


문제 개요

이진 트리(binary tree)가 하나 주어진다고 가정해 봅시다. 우리가 해야 할 일은 이 트리의 서브트리(subtree) 가운데 이진 탐색 트리(BST)의 조건을 만족하는 부분 트리를 찾아내고, 그중에서 가장 큰 BST(노드 개수가 가장 많은 BST)에 포함된 모든 노드 값의 합을 계산하는 것입니다. 최종 결과로 그 합계를 반환합니다.

입력 예시와 기대 출력

예를 들어 다음과 같은 이진 트리가 입력으로 주어진 경우를 생각해 보겠습니다.

파이썬으로 이진 트리에서 가장 큰 BST의 노드 합 구하기

이때 기대하는 출력값은 12입니다. 그 이유는 주어진 이진 트리 안에서 아래와 같은 BST를 찾을 수 있기 때문입니다.

파이썬으로 이진 트리에서 가장 큰 BST의 노드 합 구하기

이 BST를 구성하는 노드들은 4, 3, 5이며, 그 합은 4 + 3 + 5 = 12가 됩니다.

풀이 접근 방법

이 문제는 후위 순회(postorder traversal) 기반의 재귀적 접근으로 해결할 수 있습니다. 각 노드마다 왼쪽과 오른쪽 서브트리가 BST 조건을 만족하는지 확인하고, 만족한다면 노드 수를 누적하여 가장 큰 BST의 루트를 추적합니다. 전체 과정은 다음과 같습니다.

  • 변수를 초기화합니다: c := 0, m := null, value := 0
  • recurse() 함수를 정의합니다. 이 함수는 node를 인자로 받습니다.
    • node가 null이 아니라면:
      • left_val := recurse(node의 왼쪽 자식)
      • right_val := recurse(node의 오른쪽 자식)
      • count := 음의 무한대(-∞)로 초기화
      • 만약 (node.left가 null이거나 node.left.val <= node.val)이고, 동시에 (node.right가 null이거나 node.val <= node.right.val)이라면 현재 노드도 BST 조건을 만족하므로 count := left_val + right_val + 1
      • 만약 count > c라면 지금까지 찾은 최대 BST가 갱신되므로 c := count, m := node
      • count를 반환
    • node가 null이면 0을 반환
  • calculate_sum() 함수를 정의합니다. 이 함수는 root를 인자로 받습니다.
    • root가 null이 아니라면:
      • calculate_sum(root의 왼쪽 자식) 호출
      • value := value + root의 값
      • calculate_sum(root의 오른쪽 자식) 호출
  • recurse(root)를 호출하여 가장 큰 BST의 루트 m을 찾습니다.
  • calculate_sum(m)을 호출하여 해당 BST의 노드 합을 구합니다.
  • value를 반환합니다.

파이썬 구현 예시

더 나은 이해를 돕기 위해 위 알고리즘을 파이썬으로 구현한 코드를 살펴보겠습니다.

class TreeNode:
    def __init__(self, val, left=None, right=None):
        self.val = val
        self.left = left
        self.right = right

def insert(temp, data):
    que = []
    que.append(temp)
    while len(que):
        temp = que[0]
        que.pop(0)
        if not temp.left:
            if data is not None:
                temp.left = TreeNode(data)
            else:
                temp.left = TreeNode(0)
            break
        else:
            que.append(temp.left)
        if not temp.right:
            if data is not None:
                temp.right = TreeNode(data)
            else:
                temp.right = TreeNode(0)
            break
        else:
            que.append(temp.right)

def make_tree(elements):
    Tree = TreeNode(elements[0])
    for element in elements[1:]:
        insert(Tree, element)
    return Tree

def solve(root):
    c, m, value = 0, None, 0
    def recurse(node):
        if node:
            nonlocal c, m
            left_val = recurse(node.left)
            right_val = recurse(node.right)
            count = -float("inf")
            if (node.left == None or node.left.val <= node.val) and (node.right == None or node.val <= node.right.val):
                count = left_val + right_val + 1
            if count > c:
                c = count
                m = node
            return count
        return 0
    def calculate_sum(root):
        nonlocal value
        if root is not None:
            calculate_sum(root.left)
            value += root.val
            calculate_sum(root.right)
    recurse(root)
    calculate_sum(m)
    return value

tree = make_tree([1, 4, 6, 3, 5])
print(solve(tree))

실행 결과 확인

입력

tree = make_tree([1, 4, 6, 3, 5])
print(solve(tree))

출력

12

복잡도 분석

이 알고리즘은 트리의 모든 노드를 정확히 한 번씩 방문하므로 시간 복잡도는 O(n)입니다. 재귀 호출에 사용되는 스택 공간이 트리의 높이에 비례하므로, 공간 복잡도는 균형 잡힌 트리의 경우 O(log n), 편향된 트리의 경우 최악 O(n)입니다.