本人第一次写 AVL,运行模板题样例是发现意外的加入失败,不知道是什么原因。
#include<cstdio>
namespace AVL{
using namespace std;
struct node{
int h,sz,val,oz;
node *lc,*rc;
node(int _val){
h=1,sz=1;
lc=NULL,rc=NULL;
val=_val;
oz=1;
}
};
node* init_a(node*cur){
int u=1;
if(cur->lc!=NULL)u=cur->lc->h+1;
int v=1;
if(cur->rc!=NULL)v=cur->rc->h+1;
cur->h=u<v?v:u;
u=0;
if(cur->lc!=NULL)u=cur->lc->sz;
v=0;
if(cur->rc!=NULL)v=cur->rc->sz;
cur->sz=u+v+cur->oz;
return cur;
}
node* l_rotate(node* cur){
node * tmp=cur->rc;
cur->rc=tmp->lc;
tmp->lc=cur;
cur=init_a(cur);
tmp=init_a(tmp);
cur=tmp;
return cur;
}
node* r_rotate(node* cur){
node * tmp=cur->lc;
cur->lc=tmp->rc;
tmp->rc=cur;
cur=init_a(cur);
tmp=init_a(tmp);
cur=tmp;
return cur;
}
node* balance(node* cur){
cur=init_a(cur);
int u=0;
if(cur->lc!=NULL)u=cur->lc->h;
int v=0;
if(cur->rc!=NULL)v=cur->rc->h;
if(u-v==2){
//printf("r\n");
node* tmp=init_a(cur->lc);
int u=0;
if(tmp->lc!=NULL)u=tmp->lc->h;
int v=0;
if(tmp->rc!=NULL)v=tmp->rc->h;
if(u<v){
tmp=l_rotate(tmp);
//printf("lp\n");
}
return r_rotate(cur);
}
if(v-u==2){
//printf("l\n");
node* tmp=init_a(cur->rc);
int u=0;
if(tmp->lc!=NULL)u=tmp->lc->h;
int v=0;
if(tmp->rc!=NULL)v=tmp->rc->h;
if(u>v){
tmp=r_rotate(tmp);
//printf("rp\n");
}
return l_rotate(cur);
}
else return init_a(cur);
}
node* add(node* cur,int val){
if(cur==NULL){
return cur=new node(val);
}
else if(cur->val==val){
cur->oz++;
return init_a(cur);
}
else if(cur->val>val){
//printf("%d -\n",cur->val);
cur->lc=add(cur->lc,val);
cur=init_a(cur);
return balance(cur);
}
else if(cur->val<val){
//printf("%d +\n",cur->val);
cur->rc=add(cur->rc,val);
cur=init_a(cur);
return balance(cur);
}
}
node* fd(node *cur){
if(cur->lc==NULL)return cur;
else {
return fd(cur->lc);
}
}
node* fd2(node *cur){
if(cur->rc==NULL)return cur;
else {
return fd2(cur->rc);
}
}
node* dm(node* cur){
if(cur->lc==NULL)return cur->rc;
else {
cur->lc=dm(cur->lc);
return balance(cur);
}
}
node* del(node *cur,int val){
if(cur==NULL)return NULL;
if(cur->val==val){
cur->oz--;
if(cur->oz==0){
if(cur->lc==NULL){
node* tmp=cur->rc;
delete cur;
return tmp;
}
if(cur->rc==NULL){
node* tmp=cur->lc;
delete cur;
return tmp;
}
node *opt=fd(cur->rc);
dm(cur->rc);
opt->lc=cur->lc;
opt->rc=cur->rc;
delete cur;
return balance(opt);
}
return init_a(cur);
}
if(cur->val<val){
cur->rc=del(cur->rc,val);
return balance(cur);
}
if(cur->val>val){
cur->lc=del(cur->lc,val);
return balance(cur);
}
}
int rank(node* cur,int val){
if(cur==NULL)return 1;
if(cur->val==val){
int r=1;
if(cur->lc!=NULL)r+=cur->lc->sz;
return r;
}
if(cur->val<val){
int r=rank(cur->rc,val);
if(cur->lc!=NULL)r+=cur->lc->sz;
return r;
}
return rank(cur->lc,val);
}
int find_rk(node* cur,int val){
if(cur->sz<val)return -2147483647;
if(cur->lc!=NULL&&val<=cur->lc->sz)return find_rk(cur->lc,val);
else {
int r=cur->oz;
if(cur->lc!=NULL)r+=cur->lc->sz;
if(val<=r)return cur->val;
else return find_rk(cur->rc,val-r);
}
}
int lower_bound(node* cur,int val){
if(cur==NULL)return -2147483647;
if(cur->val==val){
if(cur->lc!=NULL)return fd2(cur->lc)->val;
else return -2147483647;
}
if(cur->val>val)return lower_bound(cur->lc,val);
else {
int r=lower_bound(cur->rc,val);
if(r<cur->val)r=cur->val;
return r;
}
}
int upper_bound(node* cur,int val){
if(cur==NULL)return 2147483647;
if(cur->val==val){
if(cur->rc!=NULL)return fd(cur->rc)->val;
else return 2147483647;
}
if(cur->val<val)return upper_bound(cur->rc,val);
else {
int r=upper_bound(cur->lc,val);
if(r>cur->val)r=cur->val;
return r;
}
}
void show(node* cur){
if(cur->lc!=NULL)show(cur->lc);
printf("%d %d l\n",cur->val,cur->oz);
if(cur->rc!=NULL)show(cur->rc);
}
};
using namespace AVL;
int n,op,x;
int main(){
scanf("%d",&n);
node* p=NULL;
for(int i=1;i<=n;i++){
scanf("%d%d",&op,&x);
if(op==1){
p=add(p,x);
}
else if(op==2){
p=del(p,x);
}
else if(op==3){
printf("ans=%d\n",rank(p,x));
}
else if(op==4){
printf("ans=%d\n",find_rk(p,x));
}
else if(op==5){
printf("ans=%d\n",lower_bound(p,x));
}
else{
printf("ans=%d\n",upper_bound(p,x));
}
show(p);
}
return 0;
}