自测了RE样例没问题呀
#include<iostream>
#include<algorithm>
#include<cstdlib>
using namespace std;
struct Node{
Node *zs[2];
int size,val,repcnt,prio;
Node(int val): val(val),repcnt(1),size(1){
zs[0]=zs[1]=nullptr;
prio=rand();
}
void upd_size(){
size=repcnt;
if(zs[0]!=nullptr)size+=zs[0]->size;
if(zs[1]!=nullptr)size+=zs[1]->size;
}
};
void rotate(Node *&cur,int dir){
//以下注释以左旋为例
Node *tmp=cur->zs[dir]; //将tmp设为cur的右孩子
cur->zs[dir]=tmp->zs[!dir]; //cur与tmp断裂,连接到tmp的左孩子
tmp->zs[!dir]=cur; //tmp成为根
tmp->upd_size();
cur->upd_size();//更新大小
cur=tmp;
}
void insert(Node *&cur,int val){
if(cur==nullptr){ //新建
cur=new Node(val);
return;
}else if(val==cur->val){ //发现有重复的
cur->repcnt++;
cur->size++;
}else if(val<cur->val){ //插左边
insert(cur->zs[0],val);
if(cur->zs[0]->prio<cur->prio){
//左边比根小,右旋
rotate(cur,0);
}
cur->upd_size();
}else if(val>cur->val){ //插右边
insert(cur->zs[1],val);
if(cur->zs[1]->prio<cur->prio) {
rotate(cur,1);
}
cur->upd_size();
}
}
void erase(Node *&cur,int val){
if(val<cur->val){
erase(cur->zs[0],val);
cur->upd_size();
}else if(val>cur->val){
erase(cur->zs[1],val);
cur->upd_size();
}else{
if(cur->repcnt>1){ //如果不止一个,直接减
cur->repcnt--;
cur->size--;
return;
}
uint8_t st=0;
st|=(cur->zs[0]!=nullptr);
st|=((cur->zs[1]!=nullptr)<<1);
//二进制,00都无,10有右无左,01有左无右,11都有
Node *tmp=cur;
switch(st){
case 0:
delete cur;
cur=nullptr;
break;
case 1:
cur=tmp->zs[0];
delete tmp;
break;
case 2:
cur=tmp->zs[1];
delete tmp;
break;
case 3:
int dir;
if(cur->zs[0]->prio>cur->zs[1]->prio){
dir=1;
}else{
dir=0;
}
rotate(cur,dir);
erase(cur->zs[!dir],val);
cur->upd_size();
break;
}
}
}
int frank(Node *&cur,int val){
int lesize=cur->zs[0]==nullptr?0 : cur->zs[0]->size;
//子树中比val小的结点的数量
if(val==cur->val){
return lesize+1;
}else if(val<cur->val){
return frank(cur->zs[0],val);
}else{
return lesize+cur->repcnt+frank(cur->zs[1],val);
}
}
int getv(Node *&cur,int rank){
int lesize=cur->zs[0]==nullptr?0 : cur->zs[0]->size;
if(rank<=lesize){
return getv(cur->zs[0],rank);
}else if(rank<=cur->repcnt+lesize){
return cur->val;
}else{
return getv(cur->zs[1],rank-lesize-cur->repcnt);
}
}
int pre_ans,nex_ans;
int pre(Node *&cur,int val){
if(val<=cur->val){
if(cur->zs[0]!=nullptr)return pre(cur->zs[0],val);
}else{
pre_ans=cur->val;
if(cur->zs[1]!=nullptr)pre(cur->zs[1],val);
return pre_ans;
}
}
int next(Node *&cur,int val){
if(val>=cur->val){
if(cur->zs[1]!=nullptr)return next(cur->zs[1],val);
}else{
nex_ans=cur->val;
if(cur->zs[0]!=nullptr)next(cur->zs[0],val);
return nex_ans;
}
}
int main(){
int n;
cin >> n;
Node *root=nullptr;
for(int i=1;i<=n;i++){
int opt,x;
cin >> opt >> x;
switch(opt){
case 1:insert(root,x);break;
case 2:erase(root,x);break;
case 3:cout << frank(root,x) <<"\n";break;
case 4:cout << getv(root,x) <<"\n";break;
case 5:cout << pre(root,x) <<"\n";break;
case 6:cout << next(root,x) <<"\n";break;
}
}
return 0;
}