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

파이썬으로 BST 중앙값 구하기 — O(n) 시간, O(1) 공간에 해결하는 모리스 순회 기법

이진 탐색 트리(Binary Search Tree, BST)가 주어졌을 때, 이 트리의 중앙값(median)을 구하는 문제를 생각해 봅시다.

중앙값의 정의는 노드 개수에 따라 다음과 같습니다.

  • 노드 수가 짝수일 때: 중앙값 = ((n/2번째 노드 + (n+1)/2번째 노드) / 2
  • 노드 수가 홀수일 때: 중앙값 = (n+1)/2번째 노드

예를 들어 아래와 같은 BST가 입력으로 주어지면, 출력은 7이 됩니다.

접근 방법: 모리스 순회(Morris Traversal)

일반적인 중위 순회(inorder traversal)는 재귀 호출이나 스택을 사용하기 때문에 O(n)의 추가 공간이 필요합니다. 하지만 모리스 스레드 이진 트리(Morris Threading) 기법을 활용하면 스택 없이도 중위 순회를 수행할 수 있어, 시간 복잡도 O(n), 공간 복잡도 O(1)로 중앙값을 구할 수 있습니다.

핵심 아이디어는 다음과 같습니다.

  1. 트리가 비어 있으면(root가 None이면) 0을 반환합니다.
  2. 먼저 전체 노드 개수(node_count)를 구합니다.
  3. 현재 노드(current)를 루트로 설정하고, 방문한 노드 수(count_curr)를 0으로 초기화한 뒤 트리를 순회합니다.
  4. 순회하면서 k번째 노드를 찾아, 노드 개수의 홀짝 여부에 따라 중앙값을 계산해 반환합니다.

알고리즘 단계

  • root가 None이면 → 0을 반환
  • node_count := 트리의 전체 노드 개수
  • count_curr := 0, current := root
  • current가 null이 아닌 동안 반복:
    • current.left가 null인 경우
      • count_curr를 1 증가
      • 노드 수가 홀수이고 count_curr == (node_count + 1)/2이면 → previous.data 반환
      • 그렇지 않고 노드 수가 짝수이며 count_curr == (node_count/2) + 1이면 → (previous.data + current.data) / 2 반환
      • previous := current, current := current.right로 이동
    • current.left가 존재하는 경우
      • previous := current.left
      • previous.right가 null이 아니고 current가 아닌 동안 previous := previous.right 이동
      • previous.right가 null이면 → 스레드 생성(previous.right := current, current := current.left)
      • 그렇지 않으면 → 스레드 제거(previous.right := None) 후 count_curr 증가, 홀짝 조건에 따라 중앙값 반환
      • previous := current, current := current.right로 이동

파이썬 구현 예제

아래 코드를 통해 더 자세히 이해해 보겠습니다.

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

def number_of_nodes(root):
    node_count = 0
    if (root == None):
        return node_count
    current = root
    while (current != None):
        if (current.left == None):
            node_count += 1
            current = current.right
        else:
            previous = current.left
            while (previous.right != None and previous.right != current):
                previous = previous.right
            if (previous.right == None):
                previous.right = current
                current = current.left
            else:
                previous.right = None
                node_count += 1
                current = current.right
    return node_count

def calculate_median(root):
    if (root == None):
        return 0
    node_count = number_of_nodes(root)
    count_curr = 0
    current = root
    while (current != None):
        if (current.left == None):
            count_curr += 1
            if (node_count % 2 != 0 and count_curr == (node_count + 1)//2):
                return previous.data
            elif (node_count % 2 == 0 and count_curr == (node_count//2)+1):
                return (previous.data + current.data)//2
            previous = current
            current = current.right
        else:
            previous = current.left
            while (previous.right != None and previous.right != current):
                previous = previous.right
            if (previous.right == None):
                previous.right = current
                current = current.left
            else:
                previous.right = None
                previous = previous
                count_curr += 1
                if (node_count % 2 != 0 and count_curr == (node_count + 1) // 2):
                    return current.data
                elif (node_count % 2 == 0 and count_curr == (node_count // 2) + 1):
                    return (previous.data + current.data)//2
                previous = current
                current = current.right

root = TreeNode(7)
root.left = TreeNode(4)
root.right = TreeNode(9)
root.left.left = TreeNode(2)
root.left.right = TreeNode(5)
root.right.left = TreeNode(8)
root.right.right = TreeNode(10)
print(calculate_median(root))

입력

root = TreeNode(7)
root.left = TreeNode(4)
root.right = TreeNode(9)
root.left.left = TreeNode(2)
root.left.right = TreeNode(5)
root.right.left = TreeNode(8)
root.right.right = TreeNode(10)

출력

7

정리

이 알고리즘은 모리스 순회를 두 번 수행합니다. 첫 번째 순회에서 전체 노드 개수를 세고, 두 번째 순회에서 중앙 위치의 노드 값을 찾습니다. 각 간선은 최대 상수 번만 지나가므로 전체 시간 복잡도는 O(n)이며, 재귀나 스택 없이 트리 내부의 임시 링크(thread)만 사용하므로 추가 공간은 O(1)입니다. 단, 순회 중 트리 구조가 일시적으로 변경되었다가 원복되므로, 멀티스레드 환경에서는 주의해서 사용해야 합니다.