#include<bits/stdc++.h>
using namespace std;
#define N 100005
#define INF 1000000000
#define ls(x) tr[x].ls
#define rs(x) tr[x].rs
#define cnt(x) tr[x].cnt
#define val(x) tr[x].val
#define rnd(x) tr[x].rnd
#define siz(x) tr[x].siz
#define debug cout<<"ff"<<endl
struct node{
int ls,rs,cnt,val,rnd,siz;
}tr[N];
int rt,tot;
void pushup(int rt){
siz(rt)=siz(ls(rt))+siz(rs(rt))+cnt(rt);
}
int newnode(int val){
tot++;
ls(tot)=rs(tot)=0;
cnt(tot)=1;val(tot)=val;
rnd(tot)=rand();siz(tot)=1;
return tot;
}
void build(){
rt=newnode(-INF);
rs(rt)=newnode(INF);
pushup(rt);
}
void zag(int &rt){//d=0
int tmp=rs(rt);
rs(rt)=ls(tmp);
ls(tmp)=rt;
rt=tmp;
pushup(rt);pushup(ls(rt));
}
void zig(int &rt){//d=1
int tmp=ls(rt);
ls(rt)=rs(tmp);
rs(tmp)=rt;
rt=tmp;
pushup(rt);pushup(rs(rt));
}//右旋
void insert(int &rt,int v){
if(!rt){
rt=newnode(v);
return;
}
if(val(rt)==v){
cnt(rt)++;
pushup(rt);
return;
}
if(val(rt)<=v){
insert(rs(rt),v);
if(rnd(rs(rt))<rnd(rt)) zag(rt);
pushup(rt);
}
else{
insert(ls(rt),v);
if(rnd(ls(rt))<rnd(rt)) zig(rt);
pushup(rt);
}
}
void remove(int &rt,int v){
if(!rt) return;
if(val(rt)==v){
if(cnt(rt)>1){
cnt(rt)--;
pushup(rt);
return;
}
if(ls(rt)||rs(rt)){
if(rs(rt)==0||rnd(ls(rt))<rnd(rs(rt))){
zig(rt);
remove(rs(rt),v);
}
else{
zag(rt);
remove(ls(rt),v);
}
pushup(rt);
}
else{
rt=0;
return;
}
return;
}
if(val(rt)<=v) remove(rs(rt),v);
else remove(ls(rt),v);
pushup(rt);
}
int get_val(int x){
int tmp=rt;
while(tmp){
if(x>siz(ls(tmp))&&x<=siz(ls(tmp))+cnt(tmp)) return val(tmp);
if(x<=siz(ls(tmp))) tmp=ls(tmp);
else x-=siz(ls(tmp))+cnt(tmp),tmp=rs(tmp);
}
return val(tmp);
}
int get_rank(int x){
int tmp=rt,res=0;
while(tmp){
if(x==val(tmp)) return res+siz(ls(tmp))+1;
if(x<val(tmp)) tmp=ls(tmp);
else res+=siz(ls(tmp))+cnt(tmp),tmp=rs(tmp);
}
return res;
}
int get_pre(int x){
int tmp=rt,res=0;
while(tmp){
if(val(tmp)>=x) tmp=ls(tmp);
else res=val(tmp),tmp=rs(tmp);
}
return res;
}
int get_nxt(int x){
int tmp=rt,res=0;
while(tmp){
if(val(tmp)<=x) tmp=rs(tmp);
else res=val(tmp),tmp=ls(tmp);
}
return res;
}
signed main(){
srand(time(NULL));
ios::sync_with_stdio(false);
cin.tie(0);cout.tie(0);
build();
int n;
cin>>n;
while(n--){
int opt,x;
cin>>opt>>x;
if(opt==1) insert(rt,x);
else if(opt==2) remove(rt,x);
else if(opt==3) cout<<get_rank(x)-1<<endl;
else if(opt==4) cout<<get_val(x+1)<<endl;
else if(opt==5) cout<<get_pre(x)<<endl;
else if(opt==6) cout<<get_nxt(x)<<endl;
}
return 0;
}
蒟蒻复习了一下Treap,然后上面这份代码喜提WA,调了好久没找到问题,然后看了下第一篇题解,也没啥问题。唯一的区别好像就是题解是大根堆,我是小根堆,然后我改了一下竟然过了?!
神奇,我改的几行:
if(rnd(rs(rt))<rnd(rt)) zag(rt);
if(rnd(ls(rt))<rnd(rt)) zig(rt);
if(rs(rt)==0||rnd(ls(rt))<rnd(rs(rt)))
分别在57,62,75行,均把小于换成了大于,就过了。
然后我再把题解的两个符号改一下(改成小根堆)WA。
这个东西为什么会WA,请大佬指教一下。