Official
D - 荷物の配送 / Package Delivery Editorial by admin
Qwen3-Coder-480B概要
配送依頼ごとの出発地から目的地までの最短距離を求め、その合計を答える問題。
考察
この問題では、複数の配送依頼に対してそれぞれ最短距離を求める必要があります。
最も素朴な方法は、各配送依頼ごとにダイクストラ法を実行することですが、配送依頼の数 \(N\) は最大で \(10^5\) あり、毎回ダイクストラ法を実行していると時間制限に間に合わない可能性があります(計算量的には \(O(N \cdot M \log M)\) となり、最大で \(10^5 \times 200 \log 200\) 程度)。
しかし、拠点の数 \(M\) はたかだか \(200\) と小さいことに注目します。
つまり、全頂点(拠点)から他の全頂点への最短距離を事前に計算しておくことが現実的です。
このような「全点対間の最短距離」を求めるには、各頂点を始点としてダイクストラ法を実行する方法が有効です。
これにより、前処理で \(O(M \cdot M \log M)\) 程度で全ての最短距離を求めておき、各配送依頼に対しては \(O(1)\) で答えを取得できます。
アルゴリズム
- グラフを隣接リスト形式で構築する(双方向なので両方登録)。
- 各頂点 \(i\) (拠点)を始点としたダイクストラ法を実行し、
all_dist[i][j]:= 頂点 \(i\) から \(j\) への最短距離
を前計算して保存しておく。 - 各配送依頼 \((S_i, T_i)\) に対して、
all_dist[S_i][T_i]を合計に加算する。 - 合計を出力する。
ダイクストラ法の実装では、Pythonの heapq を用いることで効率的に最小コストの頂点を取り出すことができます。
計算量
- 時間計算量: \(O(M \cdot (M + K) \log M + N)\)
- 各頂点からのダイクストラ: \(M\) 回 × \(O((M + K) \log M)\)
- 各配送依頼の処理: \(O(N)\)
- 空間計算量: \(O(M^2 + M + K)\)
all_dist: \(M \times M\)- グラフの隣接リスト: \(O(M + K)\)
実装のポイント
- 拠点番号は 1-indexed なので、配列のサイズは
M+1にしておく。 - 双方向の道路をグラフに追加するときに両方向を忘れずに登録する。
- ダイクストラ法の初期化時に、距離配列は
float('inf')で初期化し、始点だけ 0 にする。 - 入力を高速に読み込むために
sys.stdin.readを使用している(必須ではないが推奨)。
## ソースコード
```python
import heapq
import sys
from collections import defaultdict
def dijkstra(graph, start, M):
dist = [float('inf')] * (M + 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 = sys.stdin.read
data = input().split()
idx = 0
N = int(data[idx]); idx += 1
M = int(data[idx]); idx += 1
K = int(data[idx]); idx += 1
graph = defaultdict(list)
for _ in range(K):
u = int(data[idx]); idx += 1
v = int(data[idx]); idx += 1
w = int(data[idx]); idx += 1
graph[u].append((v, w))
graph[v].append((u, w))
# 全点対最短経路を前計算するため、各頂点からダイクストラを行う
all_dist = [None] * (M + 1)
for i in range(1, M + 1):
all_dist[i] = dijkstra(graph, i, M)
total = 0
for _ in range(N):
s = int(data[idx]); idx += 1
t = int(data[idx]); idx += 1
total += all_dist[s][t]
print(total)
if __name__ == "__main__":
main()
この解説は qwen3-coder-480b によって生成されました。
posted:
last update: