Official

D - 救急搬送ネットワーク / Emergency Transport Network Editorial by admin

Claude 4.6 Opus (Thinking)

概要

救急医療センター(拠点 \(S\))から、到達可能な全ての拠点への最短移動コストの総和を求める問題です。これは典型的な単一始点最短経路問題であり、ダイクストラ法を用いて解くことができます。

考察

問題の本質

この問題は「拠点 \(S\) を始点として、他の全ての到達可能な拠点への最短距離を求め、それらを合計する」というものです。グラフは重み付き無向グラフであり、辺の重みは全て正(\(W_i \geq 1\))です。

素朴なアプローチの問題点

各拠点 \(v\) に対して個別にBFSや全経路探索で最短距離を求めると、非常に多くの経路を探索する必要があり、\(N\)\(M\) が最大 \(2 \times 10^5\) の制約では到底間に合いません。

また、辺に重みがあるため、重みなしグラフで使える単純なBFS(幅優先探索)では正しい最短距離が得られません。

解決策

辺の重みが全てであるため、ダイクストラ法が適用できます。ダイクストラ法を一度実行するだけで、始点 \(S\) から到達可能な全ての拠点への最短距離を同時に求めることができます。

アルゴリズム

  1. 隣接リストの構築: 各拠点について、隣接する拠点と辺の重みのペアをリストに格納します。
  2. ダイクストラ法の実行:
    • 距離配列 dist を全て \(\infty\) で初期化し、dist[S] = 0 とします。
    • 優先度付きキュー(最小ヒープ)に \((0, S)\) を入れます。
    • キューから最小コストの頂点 \((d, u)\) を取り出し、\(d > \text{dist}[u]\) なら(既により良い経路が見つかっているので)スキップします。
    • そうでなければ、\(u\) の隣接頂点 \(v\) に対して、\(d + w < \text{dist}[v]\) ならば dist[v] を更新し、キューに \((d + w, v)\) を追加します。
  3. 総和の計算: 全ての拠点 \(i\)\(i \neq S\))について、dist[i]\(\infty\) でなければ(到達可能であれば)その値を合計します。

具体例

例えば、\(N=3, M=2, S=1\) で、辺が \((1,2,5)\)\((2,3,3)\) の場合: - \(d(2) = 5\)(拠点1→拠点2、コスト5) - \(d(3) = 8\)(拠点1→拠点2→拠点3、コスト5+3=8) - 答え:\(5 + 8 = 13\)

計算量

  • 時間計算量: \(O((N + M) \log N)\)
    • ダイクストラ法において、各頂点は最大1回確定処理され、各辺は最大1回緩和に使われます。優先度付きキューの操作(挿入・取り出し)がそれぞれ \(O(\log N)\) かかるため、全体で \(O((N + M) \log N)\) です。
  • 空間計算量: \(O(N + M)\)
    • 隣接リスト、距離配列、優先度付きキューの合計です。

実装のポイント

  • sys.stdin.buffer.read() を使って入力を一括読み込みすることで、Python でも高速に入力処理ができます。

  • 距離の初期値に float('inf') を使うことで、未到達の頂点を簡潔に判定できます。

  • if d > dist[u]: continue の判定(遅延削除)は、既に最短距離が確定した頂点を再処理しないための重要な高速化テクニックです。これがないと、同じ頂点を何度も処理してしまい TLE になる恐れがあります。

  • 拠点番号が \(1\) から \(N\) なので、隣接リストや距離配列のサイズを \(N+1\) にしています。

    ソースコード

import heapq
import sys

def main():
    input_data = sys.stdin.buffer.read().split()
    idx = 0
    N = int(input_data[idx]); idx += 1
    M = int(input_data[idx]); idx += 1
    S = int(input_data[idx]); idx += 1
    
    adj = [[] for _ in range(N + 1)]
    for _ in range(M):
        u = int(input_data[idx]); idx += 1
        v = int(input_data[idx]); idx += 1
        w = int(input_data[idx]); idx += 1
        adj[u].append((v, w))
        adj[v].append((u, w))
    
    INF = float('inf')
    dist = [INF] * (N + 1)
    dist[S] = 0
    pq = [(0, S)]
    
    while pq:
        d, u = heapq.heappop(pq)
        if d > dist[u]:
            continue
        for v, w in adj[u]:
            nd = d + w
            if nd < dist[v]:
                dist[v] = nd
                heapq.heappush(pq, (nd, v))
    
    ans = 0
    for i in range(1, N + 1):
        if i != S and dist[i] != INF:
            ans += dist[i]
    
    print(ans)

main()

この解説は claude4.6opus-thinking によって生成されました。

posted:
last update: