Official

E - 通信ネットワークの妨害 / Disruption of Communication Network Editorial by admin

gemini-3.5-flash-high

概要

この問題は、木構造の通信ネットワークにおいて、青木君が回線を1本遮断(または遮断しない)した後に、高橋君が任意の基地局から距離 \(D\) 以内の基地局のデータを最大化するように盗み出すゲームです。青木君は高橋君が盗み出せるデータの最大値を最小化しようとします。

両者が最適に行動したときの結果を、全点対最短経路と累積和を用いた \(O(N^2)\) の動的計画法(木DPの応用)によって求めます。


考察

1. ゲームの構造の整理

青木君が遮断できる回線(辺)は高々1本です。 ある辺 \(e = (u, v)\)\(u\)\(v\) の親)を遮断したとします。このとき、木は以下の2つの部分に分断されます。 - 頂点 \(v\) の部分木(以下、IN側) - それ以外の部分(以下、OUT側

高橋君は、遮断された後のネットワークで、最も多くのデータを盗み出せる基地局(侵入先)を最適に選びます。 - もし高橋君が IN側 の頂点 \(x\) に侵入した場合、獲得できる最大データ量を \(M_{in}[v]\) とします。 - もし高橋君が OUT側 の頂点 \(y\) に侵入した場合、獲得できる最大データ量を \(M_{out}[v]\) とします。

高橋君はこれらの中から最大のものを選択するため、辺 \((u, v)\) を遮断したときに高橋君が獲得するデータ量は \(\max(M_{in}[v], M_{out}[v])\) となります。

青木君はこれを最小化したいので、すべての辺についてこの値を計算し、その最小値を求めます。また、回線を「遮断しない」という選択肢もあるため、初期状態(遮断なし)での高橋君の最大獲得データ量との最小値をとったものが最終的な答えになります。

2. 愚直な方法とその限界

各辺を遮断したグラフに対して、毎回すべての頂点からBFS(幅優先探索)を行って距離 \(D\) 以内のデータ量を求めると、1つの辺の遮断につき \(O(N^2)\)、全体で \(O(N^3)\) の計算量がかかります。\(N \le 3000\) では \(N^3 \approx 2.7 \times 10^{10}\) となり、実行時間制限に間に合いません(TLE)。

3. 高速化のアイデア(差分の利用)

遮断がない状態での「頂点 \(i\) から距離 \(D\) 以内のデータ量の総和」を \(S_{all\_D}[i]\) とします。これは事前に \(O(N^2)\) で計算可能です。

\((u, v)\) を遮断したとき、ある頂点から「届かなくなる(引かれる)データ量」のみを差分として \(O(1)\) で計算できれば、各辺のシミュレーションを大幅に高速化できます。


アルゴリズム

① 事前準備(オイラーツアーと最短距離)

  1. 適当な頂点(例えば頂点 \(0\))を根として木を構築します。
  2. 行きがけ順(トポロジカル順序)の配列 order を作成します。これにより、各頂点 \(u\) の部分木に属する頂点集合は、order 上の連続する区間 \([L[u], R[u])\) として表すことができます。
  3. 全点対最短距離 \(dist[u][v]\) を、各頂点からBFSを行うことで \(O(N^2)\) で求めておきます。

② 累積和の定義

以下の2つの累積和テーブルを \(O(N^2)\) で前処理しておきます。

  • \(S_{all}[u][k]\) : 頂点 \(u\) から距離 \(k\) 以下のすべての頂点のデータ量の総和。
  • \(S_{in}[v][k]\) : 頂点 \(v\) から距離 \(k\) 以下の\(v\) の部分木内(IN側)」の頂点のデータ量の総和。

③ 遮断による影響の差分計算

\((u, v)\)\(u\)\(v\) の親)を遮断したときの影響を考えます。

(A) IN側の頂点 \(x\) に侵入する場合

