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

파이썬으로 k개의 정렬된 리스트 병합하기 – 힙(Heap) 자료구조 활용법

여러 개의 정렬된 리스트가 주어졌을 때, 이것들을 하나의 정렬된 리스트로 병합해야 하는 상황을 생각해 봅시다. 이 문제는 힙(Heap) 자료구조를 활용하면 효율적으로 해결할 수 있습니다.
예를 들어 정렬된 리스트가 [1,4,5], [1,3,4], [2,6] 세 개 있다면, 병합한 최종 결과는 [1,1,2,3,4,4,5,6]이 됩니다.

알고리즘 접근 방식

핵심 아이디어는 각 리스트의 현재 헤드 노드만 최소 힙(min-heap)에 유지하고, 가장 작은 값을 가진 노드를 꺼내 결과 리스트에 차례로 연결하는 것입니다. 노드를 꺼낼 때마다 그 노드의 다음 노드를 힙에 다시 삽입하면 됩니다. 구체적인 단계는 다음과 같습니다.

  • 최소 힙을 하나 생성합니다.

  • lists의 각 연결 리스트 l에 대해 다음을 수행합니다.

    • l이 null이 아니라면 l을 힙에 삽입합니다.

  • res := null, res_next := null로 초기화합니다.

  • 무한 루프를 돌면서 다음을 반복합니다.

    • temp := 힙에서 최솟값을 추출합니다.

    • 힙이 비어 있으면 res를 반환하고 종료합니다.

    • res가 null인 경우:

      • res := temp, res_next := temp로 설정합니다.

      • temp := temp의 다음 노드로 이동합니다.

      • temp가 null이 아니면 temp를 힙에 삽입합니다.

      • res.next := null로 설정합니다.

    • 그렇지 않은 경우:

      • res_next.next := temp로 연결한 뒤, temp := temp의 다음 노드, res_next := res_next의 다음 노드로 이동합니다.

      • temp가 null이 아니면 temp를 힙에 삽입합니다.

      • res_next.next := null로 설정합니다.

구현 예제

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

class ListNode:
   def __init__(self, data, next = None):
      self.val = data
      self.next = next

def make_list(elements):
   head = ListNode(elements[0])
   for element in elements[1:]:
      ptr = head
      while ptr.next:
         ptr = ptr.next
      ptr.next = ListNode(element)
   return head

def print_list(head):
   ptr = head
   print('[', end = "")
   while ptr:
      print(ptr.val, end = ", ")
      ptr = ptr.next
   print(']')

class Heap:
   def __init__(self):
      self.arr = []

   def getVal(self, i):
      return self.arr[i].val

   def parent(self, i):
      return (i-1)//2

   def left(self, i):
      return (2*i + 1)

   def right(self, i):
      return (2*i + 2)

   def insert(self, value):
      self.arr.append(value)
      n = len(self.arr)-1
      i = n
      while i != 0 and self.arr[i].val < self.arr[self.parent(i)].val:
         self.arr[i], self.arr[self.parent(i)] = self.arr[self.parent(i)], self.arr[i]
         i = self.parent(i)

   def heapify(self, i):
      left = self.left(i)
      right = self.right(i)
      smallest = i
      n = len(self.arr)
      if left < n and self.getVal(left) < self.getVal(smallest): smallest = left
      if right < n and self.getVal(right) < self.getVal(smallest): smallest = right
      if smallest != i:
         self.arr[i], self.arr[smallest] = self.arr[smallest], self.arr[i]
         self.heapify(smallest)

   def extractMin(self):
      n = len(self.arr)
      if n == 0:
         return '#'
      if n == 1:
         temp = self.arr[0]
         self.arr.pop()
         return temp
      root = self.arr[0]
      self.arr[0] = self.arr[-1]
      self.arr.pop()
      self.heapify(0)
      return root

class Solution(object):
   def mergeKLists(self, lists):
      heap = Heap()
      for i in lists:
         if i:
            heap.insert(i)
      res = None
      res_next = None
      while True:
         temp = heap.extractMin()
         if temp == "#":
            return res
         if not res:
            res = temp
            res_next = temp
            temp = temp.next
            if temp:
               heap.insert(temp)
            res.next = None
         else:
            res_next.next = temp
            temp = temp.next
            res_next = res_next.next
            if temp:
               heap.insert(temp)
            res_next.next = None

ob = Solution()
lists = [[1,4,5],[1,3,4],[2,6]]
lls = []
for ll in lists:
   l = make_list(ll)
   lls.append(l)
print_list(ob.mergeKLists(lls))

실행 결과

입력

[[1,4,5],[1,3,4],[2,6]]

출력

[1, 1, 2, 3, 4, 4, 5, 6]

시간 복잡도 분석

이 알고리즘의 시간 복잡도는 O(N log k)입니다. 여기서 N은 모든 리스트에 포함된 노드의 총 개수, k는 리스트의 개수입니다. 각 노드가 힙에 한 번씩 삽입되고 추출되며, 힙 연산 하나당 O(log k)의 시간이 걸리기 때문입니다. 공간 복잡도는 힙에 동시에 존재하는 노드 수에 비례하므로 O(k)입니다.