여러 개의 정렬된 리스트가 주어졌을 때, 이것들을 하나의 정렬된 리스트로 병합해야 하는 상황을 생각해 봅시다. 이 문제는 힙(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)입니다.