頂点 \(x\) から本来届くはずだったOUT側の頂点へのアクセスが、辺 \((u, v)\) の遮断によって妨げられます。 - \(x\) から \(v\) までの距離は \(dist[x][v]\) です。 - 遮断された辺を越えて \(u\) に到達した時点での、残り移動可能距離は \(rem = D - 1 - dist[x][v]\) となります。 - \(rem \ge 0\) のとき、本来届くはずだった「\(u\) から距離 \(rem\) 以内にあるOUT側の頂点」に届かなくなります。 - この「失うデータ量」は、以下のように計算できます: $\(\text{失うデータ量} = S_{all}[u][rem] - S_{in}[v][rem - 1]\)\( (\)u\( から距離 \)rem\( 以内の全頂点から、IN側にあって \)u\( から距離 \)rem\( 以内(= \)v\( から距離 \)rem-1$ 以内)の頂点を除いたもの)

したがって、獲得できるデータ量は以下になります: $\(\text{獲得データ量} = S_{all\_D}[x] - (S_{all}[u][rem] - S_{in}[v][rem - 1])\)$

(B) OUT側の頂点 \(y\) に侵入する場合

頂点 \(y\) から本来届くはずだったIN側の頂点へのアクセスが妨げられます。 - \(y\) から \(u\) までの距離は \(dist[y][u]\) です。 - 遮断された辺を越えて \(v\) に到達した時点での、残り移動可能距離は \(rem = D - 1 - dist[y][u]\) となります。 - \(rem \ge 0\) のとき、本来届くはずだった「\(v\) から距離 \(rem\) 以内にあるIN側の頂点」に届かなくなります。 - この「失うデータ量」は、定義より \(S_{in}[v][rem]\) そのものです。

したがって、獲得できるデータ量は以下になります: $\(\text{獲得データ量} = S_{all\_D}[y] - S_{in}[v][rem]\)$

④ 集計

各辺 \((u, v)\) について、IN側のすべての頂点 \(x\) に対する (A) の最大値 \(M_{in}[v]\) と、OUT側のすべての頂点 \(y\) に対する (B) の最大値 \(M_{out}[v]\) を求めます。 青木君がこの辺を遮断したときの高橋君の最大データ量は \(\max(M_{in}[v], M_{out}[v])\) となります。

すべての辺についてこの値を求め、それらの最小値(および遮断しない場合の最大値)が求める答えとなります。


計算量

時間計算量: \(O(N^2)\)

  • 全点対最短距離の計算(BFSを \(N\) 回): \(O(N^2)\)
  • 累積和 \(S_{all}, S_{in}\) の計算: \(O(N^2)\)
  • 各辺 \(v\) について、IN側の頂点(\(sz[v]\) 個)とOUT側の頂点(\(N - sz[v]\) 個)に対して \(O(1)\) の計算を行うため、1つの辺あたり \(O(N)\)。全体で \(O(N^2)\)
  • したがって、全体の時間計算量は \(O(N^2)\) となり、制限時間内に十分間に合います。

空間計算量: \(O(N^2)\)

  • 最短距離テーブル dist および累積和テーブル S_all, S_in を保持するために \(O(N^2)\) のメモリを使用します。\(N = 3000\) のとき、各テーブルの要素数は約 \(9 \times 10^6\) であり、メモリ制限(通常 1024MB 等)に余裕で収まります。

実装のポイント

  1. 部分木の判定(オイラーツアーの応用): 頂点 \(v\) の部分木に属する頂点(IN側)と属さない頂点(OUT側)の区別は、行きがけ順のインデックス \(L[v]\) と部分木サイズ \(sz[v]\) を用いて、インデックスが \([L[v], L[v] + sz[v])\) の範囲内にあるかどうかで高速に判定できます。

  2. インデックスの境界条件: 残り距離 \(rem\)\(0\) 未満になる場合や、累積和のインデックスが配列の範囲外(\(rem - 1 < 0\) など)になる場合の例外処理を丁寧に行う必要があります。提供されたコードでは、rem >= 0 の条件分岐や rem > 0 の三項演算子を用いてこれらを安全に処理しています。

    ソースコード

