75 lines
1.8 KiB
Python
75 lines
1.8 KiB
Python
class BinHeap:
|
|
def __init__(self, nodes):
|
|
self.heap = []
|
|
self.pos = [-1] * nodes
|
|
|
|
def sift_up(self, idx):
|
|
heap = self.heap
|
|
pos = self.pos
|
|
|
|
backup = heap[idx]
|
|
cur = idx
|
|
|
|
while cur > 0:
|
|
par = (cur - 1) // 2
|
|
if backup[1] >= heap[par][1]:
|
|
break
|
|
heap[cur] = heap[par]
|
|
pos[heap[cur][0]] = cur
|
|
cur = par
|
|
heap[cur] = backup
|
|
pos[backup[0]] = cur
|
|
|
|
def sift_down(self, idx):
|
|
heap = self.heap
|
|
pos = self.pos
|
|
|
|
n = len(heap)
|
|
|
|
backup = heap[idx]
|
|
cur = idx
|
|
|
|
while True:
|
|
left = cur * 2 + 1
|
|
if left >= n:
|
|
break
|
|
right = left + 1
|
|
tar = right if right < n and heap[right][1] < heap[left][1] else left
|
|
if backup[1] <= heap[tar][1]:
|
|
break
|
|
heap[cur] = heap[tar]
|
|
pos[heap[cur][0]] = cur
|
|
cur = tar
|
|
heap[cur] = backup
|
|
pos[backup[0]] = cur
|
|
|
|
def add(self, key, dist):
|
|
idx = len(self.heap)
|
|
self.heap.append([key, dist])
|
|
self.pos[key] = idx
|
|
self.sift_up(idx)
|
|
|
|
def extract_min(self):
|
|
heap = self.heap
|
|
pos = self.pos
|
|
if not heap:
|
|
return None
|
|
min_node, min_dist = heap[0]
|
|
pos[min_node] = -1
|
|
|
|
last = heap.pop()
|
|
if heap:
|
|
heap[0] = last
|
|
pos[last[0]] = 0
|
|
self.sift_down(0)
|
|
return (min_node, min_dist)
|
|
|
|
def decrease_key(self, key, new_dist):
|
|
idx = self.pos[key]
|
|
if idx == -1:
|
|
return None
|
|
node = self.heap[idx]
|
|
if new_dist >= node[1]:
|
|
return None
|
|
node[1] = new_dist
|
|
self.sift_up(idx) |