公式

D - 荷物の配送 / Package Delivery 解説 by admin

Claude 4.6 Opus (Thinking)

概要

\(M\) 個の拠点からなるグラフ上で、\(N\) 件の配送依頼それぞれについて最短距離を求め、その合計を出力する問題です。\(M\) が小さい(最大200)ことがポイントで、全拠点間の最短距離を事前に計算しておくことで効率的に解けます。

考察

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

各配送依頼ごとにダイクストラ法で最短距離を求める方法が考えられます。ダイクストラ法1回の計算量は \(O(K \log M)\) 程度なので、\(N\) 件すべてに対して実行すると \(O(N \cdot K \log M)\) となります。\(N\) が最大 \(10^5\) なので、同じ出発拠点の依頼をまとめるなどの工夫をしないと無駄な計算が多くなります。

重要な気づき:\(M\) が小さい

この問題では \(M \leq 200\) という制約があります。これは「全頂点対間最短距離」を求めるアルゴリズムである ワーシャル–フロイド法(Floyd-Warshall) が非常に適している条件です。

ワーシャル–フロイド法を使えば、\(O(M^3)\) で全拠点ペア間の最短距離をあらかじめ計算できます。\(M = 200\) のとき \(M^3 = 8 \times 10^6\) であり、十分高速です。

事前計算が終われば、各配送依頼に対してはテーブルを参照するだけ(\(O(1)\))で最短距離が分かります。

アルゴリズム

ワーシャル–フロイド法

全頂点対間の最短距離を求めるアルゴリズムです。以下の手順で動作します。

  1. 初期化: \(M \times M\) の距離テーブル dist を用意する。

    • dist[i][i] = 0(自分自身への距離は0)
    • 辺 \((u, v)\) が重み \(w\) で存在するなら dist[u][v] = dist[v][u] = w
    • それ以外は dist[i][j] = ∞
  2. 3重ループによる更新: 中継点 \(k\) を \(0\) から \(M-1\) まで順に試し、すべての頂点ペア \((i, j)\) に対して以下を更新する。 $\(\text{dist}[i][j] = \min(\text{dist}[i][j],\; \text{dist}[i][k] + \text{dist}[k][j])\)\( これは「頂点 \)i\( から \)j\( へ行くのに、頂点 \)k$ を経由した方が短いか?」を確認しています。

  3. クエリ応答: 各配送依頼 \((S_i, T_i)\) について dist[S_i][T_i] を参照し、合計を求める。

具体例

拠点が3つ、道路が「1-2(距離3)」「2-3(距離5)」「1-3(距離10)」の場合:

  • 初期テーブルでは dist[1][3] = 10
  • 中継点 \(k = 2\) を経由すると dist[1][2] + dist[2][3] = 3 + 5 = 8 < 10
  • よって dist[1][3] = 8 に更新される

計算量

  • 時間計算量: \(O(M^3 + N)\)
    • ワーシャル–フロイド法に \(O(M^3)\)、\(N\) 件のクエリ応答に \(O(N)\)
    • \(M \leq 200\) なので \(M^3 = 8 \times 10^6\) で十分高速
  • 空間計算量: \(O(M^2)\)
    • 距離テーブルの保持に必要

実装のポイント

  • 0-indexed への変換: 入力の拠点番号は1始まりなので、内部では1を引いて0始まりに統一しています。

  • 高速な入力処理: sys.stdin.buffer.read() で一括読み込みし、split() で分割することで、Python でも入力がボトルネックになりにくくしています。

  • 枝刈り: dist[i][k] == INF のとき内側のループをスキップすることで、不要な計算を省いています。

  • 変数のローカル参照: dk = dist[k]、di = dist[i] のようにリストの参照をローカル変数に取ることで、Python のリストアクセスを高速化しています。

    ソースコード

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

    INF = float('inf')
    dist = [[INF] * M for _ in range(M)]
    for i in range(M):
        dist[i][i] = 0

    for _ in range(K):
        u = int(input_data[idx]) - 1; idx += 1
        v = int(input_data[idx]) - 1; idx += 1
        w = int(input_data[idx]); idx += 1
        if w < dist[u][v]:
            dist[u][v] = w
            dist[v][u] = w

    # Floyd-Warshall
    for k in range(M):
        dk = dist[k]
        for i in range(M):
            dik = dist[i][k]
            if dik == INF:
                continue
            di = dist[i]
            for j in range(M):
                val = dik + dk[j]
                if val < di[j]:
                    di[j] = val

    total = 0
    for _ in range(N):
        s = int(input_data[idx]) - 1; idx += 1
        t = int(input_data[idx]) - 1; idx += 1
        total += dist[s][t]

    print(total)

main()

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

投稿日時:
最終更新: