公式

D - 中継都市の重要度 / Importance of Relay Cities 解説 by admin

GLM 5.2 (High, OpenRouter)

概要

すべての頂点間の最短距離を求めた上で、各頂点 \(k\) について「その頂点 \(k\) を内部の頂点として経由するような最短経路の始点・終点のペア \((i, j)\) の個数」を数える問題です。

考察

まず、「都市 \(i\) から都市 \(j\) への最短経路のうち、都市 \(k\) を内部の頂点として含むものが少なくとも1つ存在する」という条件を考えます。

都市 \(k\) を経由して \(i\) から \(j\) へ行く経路の距離は、\(i \to k\) の最短距離と \(k \to j\) の最短距離の和、すなわち \(d(i, k) + d(k, j)\) となります。この経路が \(i \to j\) の最短経路になるための条件は、以下の等式が成り立つことです。 $\( d(i, k) + d(k, j) = d(i, j) \)\( この等式が成り立つとき、\)i \to k\( の最短経路と \)k \to j\( の最短経路を繋げたものが \)i \to j\( の最短経路となります。これにより、\)k\( を内部の頂点として含む最短経路が必ず存在することになります。逆に、この等式が成り立たない場合は \)k$ を経由するとどうしても最短経路より距離が長くなってしまうため、条件を満たす最短経路は存在しません。

したがって、各頂点 \(k\) について、すべての順序付きペア \((i, j)\)\(i \neq j \neq k\))を調べ、\(d(i, k) + d(k, j) = d(i, j)\) となるペアの数を数えればよいことがわかります。

\(N \leq 250\) と小さいため、すべての頂点間の最短距離を求めるワーシャルフロイド法を用いて \(O(N^3)\) で前計算しておき、その後各 \(k\) について \(O(N^2)\) の判定を行う全体 \(O(N^3)\) のアプローチで十分に間に合います。

アルゴリズム

  1. ワーシャルフロイド法: すべての頂点間の最短距離 \(d[i][j]\) を求めます。初期状態では、直接道路がある場合はその通行料、ない場合は無限大(\(\infty\))とします。
  2. 重要度の計算: 各都市 \(k\) について、以下を行います。
    • すべての都市 \(i, j\)\(i \neq k, j \neq k, i \neq j\))についてループを回します。
    • \(d[i][k]\)\(d[k][j]\) がどちらも \(\infty\) でない(経路が存在する)ことを確認します。
    • \(d[i][k] + d[k][j] == d[i][j]\) が成り立つ場合、ペア \((i, j)\) は都市 \(k\) の重要度に寄与するのでカウントを \(1\) 増やします。
    • すべての \(i, j\) について調べ終わったら、カウントした値を都市 \(k\) の重要度として出力します。

計算量

  • 時間計算量: \(O(N^3)\)
    • ワーシャルフロイド法による最短距離計算が \(O(N^3)\)
    • \(k\) について全ペア \((i, j)\) を調べる部分が \(O(N^3)\)
    • \(N \leq 250\) なので、全体で約 \(1.5 \times 10^7\) 回の操作となり、Pythonでも十分に実行制限時間内に終わります。
  • 空間計算量: \(O(N^2)\)
    • 最短距離テーブル \(d\) を保持するため。

実装のポイント

Pythonで \(O(N^3)\) の計算を実行する際、そのまま3重ループを書くと定数倍が重くTLE(Time Limit Exceeded)する可能性があります。そこで以下の工夫をしています。 - 内側のループで使う配列を変数にキャッシュして参照を高速化します(例: di = d[i], dk = d[k])。 - 距離が \(\infty\) の場合は早期に continue することで無駄な加算・比較をスキップします。 - リスト内の要素に何度もアクセスするのを防ぐため、ループの外に変数として取り出しておきます(例: dik = d[i][k], dkj = dk[j])。

ソースコード

import sys

def main():
    input = sys.stdin.readline
    N, M = map(int, input().split())
    INF = float('inf')
    d = [[INF] * (N + 1) for _ in range(N + 1)]
    for i in range(1, N + 1):
        d[i][i] = 0
    for _ in range(M):
        u, v, w = map(int, input().split())
        d[u][v] = w
    # Floyd-Warshall
    for k in range(1, N + 1):
        dk = d[k]
        for i in range(1, N + 1):
            dik = d[i][k]
            if dik == INF:
                continue
            di = d[i]
            dik_plus = dik
            for j in range(1, N + 1):
                v2 = dik_plus + dk[j]
                if v2 < di[j]:
                    di[j] = v2
    out = []
    for k in range(1, N + 1):
        dk = d[k]
        cnt = 0
        for i in range(1, N + 1):
            if i == k:
                continue
            dik = d[i][k]
            if dik == INF:
                continue
            di = d[i]
            for j in range(1, N + 1):
                if j == k or j == i:
                    continue
                dkj = dk[j]
                if dkj == INF:
                    continue
                if dik + dkj == di[j]:
                    cnt += 1
        out.append(str(cnt))
    sys.stdout.write('\n'.join(out))

main()

この解説は or-glm-5.2-high によって生成されました。

投稿日時:
最終更新: