Day 083-Q2 — Radix Heap(最短路の整数ラベリング・$O(M + N \log W)$)

2026-07-06 赤色 Master / Phase 8+ ★★★★★★★★★ Dial's Algorithm・Radix Heap・bucket queue

問題

$N$ 頂点 $M$ 辺の重み付き有向グラフが与えられる。辺重みはすべて非負整数で最大値は $W$。頂点 $1$ から全頂点への最短距離を求めよ。

$W \le 10^6$, $N \le 10^6$, $M \le 5 \times 10^6$ であり、Radix Heap を用いて $O(M + N \log W)$ で解け。

制約

パラメータ範囲備考
$N$$\le 10^6$頂点数
$M$$\le 5 \times 10^6$辺数
$W$$\le 10^6$辺重みの最大値

入出力例

入力例1

5 7 100
1 2 10
1 3 3
2 3 1
3 2 4
2 4 2
3 5 8
4 5 5

出力例1

0 7 3 9 14

概念図: Radix Heap のバケット構造

Radix Heap — last=5 のときのバケット割り当て バケット B[0]: key==last B[1]: MSB(key XOR last)==0 B[2]: MSB(key XOR last)==1 B[3]: MSB(key XOR last)==2 B[k]: MSB(key XOR last)==k-1 last = 5 = 0b101 key binary XOR with 5 B 5 101 000 0 4 100 001 1 6 110 011 2 7 111 010 2 9 1001 1100 4 20 10100 10001 5 pop_min 操作: 1. B[0] が空なら最小バケットを探す 2. そのバケットの最小値を new_last に 3. バケット内要素を再分散してB[0]から取り出す 計算量比較 手法 計算量 空間 Binary Heap O(M log N) O(N) Dial's Alg O(M + NW) O(W) Radix Heap O(M + N log W) O(log W)

ヒント

ヒント1(方向性)

Radix Heap は距離値のビット表現に基づいてバケットを分割する。最後に取り出した最小値 $last$ との XOR の MSB 位置によってバケットが決まる。Dijkstra の単調増加性(pop する距離が非減少)を利用して各要素のバケット移動を $O(\log W)$ 以内に抑える。

ヒント2(アプローチ)

push(key, val): key == last なら B[0] へ、それ以外は msb(key XOR last) + 1 番目のバケットへ。

pop_min(): B[0] が空なら最小の空でないバケット $b$ を探し、その中の最小値を新たな last として全要素を再分散。

ヒント3(ほぼ答え)
def msb(x):
    return x.bit_length() - 1

class RadixHeap:
    K = 21
    def __init__(self):
        self.last = 0
        self.buckets = [[] for _ in range(self.K + 2)]
        self._size = 0

    def push(self, key, val):
        b = 0 if key == self.last else min(msb(key ^ self.last) + 1, self.K + 1)
        self.buckets[b].append((key, val))
        self._size += 1

    def pop_min(self):
        while not self.buckets[0]:
            for b in range(1, self.K + 2):
                if self.buckets[b]:
                    self.last = min(k for k, _ in self.buckets[b])
                    tmp = self.buckets[b]; self.buckets[b] = []
                    for k, v in tmp:
                        nb = 0 if k == self.last else min(msb(k ^ self.last)+1, self.K+1)
                        self.buckets[nb].append((k, v))
                    break
        self._size -= 1
        return self.buckets[0].pop()

模範解答

import sys
input = sys.stdin.readline

def msb(x):
    return x.bit_length() - 1

class RadixHeap:
    K = 21
    def __init__(self):
        self.last = 0
        self.buckets = [[] for _ in range(self.K + 2)]
        self._size = 0

    def push(self, key, val):
        if key == self.last:
            b = 0
        else:
            b = msb(key ^ self.last) + 1
            if b > self.K + 1:
                b = self.K + 1
        self.buckets[b].append((key, val))
        self._size += 1

    def pop_min(self):
        if self._size == 0:
            return None
        while not self.buckets[0]:
            for b in range(1, self.K + 2):
                if self.buckets[b]:
                    self.last = min(k for k, _ in self.buckets[b])
                    tmp = self.buckets[b]
                    self.buckets[b] = []
                    for item in tmp:
                        k, v = item
                        if k == self.last:
                            nb = 0
                        else:
                            nb = msb(k ^ self.last) + 1
                            if nb > self.K + 1:
                                nb = self.K + 1
                        self.buckets[nb].append(item)
                    break
        self._size -= 1
        return self.buckets[0].pop()

    def empty(self):
        return self._size == 0

def solve():
    N, M, W = map(int, input().split())
    graph = [[] for _ in range(N + 1)]
    for _ in range(M):
        u, v, w = map(int, input().split())
        graph[u].append((v, w))

    INF = float('inf')
    dist = [INF] * (N + 1)
    dist[1] = 0

    rh = RadixHeap()
    rh.push(0, 1)

    while not rh.empty():
        item = rh.pop_min()
        if item is None:
            break
        d, u = item
        if d > dist[u]:
            continue
        for v, w in graph[u]:
            nd = d + w
            if nd < dist[v]:
                dist[v] = nd
                rh.push(nd, v)

    ans = [str(dist[i]) if dist[i] != INF else '-1' for i in range(1, N + 1)]
    print(' '.join(ans))

solve()

Step-by-Step 解説

Step 1: Radix Heap の設計原理

距離値の単調増加性(Dijkstra の特性)を利用。最後に pop した値 last との XOR の MSB 位置でバケットを決定する。

バケット条件key の範囲
B[0]key == last単一値
B[k] (k≥1)msb(key XOR last) == k-1$[last + 2^{k-1}, last + 2^k)$

Step 2: push 操作 $O(1)$

def push(self, key, val):
    b = 0 if key == self.last else msb(key ^ self.last) + 1
    self.buckets[b].append((key, val))

Step 3: pop_min と再分散

B[0] が空の場合、最小の空でないバケットを見つけ、その中の最小値を新たな last として再分散する。各再分散で要素は必ず番号の小さいバケットに移動するため、全体の再分散コストは $O(N \log W)$。

Step 4: Dijkstra との統合

通常の Dijkstra と同様だが heappush/heappoprh.push/rh.pop_min に置き換えるだけ。

よくあるミス

ミス原因正しい書き方
msb(0) でエラーbit_length() は 0 のとき 0key == last の場合を先に処理
バケット数不足$W > 2^K$ の場合の overflowK = W.bit_length() + 1 以上
last 更新忘れ再分散時に最小値で更新が必要self.last = min(...)
Dial's Algorithm の空間不足バケット数 W+1 で MLERadix Heap で $O(\log W)$ バケット

次のステップ

  • 発展問題: 辺重みが $[0, W]$ で $W = O(N)$ のとき Dial's Algorithm が $O(N + M)$ になることを証明し実装せよ
  • 参考: Ahuja, Mehlhorn, Orlin, Tarjan (1990) "Faster Algorithms for the Shortest Path Problem"

自己評価

理解度: / /

自分の回答:

気づき・メモ: