深圳技术大学第六届程序设计竞赛 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\) 的绝对值,分两部分计算。
-
\(a_v \gt a_u\) ,记符合条件的 \(a_v\) 数量为 \(cnt\),这部分贡献为 \(\sum a_v - a_u \cdot cnt\) 。
-
\(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);
}

浙公网安备 33010602011771号