萌新初学treap,模板36pts,码风清晰求调
查看原帖
萌新初学treap,模板36pts,码风清晰求调
421265
eastcloud楼主2022/7/15 19:44
#include<iostream>
#include<cstdio>
#include<algorithm>
#include<cstring>
#include<cmath>
#include<cstdlib>
using namespace std;
struct Node{
	int l,r;
	int val,dat;
	int cnt,size;
}tr[100010];
int tot,root;
int New(int val){
	tr[++tot].val=val;
	tr[tot].dat=rand();
	tr[tot].cnt=1;
	tr[tot].size=1;
	return tot;
}
void push_up(int x){
	tr[x].size=tr[tr[x].l].size+tr[tr[x].r].size+tr[x].cnt;
}
void build(){
	New(-(1<<30));
	New(1<<30);
	tr[1].r=2;
	root=1;
	push_up(root);
}
void zig(int &p){
	int q=tr[p].l;
	tr[p].l=tr[q].r;
	tr[q].r=p;
	p=q;
	push_up(tr[p].r);
	push_up(p);
}
void zag(int &p){
	int q=tr[p].r;
	tr[p].r=tr[q].l;
	tr[q].l=p;
	p=q;
	push_up(tr[p].l);
	push_up(p);
}
void insert(int &p,int val){
	if(!p)p=New(val);
	else if(val==tr[p].val){
		tr[p].cnt++;
		push_up(p);
	}
	else if(val<tr[p].val){
		insert(tr[p].l,val);
		if(tr[p].dat<tr[tr[p].l].dat) zig(p);
	}
	else if(val>tr[p].val){
		insert(tr[p].r,val);
		if(tr[p].dat<tr[tr[p].r].dat) zag(p);
	}
}
void remove(int &p,int val){
	if(!p) return;
	if(tr[p].val==val){
		if(tr[p].cnt>1){
			tr[p].cnt--;
			return;
		}
		if(tr[p].l || tr[p].r){
			if(!tr[p].r || tr[tr[p].l].dat>tr[tr[p].r].dat){
				zig(p);
				remove(tr[p].r,val);
			}
			else{
				zag(p);
				remove(tr[p].l,val);
			}
		}
		else p=0;
		return;
	}
	val<tr[p].val?remove(tr[p].l,val):remove(tr[p].r,val);
	push_up(p);
}
int get_rank(int p,int val){
	if(!p) return 0; 
	else if(val==tr[p].val) return tr[tr[p].l].size+1;
	else if(val<tr[p].val) return get_rank(tr[p].l,val);
	return get_rank(tr[p].r,val)+tr[p].cnt+tr[tr[p].l].size;
}
int get_val(int p,int rank){
	if(p==0) return 1<<30;
	else if(tr[tr[p].l].size>=rank) return get_val(tr[p].l,rank);
	else if(tr[tr[p].l].size+tr[p].cnt>=rank) return tr[p].val;
	return get_val(tr[p].r,rank-tr[tr[p].l].size-tr[p].cnt);
}
int get_pre(int p,int val){
	int ans=1;
	while(p){
		if(val==tr[p].val){
			if(tr[p].l){
				p=tr[p].l;
				while(tr[p].r) p=tr[p].r;
				ans=p;
			}
			break;
		}
		if(tr[p].val<val && tr[p].val>tr[ans].val) ans=p;
		p = val<tr[p].val?tr[p].l:tr[p].r;
	}
	return tr[ans].val;
}
int get_next(int p,int val){
	int ans=2;
	while(p){
		if(val==tr[p].val){
			if(tr[p].r){
				p=tr[p].r;
				while(tr[p].l) p=tr[p].l;
				ans=p;
			}
			break;
		}
		if(tr[p].val>val && tr[p].val<tr[ans].val) ans=p;
		p=val<tr[p].val?tr[p].l:tr[p].r;
	}
	return tr[ans].val;
}
int main(){
	int n,opt,x;
	cin>>n;
	build();
	for(int i=1;i<=n;i++){
		cin>>opt>>x;
		if(opt==1) insert(root,x);
		else if(opt==2) remove(root,x);
		else if(opt==3) cout<<get_rank(root,x)-1<<endl;
		else if(opt==4) cout<<get_val(root,x+1)<<endl;
		else if(opt==5) cout<<get_pre(root,x)<<endl;
		else cout<<get_next(root,x)<<endl;
	}
}
2022/7/15 19:44
加载中...