树上启发式合并

https://www.luogu.com.cn/problem/U41492

题意概述

给定一棵 \(n\) 个结点的有根树,根为 \(1\) 。每个结点 \(i\) 有一个颜色 \(c_i\)

\(m\) 次询问,每次询问求以 \(x\) 为根的子树中不同颜色的种类数。

思路

首先做重链剖分,只需要轻重儿子信息。

\(dfs\) 带一个参数 \(keep\) ,表示是否保留当前贡献。对当前节点 \(u\),先 \(dfs\) 所有轻儿子,不保留贡献;
然后 \(dfs\) 重儿子,保留贡献;最后在加上轻儿子和自身的贡献,记录答案。如果 \(keep\)\(0\),把当前贡献删掉。

首先优先遍历所有轻儿子,这时候都是不保留贡献的,相当于独立计算了所有轻儿子的答案。然后遍历重儿子,带贡献,相当于只计算了重儿子的答案,并算上贡献。

至此计算完了所有子树的答案,现在需要计算当前节点的答案。

重儿子的贡献已经算过了,再次遍历所有轻儿子计算贡献即可,但注意不能通过调用 \(dfs\) 函数,因为这样会污染计算好的答案,需要重写一个函数 \(change\) ,用于计算贡献和删除贡献。

最后 \(keep\)\(0\) 的情况,调用 \(change(u)\),删除所有贡献。

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

代码

//author:kzssCCC

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


void solve(){
	int n;
	cin >> n;

	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<int> c(n+1);
	for (int i=1;i<=n;i++){
		cin >> c[i];
	}

	vector<int> son(n+1),sz(n+1);

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

		int pos = -1;
		int mx = 0;

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

			dfs(v,u);
			sz[u] += sz[v];

			if (sz[v]>mx){
				mx = sz[v];
				pos = v;
			}
		}

		son[u] = pos;
	};

	dfs(1,-1);

	vector<int> res(n+1);
	int len = *max_element(c.begin()+1,c.end());
	vector<int> cnt(len+1);
	int cur = 0;

	function<void(int,int,int)> change = [&](int u,int par,int op){
		if (op==0){
			if (--cnt[c[u]]==0){
				cur--;
			}		
		}	
		else{
			if (cnt[c[u]]++==0){
				cur++;
			}
		}

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

			change(v,u,op);
		}
	};

	function<void(int,int,int)> dfs2 = [&](int u,int par,int keep){
		for (auto& v:adj[u]){
			if (v==par || v==son[u]) continue;

			dfs2(v,u,0);				
		}

		if (son[u]!=-1){
			dfs2(son[u],u,1);
		}		

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

			change(v,u,1);				
		}

		if (cnt[c[u]]++==0){
			cur++;
		}	

		res[u] = cur;

		if (keep==0){
			change(u,par,0);
		}		
	};

	dfs2(1,-1,0);

	int m;
	cin >> m;

	while (m--){
		int x;
		cin >> x;

		cout << res[x] << '\n';
	}
}

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

	return 0;
}
posted @ 2026-05-12 14:44  kzssCCC  阅读(8)  评论(0)    收藏  举报