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

이때 기대하는 출력값은 12입니다. 그 이유는 주어진 이진 트리 안에서 아래와 같은 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을 반환
- node가 null이 아니라면:
- calculate_sum() 함수를 정의합니다. 이 함수는 root를 인자로 받습니다.
- root가 null이 아니라면:
- calculate_sum(root의 왼쪽 자식) 호출
- value := value + root의 값
- calculate_sum(root의 오른쪽 자식) 호출
- root가 null이 아니라면:
- 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)입니다.