P3605 [USACO17JAN] Promotion Counting P 解题报告

P3605 [USACO17JAN] Promotion Counting P 解题报告 (线段树合并+离散化)

\(\quad\)中考结束,又回来训练了,于是,在2026/7/6这天,教练讲了 线段树合并 这么个算法,所以为了方便记忆,就打算写一篇模板题的题解

1. 原题 P3605 [USACO17JAN] Promotion Counting P

2. 概述题意

给定 \(n\) 个节点的一棵树,每个节点有一个权值,问以每个节点为根的子树中,有多少个节点的权值大于当前节点

3. 算法介绍

前置知识:权值线段树,动态开点线段树,线段树合并,离散化
这里只详细介绍线段树合并

线段树合并

“合并”——核心词义指将两个或多个独立实体结合为一的行为。——百度百科

线段树合并,也就顾名思义,将两棵已知的线段树合并为一棵线段树的算法,相应的,有线段树合并就有线段树分裂。因为线段树可以将两棵线段树变成一棵,所以线段树合并可以用来解决图上或者树上信息处理类问题。
该算法常常用于动态开点的权值线段树

算法过程

我们设一棵线段树是 \(X\) 根节点为 ,一棵是 \(Y\) ,首先我们各自让其从根节点向下递归,我们知道动态开点的线段树是形态各异的,所以我们可能会遇到如下三种情况:

  1. 线段树 \(X\) 的节点 线段树 \(Y\) 没有:此时我们直接返回线段树 \(X\) 的节点即可

  2. 线段树 \(Y\) 的节点 线段树 \(X\) 没有:此时我们直接返回线段树 \(X\) 的节点即可

  3. 线段树 \(X\) 、 \(Y\) 都有该节点:此时我们将信息合并,并继续向下递归,再用子节点更新当前节点即可

一般情况下线段树合并不会去开一棵新的线段树,所以我们一般将信息全部合并至 \(X\) 线段树,所以普通线段树合并的板子长这样:

int merge(int x, int y, int l, int r){
	if(!x) return y;
	if(!y) return x;
	
	if(l == r){
		// 将叶子结点的信息合并
		return ; 
	}
	int mid = l + r >> 1;
	tr[x].ls = merge(tr[x].ls, tr[y].ls, l, mid);
	tr[x].rs = merge(tr[x].rs, tr[y].rs, mid + 1, r);
	pushup(x);
	return x;
} 

时间复杂度

虽然这个代码看起来像是暴力,时间复杂度会很大,但实际上操作的都是动态开点的权值线段树,所以时间复杂度可以看做近似 \(O(nlogn)\)

4. 题解

首先,我们如何快速( \(O(logn)\) )知道一堆数中有多少个数比 \(x\) 大呢?没错!权值线段树!如果我们能够建出权值线段树,查询大于 \(x\) 的数的个数就是 query(x + 1, MAX)

经过对数据范围的观察,我们发现朴素的权值线段树会炸空间,所以我们考虑离散化或动态开点(本篇全用了)

那么我们现在考虑如何在一棵树中求出 以 \(u\) 子树中,他的子节点权值有多少个比 \(u\) 大的呢? 我们可以首先给每个节点建一棵线段树,然后我们通过对这颗树进行dfs,再通过dfs递归的方式,不断将子节点的线段树与当前节点的线段树进行合并,合并完以后直接对节点的线段树 query 询问求出答案即可,这个过程就用到了线段树合并

算法过程

  1. 首先将点权,也就是牛的能力值存下来离散化,然后对每个节点单独建立一棵权值线段树

  2. 将这棵树读入,跑dfs,在dfs递归调用结束后,将其子节点的线段树与当前节点的线段树进行合并

  3. 遍历完子节点并对其线段树合并以后,将解求出,并存下来,最后输出即可

代码

废话说了这么多,直接粘上代码~

#include<bits/stdc++.h>

using namespace std;

const int N = 1e6 + 10;
typedef long long LL;
typedef pair<int, int> PII;

int n, nn, v[N], p[N];
vector<int> e[N];
int ans[N];

int root[N], tot = 0;
struct Node{
	int ls, rs, sum;
}tr[N << 2];

void pushup(int u){
	tr[u].sum = tr[tr[u].ls].sum + tr[tr[u].rs].sum;
}

void modify(int &u, int l, int r, int pos, int x){
	if(!u) u = ++ tot;
	if(l == r){
		tr[u].sum += x;
		return ;
	}
	
	int mid = l + r >> 1;
	if(pos <= mid) modify(tr[u].ls, l, mid, pos, x);
	else modify(tr[u].rs, mid + 1, r, pos, x);
	pushup(u);
}

int query(int u, int l, int r, int x){ // 相当于普通query的query(u, l, r, L, R = nn);  
	if(l == r) return 0;
	
	int mid = l + r >> 1;
	if(x <= mid) return tr[tr[u].rs].sum + query(tr[u].ls, l, mid, x);
	return query(tr[u].rs, mid + 1, r, x);
}

int merge(int x, int y, int l, int r){
    if(!x) return y;
    if(!y) return x;
    
	if(l == r){
		tr[x].sum += tr[y].sum;
	}
	
	int mid = l + r >> 1;
    tr[x].ls = merge(tr[x].ls, tr[y].ls, l, mid);
    tr[x].rs = merge(tr[x].rs, tr[y].rs, mid + 1, r);
    pushup(x);
    return x;
}

void dfs(int u){
	for(int j : e[u]){
		dfs(j);
		root[u] = merge(root[u], root[j], 1, nn);
	}
	ans[u] = query(root[u], 1, nn, p[u]);
}

int main(){
	
	cin>>n;
	
	for(int i = 1 ; i <= n ; i ++){
		cin>>p[i];
		v[i] = p[i];
	}
	
	sort(v + 1, v + n + 1);
	nn = unique(v + 1, v + n + 1) - v - 1;
	for(int i = 1 ; i <= n ; i ++){
		p[i] = lower_bound(v + 1, v + nn + 1, p[i]) - v;
		modify(root[i], 1, nn, p[i], 1);
	} 
	
	for(int i = 2 ; i <= n ; i ++){
		int x;
		cin>>x;
		e[x].push_back(i);
	}
	
	dfs(1);
	for(int i = 1 ; i <= n ; i ++) cout<<ans[i]<<'\n';
	
	return 0;
} 
posted @ 2026-07-10 20:42  神烦doge  阅读(23)  评论(0)    收藏  举报