第九届河北省大学生程序设计竞赛 B题思路分享(点分治)

题意概述

给定一棵树,\(n\) 个节点,其中 \(m\) 个节点上有石楠花,可以移除最多 \(k\) 个石楠花节点。

求最终树上的点与最近的石楠花节点距离的最大值。

\(1 \leq m \leq n \leq 10^5, 0 \leq k < m\)

思路

考虑二分答案 \(mid\)

只要存在一个节点,可以通过删除石楠花节点,让所有距离该节点 \(\le mid\) 的节点都不是石楠花节点,该 \(mid\) 可行。

实现上,对于每个石楠花节点,把距离该节点 \(\le mid\) 的节点权值 \(+1\),如果存在某个节点权值 \(\le k\) ,即可行。该操作可以通过点分治实现。

处理某个重心 \(r\),按距离分组 \(dfs\) 收集所有节点到 \(wait\),同时记录每个距离的石楠花节点数量。

由于操作是相互的,需要考虑容斥。因此一次 \(dfs\) 之后只做信息合并,不计算。

先不考虑重复部分,双指针给所有 \(wait\) 中元素计算贡献。

注意到只有同一子树内会产生重复,可以清空信息,对一棵子树 \(dfs\) 后,再次双指针给 \(wait\) 中所有元素减去产生的贡献。

时间复杂度 \(\mathcal{O}(n \log^2 n)\)

代码

//author:kzssCCC

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


void solve(){
	int n,m,k;
	cin >> n >> m >> k;

	vector<vector<int>> adj(n+1);
	for (int i=0;i<n-1;i++){
		int u,v;
		cin >> u >> v;

		adj[u].push_back(v);	
		adj[v].push_back(u);	
	}

	vector<bool> f(n+1,false);

	for (int i=1;i<=m;i++){
		int v;
		cin >> v;
		f[v] = true;
	}

	auto check = [&](int mid){
		vector<int> sz(n+1);
		vector<bool> vis(n+1,false);
		vector<int> cnt(n+1);
		vector<vector<int>> wait(mid+1);
		vector<int> dp(mid+1),pre(mid+1);

		int mxd = -1;
		
		function<void(int,int)> getsz = [&](int u,int par){
			sz[u] = 1;

			for (auto& v:adj[u]){
				if (v==par || vis[v]) continue;

				getsz(v,u);
				sz[u] += sz[v];
			}
		};

		auto getsent = [&](int u){
			getsz(u,-1);
			int par = -1;
			int half = sz[u]>>1;

			while (1){
				bool ok = true;

				for (auto& v:adj[u]){
					if (v==par || vis[v]) continue;

					if (sz[v]>half){
						ok = false;
						par = u;
						u = v;
						break;
					}	
				}

				if (ok) break;
			}

			return u;
		};

		function<void(int,int,int)> dfs = [&](int u,int par,int dis){
			if (dis>mid) return;
			mxd = max(mxd,dis);

			if (f[u]){
				dp[dis]++;
			}
			wait[dis].push_back(u);

			for (auto& v:adj[u]){
				if (v==par || vis[v]) continue;

				dfs(v,u,dis+1);
			}
		};

		function<void(int)> work = [&](int u){
			vis[u] = true;

			for (int i=0;i<=mxd;i++){
				dp[i] = 0;
				wait[i].clear();	
			}

			mxd = 0;

			if (f[u]){
				dp[0]++;
			}
			wait[0].push_back(u);

			for (auto& v:adj[u]){
				if (vis[v]) continue;

				dfs(v,u,1);
			}

			for (int i=0;i<=mxd;i++){
				pre[i] = (i-1>=0?pre[i-1]:0)+dp[i];
			}

			int j = mxd;

			for (int i=0;i<=mxd;i++){
				while (j>=0 && i+j>mid){
					j--;
				}

				if (j==-1) break;

				for (auto& tt:wait[i]){
					cnt[tt] += pre[j];
				}
			}

			for (auto& v:adj[u]){
				if (vis[v]) continue;

				for (int i=0;i<=mxd;i++){
					dp[i] = 0;
					wait[i].clear();
				}

				mxd = -1;
				dfs(v,u,1);	

				for (int i=0;i<=mxd;i++){
					pre[i] = (i-1>=0?pre[i-1]:0)+dp[i];
				}

				j = mxd;
				for (int i=0;i<=mxd;i++){
					while (j>=0 && i+j>mid){
						j--;
					}

					if (j==-1) break;

					for (auto& tt:wait[i]){
						cnt[tt] -= pre[j];
					}
				}
			}

			for (auto& v:adj[u]){
				if (vis[v]) continue;
				
				work(getsent(v));			
			}
		};

		work(getsent(1));

		for (int i=1;i<=n;i++){
			if (cnt[i]<=k){
				return true;
			}
		}

		return false;
	};

	int l=0,r=n;
	while (l<=r){
		int mid = l+r >> 1;

		if (check(mid)){
			l = mid+1;
		}
		else{
			r = mid-1;
		}
	}

	cout << r+1 << '\n';
}

int main(){
	ios::sync_with_stdio(false);
	cin.tie(0);
	
	int t = 1;
	// cin >> t;
	while (t--) solve();

	return 0;
}
posted @ 2026-05-11 14:25  kzssCCC  阅读(18)  评论(0)    收藏  举报