AVL求助
  • 板块学术版
  • 楼主OtterZ
  • 当前回复1
  • 已保存回复1
  • 发布时间2023/7/19 18:03
  • 上次更新2023/11/3 08:49:42
查看原帖
AVL求助
609565
OtterZ楼主2023/7/19 18:03

本人第一次写 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;
}
2023/7/19 18:03
加载中...