公式

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

gemini-3.1-pro-thinking

概要

拠点 \(S\) から到達可能なすべての拠点への最短距離(最小の移動コスト)を求め、その総和を計算する問題です。グラフ理論における「単一始点最短経路問題」に帰着されます。

考察

問題文にある「拠点 \(S\) から拠点 \(v\) への全ての経路の移動コストのうち最小のもの」とは、グラフにおける始点 \(S\) から頂点 \(v\) への最短距離を意味しています。

最短距離を求めるアルゴリズムとして、単純な幅優先探索(BFS)がありますが、これはすべての道路のコストが同じ場合にしか使えません。今回は道路ごとにコスト \(W_i\) が異なるため、BFSでは正しい答えを導くことができません。

また、ベルマンフォード法というアルゴリズムもありますが、計算量が \(O(NM)\) となり、今回の制約(\(N, M \le 2 \times 10^5\))では実行時間制限(TLE)に引っかかってしまいます。

道路のコストがすべて正の値(\(1 \le W_i \le 10^4\))であることに注目すると、この問題はダイクストラ法(Dijkstra’s algorithm)を用いることで高速かつ正確に解くことができます。

アルゴリズム

ダイクストラ法は、「現在判明している中で最も近い頂点から順に距離を確定させていく」アルゴリズムです。効率よく「最も近い頂点」を見つけるために、優先度付きキュー(ヒープ)を使用します。

具体的な手順は以下の通りです: 1. 拠点 \(S\) から各拠点への最短距離を管理する配列 dist を用意し、初期値を非常に大きな値(\(\infty\))にします。拠点 \(S\) 自身の距離は dist[S] = 0 とします。 2. 優先度付きキューに、始点の情報 (距離 0, 拠点 S) を入れます。 3. キューが空になるまで以下を繰り返します: - キューから (暫定距離 d, 拠点 u) を取り出します。 - もし d が現在の dist[u] より大きければ、古い情報なので無視します(スキップ)。 - 拠点 \(u\) から繋がっている各拠点 \(v\) について、拠点 \(u\) を経由した場合の距離 nd = d + w を計算します。 - もし nd が現在の dist[v] より小さければ、最短距離を更新(dist[v] = nd)し、キューに (nd, 拠点 v) を追加します。 4. 全ての探索が終わった後、到達可能な拠点(dist\(\infty\) ではない、かつ \(S\) 以外の拠点)の距離の総和を計算して出力します。

計算量

  • 時間計算量: \(O((N + M) \log N)\)
    • 各頂点と辺を1回ずつ確認し、優先度付きキューへの追加・取り出しに \(O(\log N)\) かかるため、全体でこの計算量になります。最大ケースでも十分に高速です。
  • 空間計算量: \(O(N + M)\)
    • グラフを表現する隣接リストと、距離を管理する配列、優先度付きキューを保持するためのメモリ空間です。

実装のポイント

  • 優先度付きキューの利用 Pythonでは標準ライブラリの heapq を用います。heapq は常に最小の値を取り出せるデータ構造(最小ヒープ)なので、(距離, 頂点) のタプルを保存することで、自動的に距離が最小のものから取り出すことができます。

  • 枝刈り(不要な探索の省略) キューから取り出した際、if d > dist[u]: continue という処理を入れることが重要です。これにより、すでに最短距離が更新された後の「古い(長い)経路の情報」を処理する無駄を省き、実行時間を大幅に短縮できます。

  • 高速な入力の受け取り 入力サイズが非常に大きいため、input() を何度も呼ぶと遅くなります。sys.stdin.read().split() を使って入力を一括で読み込むことで、入力処理のボトルネックを解消しています。

    ソースコード

import sys
import heapq

def solve():
    input_data = sys.stdin.read().split()
    if not input_data:
        return
    
    N = int(input_data[0])
    M = int(input_data[1])
    S = int(input_data[2])
    
    adj = [[] for _ in range(N + 1)]
    
    idx = 3
    for _ in range(M):
        u = int(input_data[idx])
        v = int(input_data[idx+1])
        w = int(input_data[idx+2])
        adj[u].append((v, w))
        adj[v].append((u, w))
        idx += 3
        
    INF = 10**18
    dist = [INF] * (N + 1)
    dist[S] = 0
    
    hq = [(0, S)]
    
    while hq:
        d, u = heapq.heappop(hq)
        
        if d > dist[u]:
            continue
            
        for v, w in adj[u]:
            nd = d + w
            if nd < dist[v]:
                dist[v] = nd
                heapq.heappush(hq, (nd, v))
                
    ans = sum(d for i, d in enumerate(dist) if i != 0 and i != S and d != INF)
    print(ans)

if __name__ == '__main__':
    solve()

この解説は gemini-3.1-pro-thinking によって生成されました。

投稿日時:
最終更新: