深圳技术大学第六届程序设计竞赛 N题思路分享(树上启发式合并,权值线段树)

https://cpc.csgrandeur.cn/csgoj/problemset/problem?pid=1524

题意概述

给定一棵有根树,根为 \(1\) ,节点 \(i\) 的力量值为 \(a_i\)

对于树上任意一个子树 \(T(u)\)(以 \(u\) 为根的子树),设其包含的节点集合为 \(S(u)\),定义该子树的 羁绊分歧度 为:

\[D(u) = \sum_{\{x,y\} \subseteq S(u),\, x \ne y} |a_x - a_y| \]

即子树内所有无序对 \((x, y)\) 力量值之差的绝对值之和。

对每个 \(u = 1, 2, \dots, n\) 输出 \(D(u) \bmod 998244353\)

思路

考虑树上启发式合并。

在计算节点 \(u\) 贡献时,需要加上当前所有 \(a_v\)\(a_u\) 的绝对值,分两部分计算。

  1. \(a_v \gt a_u\) ,记符合条件的 \(a_v\) 数量为 \(cnt\),这部分贡献为 \(\sum a_v - a_u \cdot cnt\)

  2. \(a_v \lt a_u\) ,同理,为 \(a_u \cdot cnt - \sum a_v\)

可以用权值线段树维护,先对 \(a\) 离散化,线段树每个节点维护离散化前 \(a_i\) 的和,以及 \(a_i\) 的个数。

但要注意增加和删除贡献的顺序,想象成栈,增加贡献把节点压入栈,删除贡献先删除栈顶的。

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

代码

//author:kzssCCC

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

const int MOD = 998244353;
vector<ll> reff;

class node{
public:
	ll sum=0,cnt=0;
};

class segmentTree{
public:
	int n;
	vector<node> seg;

	segmentTree(int _n){
		n = _n;
		seg = vector<node>(4*n+1);
	}

	node merge(node p1,node p2){
		node temp;
		temp.sum = (p1.sum+p2.sum)%MOD;
		temp.cnt = p1.cnt+p2.cnt;
		return temp;
	}

	void build(vector<ll>& a){
		build(1,1,n,a);
	}	

	void build(int rt,int l,int r,vector<ll>& a){
		if (l==r){
			//

			return;
		}	

		int mid = l+r >> 1;
		build(rt<<1,l,mid,a);
		build(rt<<1|1,mid+1,r,a);

		seg[rt] = merge(seg[rt<<1],seg[rt<<1|1]);			
	}

	void push_down(int rt,int l,int r){

	}

	void update(int pos,ll val){
		update(1,1,n,pos,val);
	}

	void update(int rt,int l,int r,int pos,ll val){
		if (l==r){
			seg[rt].sum = (seg[rt].sum+val*reff[pos]+MOD)%MOD;
			seg[rt].cnt += val;

			return;
		}		

		int mid = l+r >> 1;
		push_down(rt,l,r);

		if (pos<=mid){
			update(rt<<1,l,mid,pos,val);
		}
		else{
			update(rt<<1|1,mid+1,r,pos,val);
		}

		seg[rt] = merge(seg[rt<<1],seg[rt<<1|1]);
	}

	void update_range(int x,int y,ll val){
		update_range(1,1,n,x,y,val);
	}

	void update_range(int rt,int l,int r,int x,int y,ll val){
		if (r<x || l>y){
			return;
		}

		if (x<=l && y>=r){
			//


			return;
		}

		int mid = l+r >> 1;
		push_down(rt,l,r);

		update_range(rt<<1,l,mid,x,y,val);
		update_range(rt<<1|1,mid+1,r,x,y,val);

		seg[rt] = merge(seg[rt<<1],seg[rt<<1|1]);
	}


	node query(int pos){
		if (pos<1 || pos>n) return {};
		return query(1,1,n,pos);
	}

	node query(int rt,int l,int r,int pos){
		if (l==r){
			return seg[rt];
		}		

		int mid = l+r >> 1;
		push_down(rt,l,r);

		if (pos<=mid){
			return query(rt<<1,l,mid,pos);
		}
		else{
			return query(rt<<1|1,mid+1,r,pos);
		}
	}


	node query_range(int l,int r){
		if (l<1 || l>n || r<1 || r>n || l>r) return {};
		return query_range(1,1,n,l,r);
	}

	node query_range(int rt,int l,int r,int x,int y){
		if (r<x || l>y){
			return {};
		}

		if (x<=l && y>=r){
			return seg[rt];
		}

		int mid = l+r >> 1;
		push_down(rt,l,r);

		return merge(query_range(rt<<1,l,mid,x,y),query_range(rt<<1|1,mid+1,r,x,y));
	}
};

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

	vector<ll> a(n+1);
	for (int i=1;i<=n;i++){
		cin >> a[i];
	}

	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);
	}

	auto uni = a;
	sort(uni.begin()+1,uni.end());
	uni.erase(unique(uni.begin()+1,uni.end()),uni.end());

	int m = uni.size()-1;
	reff = vector<ll>(m+1);

	for (int i=1;i<=n;i++){
		int pos = lower_bound(uni.begin()+1,uni.end(),a[i])-uni.begin();

		reff[pos] = a[i];
		a[i] = pos;
	}

	segmentTree sg(m);

	vector<int> sz(n+1),son(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<ll> res(n+1);
	ll cur = 0;
	
	function<void(int,int,int)> change = [&](int u,int par,int op){
		if (op==1){
			auto t1 = sg.query_range(a[u]+1,m);			
			auto t2 = sg.query_range(1,a[u]-1);

			ll add1 = (t1.sum-t1.cnt*reff[a[u]]%MOD+MOD)%MOD;
			ll add2 = (t2.cnt*reff[a[u]]%MOD-t2.sum+MOD)%MOD;

			cur = ((cur+add1)%MOD+add2)%MOD;
			sg.update(a[u],1);					
		}

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

				change(v,u,op);
			}
		}
		else{
			int len = adj[u].size();
			for (int j=len-1;j>=0;j--){
				int v = adj[u][j];
				if (v==par) continue;

				change(v,u,op);
			}
		}

		if (op==0){
			auto t1 = sg.query_range(a[u]+1,m);			
			auto t2 = sg.query_range(1,a[u]-1);		

			ll add1 = (t1.sum-t1.cnt*reff[a[u]]%MOD+MOD)%MOD;
			ll add2 = (t2.cnt*reff[a[u]]%MOD-t2.sum+MOD)%MOD;

			cur = (cur-add1+MOD)%MOD;
			cur = (cur-add2+MOD)%MOD;

			sg.update(a[u],-1);			
		}
	};

	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);
		}

		{
			auto t1 = sg.query_range(a[u]+1,m);			
			auto t2 = sg.query_range(1,a[u]-1);

			ll add1 = (t1.sum-t1.cnt*reff[a[u]]%MOD+MOD)%MOD;
			ll add2 = (t2.cnt*reff[a[u]]%MOD-t2.sum+MOD)%MOD;

			cur = ((cur+add1)%MOD+add2)%MOD;
			sg.update(a[u],1);
		}

		res[u] = cur;

		if (keep==0){
			{
				auto t1 = sg.query_range(a[u]+1,m);			
				auto t2 = sg.query_range(1,a[u]-1);		

				ll add1 = (t1.sum-t1.cnt*reff[a[u]]%MOD+MOD)%MOD;
				ll add2 = (t2.cnt*reff[a[u]]%MOD-t2.sum+MOD)%MOD;

				cur = (cur-add1+MOD)%MOD;
				cur = (cur-add2+MOD)%MOD;

				sg.update(a[u],-1);	
			}

			int len = adj[u].size();
			for (int j=len-1;j>=0;j--){
				int v = adj[u][j];
				if (v==par || v==son[u]) continue;

				change(v,u,0);
			}

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

	dfs2(1,-1,1);	

	for (int i=1;i<=n;i++){
		cout << res[i] << ' ';
	}
	cout << '\n';
}

int main(){
	int size(64 << 20);  // 64 MB
    __asm__("movq %0, %%rsp\n" :: "r"((char*)malloc(size) + size));

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

	exit(0);
}
posted @ 2026-05-12 16:09  kzssCCC  阅读(28)  评论(0)    收藏  举报