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))
投稿日時:
最終更新:
