Files
2026-03-02 20:55:36 +09:00

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)