Official

D - 届け物 / Delivery Editorial by admin

Claude 4.6 Opus (Thinking)

概要

地点 \(S\) から地点 \(G\) を経由して地点 \(T\) に到達する最短時間を求める問題です。\(S \to G\) の最短距離と \(G \to T\) の最短距離を足し合わせれば答えになります。

考察

高橋君は \(S \to G \to T\) の順に移動する必要があります。ここで重要な気づきは、経由地 \(G\) で経路を2つに分割できるということです。

  • 前半: \(S\) から \(G\) への最短経路
  • 後半: \(G\) から \(T\) への最短経路

同じ道路や地点を複数回通ってもよいと問題文にあるため、前半と後半の経路が重なっていても問題ありません。つまり、前半・後半をそれぞれ独立に最短距離を求め、その和を取れば最適解が得られます。

素朴なアプローチとの比較

\(S\) から \(G\) を経由して \(T\) に行く全経路を列挙する」ようなアプローチでは、組合せ爆発により到底間に合いません。しかし、上記の分割の観察により、最短経路問題を2回解くだけで十分です。

アルゴリズム

  1. グラフを隣接リストとして構築する。
  2. ダイクストラ法を用いて、始点 \(S\) から全頂点への最短距離 \(\mathrm{dist\_s}\) を計算する。
  3. 同様に、始点 \(G\) から全頂点への最短距離 \(\mathrm{dist\_g}\) を計算する。
  4. \(\mathrm{dist\_s}[G] + \mathrm{dist\_g}[T]\) が答え。ただし、どちらかが \(\infty\)(到達不能)の場合は \(-1\) を出力する。

具体例: \(S=1, G=3, T=5\) のとき、ダイクストラで \(1\) から \(3\) への最短距離が \(4\)\(3\) から \(5\) への最短距離が \(7\) と求まれば、答えは \(4 + 7 = 11\) です。

なぜダイクストラ法を2回で済むのか

\(G\) を始点としたダイクストラ1回で \(G \to T\) の最短距離が分かります。\(S\) を始点としたダイクストラ1回で \(S \to G\) の最短距離が分かります。合計2回のダイクストラで必要な情報がすべて揃います。

計算量

  • 時間計算量: \(O((N + M) \log N)\)
    • ダイクストラ法を2回実行するため \(O(2 \times (N + M) \log N)\) ですが、定数倍を無視すれば \(O((N + M) \log N)\) です。
  • 空間計算量: \(O(N + M)\)
    • グラフの隣接リストに \(O(N + M)\)、距離配列に \(O(N)\) を使用します。

実装のポイント

  • 入力の高速化: sys.stdin.buffer.read() でまとめて読み込み、split して処理することで Python でも十分高速に動作します。

  • 到達不能の判定: 距離が float('inf') のままであれば到達不能と判断し、\(-1\) を出力します。

  • 辺のコストが最大 \(10^9\)、頂点数が最大 \(10^5\) なので、最短距離の合計は最大で約 \(10^{14}\) 程度になり得ます。Python では整数オーバーフローの心配はありませんが、他の言語では 64 ビット整数型を使う必要があります。

  • グラフは双方向(無向グラフ)なので、各辺について両方向を隣接リストに追加します。

    ソースコード

import heapq
import sys

def dijkstra(graph, start, n):
    dist = [float('inf')] * (n + 1)
    dist[start] = 0
    pq = [(0, start)]
    while pq:
        d, u = heapq.heappop(pq)
        if d > dist[u]:
            continue
        for v, w in graph[u]:
            if dist[u] + w < dist[v]:
                dist[v] = dist[u] + w
                heapq.heappush(pq, (dist[v], v))
    return dist

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
    G = int(input_data[idx]); idx += 1
    T = int(input_data[idx]); idx += 1

    graph = [[] for _ in range(N + 1)]
    for _ in range(M):
        u = int(input_data[idx]); idx += 1
        v = int(input_data[idx]); idx += 1
        c = int(input_data[idx]); idx += 1
        graph[u].append((v, c))
        graph[v].append((u, c))

    dist_s = dijkstra(graph, S, N)
    dist_g = dijkstra(graph, G, N)

    sg = dist_s[G]
    gt = dist_g[T]

    if sg == float('inf') or gt == float('inf'):
        print(-1)
    else:
        print(sg + gt)

main()

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

posted:
last update: