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

Python으로 두 개의 BST에서 주어진 합을 만족하는 쌍 찾기


문제 이해

두 개의 이진 탐색 트리(Binary Search Tree, BST)와 하나의 목표 합(sum)이 주어졌을 때, 두 요소의 합이 주어진 값과 같으면서 각 쌍의 요소들이 반드시 서로 다른 BST에 속하도록 하는 모든 쌍을 찾아야 합니다.

예를 들어 합이 12로 주어진 경우를 살펴보겠습니다.

Python으로 두 개의 BST에서 주어진 합을 만족하는 쌍 찾기

이때 출력은 [(6, 6), (7, 5), (9, 3)]이 됩니다.

접근 방법

이 문제의 핵심 아이디어는 간단합니다. BST를 중위 순회(inorder traversal)하면 항상 오름차순으로 정렬된 리스트를 얻을 수 있습니다. 따라서 두 트리를 각각 중위 순회해 정렬된 두 배열을 만든 뒤 투 포인터(two pointer) 기법을 적용하면, 선형 시간 안에 조건을 만족하는 모든 쌍을 효율적으로 찾을 수 있습니다.

해결 절차는 다음과 같습니다.

  1. solve() 함수 정의 — 인자로 trav1, trav2, Sum을 받습니다.
  2. left := 0으로 초기화합니다.
  3. right := trav2의 크기 − 1로 초기화합니다.
  4. res := 결과를 담을 새로운 리스트를 생성합니다.
  5. left < trav1의 크기이고 right ≥ 0인 동안 다음을 반복합니다.
    • trav1[left] + trav2[right] == Sum인 경우: (trav1[left], trav2[right])를 res 끝에 추가하고, left는 1 증가, right는 1 감소시킵니다.
    • trav1[left] + trav2[right] < Sum인 경우: 합을 키우기 위해 left를 1 증가시킵니다.
    • 그 외(합이 Sum보다 큰 경우): 합을 줄이기 위해 right를 1 감소시킵니다.
  6. 반복이 종료되면 res를 반환합니다.

메인 루틴에서는 다음 작업을 수행합니다.

  • trav1, trav2라는 두 개의 새로운 리스트를 생성합니다.
  • trav1 := 첫 번째 트리(tree1)의 중위 순회 결과
  • trav2 := 두 번째 트리(tree2)의 중위 순회 결과
  • solve(trav1, trav2, Sum)을 반환합니다.

구현 예제 (Python)

다음 구현을 통해 동작 과정을 더 자세히 이해할 수 있습니다.

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

def insert(root, key):
    if root == None:
        return TreeNode(key)
    if root.data > key:
        root.left = insert(root.left, key)
    else:
        root.right = insert(root.right, key)
    return root

def storeInorder(ptr, traversal):
    if ptr == None:
        return
    storeInorder(ptr.left, traversal)
    traversal.append(ptr.data)
    storeInorder(ptr.right, traversal)

def solve(trav1, trav2, Sum):
    left = 0
    right = len(trav2) - 1
    res = []
    while left < len(trav1) and right >= 0:
        if trav1[left] + trav2[right] == Sum:
            res.append((trav1[left], trav2[right]))
            left += 1
            right -= 1
        elif trav1[left] + trav2[right] < Sum:
            left += 1
        else:
            right -= 1
    return res

def get_pair_sum(root1, root2, Sum):
    trav1 = []
    trav2 = []
    storeInorder(root1, trav1)
    storeInorder(root2, trav2)
    return solve(trav1, trav2, Sum)

root1 = None
for element in [9, 11, 4, 7, 2, 6, 15, 14]:
    root1 = insert(root1, element)

root2 = None
for element in [6, 19, 3, 2, 4, 5]:
    root2 = insert(root2, element)

Sum = 12
print(get_pair_sum(root1, root2, Sum))

입력

[9,11,4,7,2,6,15,14], [6,19,3,2,4,5], 12

출력

[(6, 6), (7, 5), (9, 3)]

복잡도 분석

시간 복잡도: O(n + m) — n과 m은 각각 첫 번째, 두 번째 트리의 노드 수입니다. 중위 순회에 O(n + m), 투 포인터 탐색에 최대 O(n + m)이 소요됩니다.

공간 복잡도: O(n + m) — 두 트리의 중위 순회 결과를 저장하기 위한 추가 공간이 필요합니다.