import sys


def solve():
    input = sys.stdin.read
    data = input().split()
    if not data:
        return

    N = int(data[0])
    D = int(data[1])

    V = [int(x) for x in data[2 : N + 2]]

    adj = [[] for _ in range(N)]
    idx = N + 2
    for _ in range(N - 1):
        u = int(data[idx]) - 1
        v = int(data[idx + 1]) - 1
        adj[u].append(v)
        adj[v].append(u)
        idx += 2

    # BFS to find parent and topological order
    parent = [-1] * N
    order = []
    stack = [0]
    visited = [False] * N
    visited[0] = True
    while stack:
        u = stack.pop()
        order.append(u)
        for v in adj[u]:
            if not visited[v]:
                visited[v] = True
                parent[v] = u
                stack.append(v)

    L = [0] * N
    for i, u in enumerate(order):
        L[u] = i

    sz = [1] * N
    for u in reversed(order):
        p = parent[u]
        if p != -1:
            sz[p] += sz[u]

    R = [L[u] + sz[u] for u in range(N)]

    # All pairs shortest path
    dist = [[-1] * N for _ in range(N)]
    for start in range(N):
        d = dist[start]
        d[start] = 0
        queue = [start]
        head = 0
        while head < len(queue):
            u = queue[head]
            head += 1
            curr_d = d[u]
            for v in adj[u]:
                if d[v] == -1:
                    d[v] = curr_d + 1
                    queue.append(v)

    # S_all[u][k]
    S_all = [[0] * (N + 1) for _ in range(N)]
    for u in range(N):
        su = S_all[u]
        du = dist[u]
        for v in range(N):
            su[du[v]] += V[v]
        for k in range(1, N + 1):
            su[k] += su[k - 1]

    # Precompute S_all_D
    S_all_D = [S_all[i][D] for i in range(N)]

    # S_in[v][k]
    S_in = [[0] * (N + 1) for _ in range(N)]
    for v in range(N):
        sv = S_in[v]
        dv = dist[v]
        l = L[v]
        r = R[v]
        for idx in range(l, r):
            x = order[idx]
            sv[dv[x]] += V[x]
        for k in range(1, N + 1):
            sv[k] += sv[k - 1]

    M_in = [0] * N
    M_out = [0] * N

    for v in range(1, N):
        u = parent[v]
        dv = dist[v]
        su = S_all[u]
        sv = S_in[v]

        # M_in[v]
        m_in = 0
        l = L[v]
        r = R[v]
        for idx in range(l, r):
            x = order[idx]
            rem = D - 1 - dv[x]
            val = S_all_D[x]
            if rem >= 0:
                val -= su[rem] - (sv[rem - 1] if rem > 0 else 0)
            if val > m_in:
                m_in = val
        M_in[v] = m_in

        # M_out[v]
        m_out = 0
        du = dist[u]
        for idx in range(0, l):
            y = order[idx]
            rem = D - 1 - du[y]
            val = S_all_D[y]
            if rem >= 0:
                val -= sv[rem]
            if val > m_out:
                m_out = val
        for idx in range(r, N):
            y = order[idx]
            rem = D - 1 - du[y]
            val = S_all_D[y]
            if rem >= 0:
                val -= sv[rem]
            if val > m_out:
                m_out = val
        M_out[v] = m_out

    ans = max(S_all_D)
    for v in range(1, N):
        cand = M_in[v]
        if M_out[v] > cand:
            cand = M_out[v]
        if cand < ans:
            ans = cand

    print(ans)


if __name__ == "__main__":
    solve()

この解説は gemini-3.5-flash-high によって生成されました。

posted:
last update: