Official

D - 通信ネットワークの構築 / Construction of a Communication Network Editorial by admin

claude4.8opus-high

概要

すべての拠点を連結にするケーブル集合のうち、「総コスト+\(K\times\)最大コスト」を最小化する問題です。実は最小全域木(MST)を求めるだけで解けます。

考察

採用する辺の集合は「全域木」でよい

まず、選ぶケーブル集合 \(S\) にサイクルがあると無駄です。サイクルを構成する辺を1本取り除いても連結性は保たれますが、辺を1本減らすので総コストは確実に下がり、最大コストも上がることはありません。つまり負荷指数は増えません。

したがって、最適な \(S\) は必ず「全域木(\(N-1\) 本の辺で全拠点を連結するもの)」になります。

全域木の中でどれを選ぶか

負荷指数は次の2つの項からなります。

\[ \underbrace{\sum_{i \in S} c_i}_{\text{総コスト}} \;+\; K \times \underbrace{\max_{i \in S} c_i}_{\text{最大コスト}} \]

一見すると「総コスト」と「最大コスト」のトレードオフを考えて、\(K\) の値に応じてどちらを優先するか調整する必要がありそうに見えます。しかしここに重要な性質があります。

最小全域木(MST)は、総コストを最小にすると同時に、含まれる辺の最大コストも最小にする

これはMSTが「最小ボトルネック全域木」でもあるという有名な性質です。つまりMSTを取れば、

  • 総コスト \(\sum c_i\) → 全域木の中で最小
  • 最大コスト \(\max c_i\) → 全域木の中で最小

両方が同時に達成されます

両方の項が同時に最小になるのですから、\(K \geq 0\) がどんな値であっても、その和である負荷指数もMSTで最小になります。よって \(K\) の値で場合分けする必要はなく、単純にMSTを1つ求めればよいのです。

アルゴリズム

クラスカル法でMSTを構築します。

  1. すべての辺をコスト \(c\) の昇順にソートする。
  2. コストの小さい辺から順に見て、その辺の両端がまだ連結でなければ採用する(Union-Find で連結判定)。
  3. 採用した辺で総コストを加算していく。コスト昇順に処理しているため、最後に採用した辺のコストがMSTの最大コストになる。
  4. 採用した辺が \(N-1\) 本になったら終了。

最後に「総コスト + \(K\times\) 最大コスト」を出力します。

具体例として、コスト \(1, 2, 5\) の辺でMSTが作れたとすると、総コストは \(1+2+5=8\)、最大コストは \(5\) なので、\(K=3\) なら負荷指数は \(8 + 3\times5 = 23\) となります。

計算量

  • 時間計算量: \(O(M \log M)\)(辺のソートが支配的。Union-Find の操作はほぼ定数)
  • 空間計算量: \(O(N + M)\)

実装のポイント

  • 辺を (c, u, v) のタプルにしてソートすると、コスト基準で簡単に昇順ソートできます。

  • コスト昇順で採用するため、採用するたびに mx = c と更新すれば、ループ終了時には自動的にMST内の最大コストが入っています(別途 max を計算する必要はありません)。

  • Union-Find は経路圧縮を入れて高速化しています。

  • \(N, M\) が最大 \(2\times10^5\) と大きいので、入力は sys.stdin.buffer.read() で一括読み込みすると高速です。

  • 答えは \(K\)\(c\) が大きい場合でも64bit整数の範囲に収まりますが、Python では多倍長整数なのでオーバーフローの心配は不要です。

    ソースコード

import sys

def main():
    data = sys.stdin.buffer.read().split()
    idx = 0
    N = int(data[idx]); idx += 1
    M = int(data[idx]); idx += 1
    K = int(data[idx]); idx += 1
    edges = []
    for _ in range(M):
        u = int(data[idx]); v = int(data[idx+1]); c = int(data[idx+2]); idx += 3
        edges.append((c, u, v))
    edges.sort()

    parent = list(range(N + 1))

    def find(x):
        while parent[x] != x:
            parent[x] = parent[parent[x]]
            x = parent[x]
        return x

    total = 0
    mx = 0
    cnt = 0
    for c, u, v in edges:
        ru = find(u)
        rv = find(v)
        if ru != rv:
            parent[ru] = rv
            total += c
            mx = c  # 昇順処理なので最後に採用した辺が最大辺
            cnt += 1
            if cnt == N - 1:
                break

    print(total + K * mx)

main()

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

posted:
last update: