公式

E - 通信ネットワークの妨害 / Disruption of Communication Network 解説 by kyopro_friends


青木君が遮断する回線を \(e\) 、高橋君が侵入する基地局を \(v\) としたときの結果を \(f(e,v)\) とします。全ての \((e,v)\) についての \(f(e,v)\) が求まっていれば、答えである \(\min_e \max_v f(e,v)\)\(O(N^2)\) で求めることができます。

\(v\) を固定したとき、全ての \(e\) についての \(f(e,v)\) を高速に求めることを考えます。\(v\) を根としてDFSすることで、全ての頂点 \(v'\) について「\(v'\) の部分木のうち\(v\) からの距離が \(D\) 以下である頂点のデータの合計量」を求めることができます。これを \(g_v(v')\) とします。

回線が遮断されないときに盗み出せるデータの合計量は \(g_v(v)\) であり、 \(v'\) とその親を結ぶ回線 \(e_{v'}\) を遮断したときに盗めなくなるデータの合計量は \(g_v(v')\) であることから \(f(e_{v'},v)=g_v(v)-g_v(v')\) がわかります。よって全ての \(e\) に対する \(f(e,v)\)\(O(N)\) で求めることができました。

以上を全ての \(v\) に対して行うことで、全ての \(f(e,v)\)\(O(N^2)\) で求めることができるため、全体で \(O(N^2)\) で答えを求めることができます。

実装例 (C++)

#include<bits/stdc++.h>
using namespace std;
using ll = long long;

int main(){
  int n, d;
  cin >> n >> d;
  vector<int> V(n);
  for(int i=0; i<n; i++) cin >> V[i];
  vector<vector<pair<int,int>>> G(n);
  for(int i=0; i<n-1; i++){
    int a, b;
    cin >> a >> b;
    a--, b--;
    G[a].push_back({b,i});
    G[b].push_back({a,i});
  }

  vector<ll> gv(n);
  vector<int> ev(n);
  auto dfs = [&](auto self, int v, int d, int pre) -> void{
    ll ret = 0;
    for(auto[vv, i]: G[v]){
      if(vv == pre){
        ev[v] = i;
      }else{
        self(self, vv, d-1, v);
        if(d > 0){
          ret += gv[vv];
        }
      }
    }
    if(d >= 0){
      ret += V[v];
    }
    gv[v] = ret;
  };

  vector<ll> ans(n-1);  // ans[i] = max_v f(e[i],v)
  for(int v=0; v<n; v++){
    dfs(dfs, v, d, -1);
    for(int vv=0; vv<n; vv++){
      if(v == vv){
        continue;
      }
      ans[ev[vv]] = max(ans[ev[vv]], gv[v] - gv[vv]);
    }
  }
  cout << *min_element(ans.begin(), ans.end()) << endl;
}

実装例 (Python)

import sys
sys.setrecursionlimit(10**9)

N, D = map(int, input().split())
V = list(map(int, input().split()))
G = [[] for _ in range(N)]
for i in range(N-1):
  A, B = map(int, input().split())
  A -= 1
  B -= 1
  G[A].append((B, i))
  G[B].append((A, i))

gv = [0] * N
ev = [-1] * N
def dfs(v, d, pre):
  ret = 0
  for vv, i in G[v]:
    if vv == pre:
      ev[v] = i
    else:
      dfs(vv, d-1, v)
      if d > 0:
        ret += gv[vv]
  if d >= 0:
    ret += V[v]
  gv[v] = ret

ans = [0] * (N-1)  # ans[i] = max_v f(e[i],v)
for v in range(N):
  dfs(v, D, -1)
  for vv in range(N):
    if v == vv:
      continue
    ans[ev[vv]] = max(ans[ev[vv]], gv[v] - gv[vv])

print(min(ans))

投稿日時:
最終更新: