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

Python으로 풀어보는 최대 평균 서브트리(Subtree) 문제

문제 개요

이진 트리(binary tree)의 루트 노드가 주어졌을 때, 해당 트리에 속한 모든 서브트리 중에서 평균값이 가장 큰 서브트리의 평균을 구하는 문제입니다.

예를 들어 아래와 같은 트리가 있다고 가정해 보겠습니다.

Python으로 풀어보는 최대 평균 서브트리(Subtree) 문제

이 경우 출력 결과는 6입니다. 그 이유는 다음과 같습니다.

  • 노드 5를 루트로 하는 서브트리: (5 + 6 + 1) / 3 = 4
  • 노드 6만 있는 서브트리: 6 / 1 = 6
  • 노드 1만 있는 서브트리: 1 / 1 = 1

세 값 중 가장 큰 값인 6이 정답이 됩니다.

풀이 접근 방식

이 문제는 후위 순회(postorder traversal)를 활용하면 효율적으로 해결할 수 있습니다. 핵심 아이디어는 각 노드마다 자신을 포함한 서브트리의 노드 개수합계를 함께 계산하면서 올라오는 것입니다.

알고리즘 단계

  1. 결괏값을 저장할 변수 ans를 0으로 초기화합니다.
  2. solve(root)라는 재귀 함수를 정의합니다.
  3. 루트가 None이면 (노드 개수, 합계) 쌍인 (0, 0)을 반환합니다.
  4. 왼쪽 자식과 오른쪽 자식에 대해 재귀적으로 solve()를 호출합니다.
  5. 현재 서브트리의 노드 개수 c = 왼쪽 개수 + 오른쪽 개수 + 1로 계산합니다.
  6. 현재 서브트리의 합계 s = 왼쪽 합 + 오른쪽 합 + 현재 노드의 값으로 계산합니다.
  7. ansmax(ans, s / c)로 갱신합니다.
  8. (c, s) 쌍을 상위 호출로 반환합니다.

모든 탐색이 끝나면 ans에 최대 평균값이 저장되며, 이를 반환하면 됩니다.

Python 구현 예제

아래는 위 알고리즘을 실제로 구현한 전체 코드입니다.

class TreeNode:
    def __init__(self, data, left=None, right=None):
        self.data = data
        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

class Solution(object):
    def helper(self, node):
        if not node:
            return 0, 0
        left_sum, left_count = self.helper(node.left)
        right_sum, right_count = self.helper(node.right)
        self.ans = max(self.ans,
                       (left_sum + right_sum + node.data) /
                       (left_count + right_count + 1))
        return left_sum + right_sum + node.data, left_count + right_count + 1

    def maximumAverageSubtree(self, root):
        self.ans = 0
        self.helper(root)
        return self.ans

ob = Solution()
root = make_tree([5, 6, 1])
print(ob.maximumAverageSubtree(root))

실행 결과 확인

입력

[5, 6, 1]

출력

6.0

정리 및 복잡도 분석

이 풀이는 트리의 모든 노드를 정확히 한 번씩 방문하므로 시간 복잡도는 O(N), 재귀 호출 스택 깊이 때문에 공간 복잡도는 최악의 경우(편향된 트리) O(N), 균형 잡힌 트리의 경우 O(log N)입니다.

핵심 포인트는 각 노드에서 서브트리의 합계와 노드 개수를 동시에 반환함으로써, 매번 서브트리를 다시 순회하는 비효율적인 O(N²) 접근을 피했다는 점입니다. 이러한 '부분 결과를 반환하며 올려보내는' 패턴은 트리 관련 문제에서 매우 자주 활용되므로 꼭 익혀두시기 바랍니다.