平衡树模板求调,好像有RE....对着老师打的模板对了好久也不知道哪儿错了
查看原帖
平衡树模板求调,好像有RE....对着老师打的模板对了好久也不知道哪儿错了
526895
WYZ20030051楼主2023/7/2 17:47
#include<iostream>
#include<cstdio>
#include<cmath>
#include<string>
#include<cstring>
#include<algorithm>
#include<cassert>
#include<stack>
#include<queue>
#include<vector>
#include<map>
#include<cstdlib>
using namespace std;
#define ll long long
#define ull unsigned long long
#define INF 0x3f3f3f3f3f3f3f
int read()
{
	int now=0,nev=1; 
	char c=getchar();
	while(c<'0' || c>'9') 
	{ 
		if(c=='-') 
			nev=-1; 
		c=getchar();
	}
	while(c>='0' && c<='9') 
	{ 
		now=(now<<1)+(now<<3)+(c&15); 
		c=getchar(); 
	}
	return now*nev;
}
const int MAXN=1e5+10;
int n;
int root,tt;
struct node
{
	int l,r;//左右儿子 
	int p;//随机数p,用于保持树的平衡,由于插入点和删除点后平衡树可能会失衡,所以需要一个随机数维护失衡 
	int w;//当前节点权值 
	int size;//子树大小
	int c;//用作计数器 
}tr[MAXN]; 
void rotate_left(int &k)//左旋转,即逆时针旋转,会把一个节点的右子树转到左子树上 
{
	tr[k].r=tr[tr[k].r].l;//将当前节点的右子树转到当前节点的右子树的左子树上,即把右下角的给接到了左下角 
	tr[tr[k].r].l=k;//将当前节点的右子树的左子树转成节点 
	tr[tr[k].r].size=tr[k].size;//同时更新左右子树的大小 
	tr[k].size=tr[tr[k].l].size+tr[tr[k].r].size+tr[k].c;//当前节点的子树大小为左子树大小加右子树大小加上当前点的大小 
	k=tr[k].r;//更新节点,与第一第二行的旋转操作相呼应 
}
void rotate_right(int &k)//右旋转,即顺时针旋转,会把一个节点的左子树转到右子树上,代码是把左旋转中的l与r互换一下 
{
	tr[k].l=tr[tr[k].l].r;
	tr[tr[k].l].r=k;
	tr[tr[k].l].size=tr[k].size;
	tr[k].size=tr[tr[k].l].size+tr[tr[k].r].size+tr[k].c;
	k=tr[k].l;
}
void insert_x(int &k,int x)//插入,k为节点编号,x为插入点的大小 
{
	if(!k)//加点操作,只要当前节点不为0就可以加 
	{
		k=++tt;
		tr[k].w=x;//有点类似链式前向星的加点 
		tr[k].p=rand();//roud函数可用于产生随机数 
		tr[k].c=1;//以下四行为初始化,注意赋值的大小 
		tr[k].size=1;
		tr[k].l=0;
		tr[k].r=0;
		return ;
	}
	tr[k].size++;//插入点后子树大小也会相应的大一点 
	if(tr[k].w==x)
		tr[k].c++;//计数器此时自加1 
	if(x<tr[k].w)//若当前节点的值比目标值大,则到该节点的右子树插入点 
	{
		insert_x(tr[k].l,x);//到左子树插入点 
		if(tr[tr[k].l].p<tr[k].p)//如果树不平衡了,就进行旋转 
			rotate_right(k);//由于插入节点后左子树不平衡,所以就右旋转使其平衡 
	}
	else//若当前节点的值比目标值小,则到该节点的左子树插入点 
	{
		insert_x(tr[k].r,x);
		if(tr[tr[k].r].p<tr[k].p)
			rotate_left(k);//同上,由于插入节点后右子树不平衡,所以就左旋转使其平衡 
	}
}
void delete_x(int &k,int x)//删除 
{
	if(tr[k].size==x)
	{
		if(tr[k].size>1) 
		{
			tr[k].size--;
			tr[k].c--;
		}
		else if(!tr[k].l || !tr[k].r)//若只有一个子树 
			k=tr[k].l+tr[k].r;//用子树将当前点更新一下,也就是删除了当前点后,直接让子树与当前点的父节点连接 
		else if(tr[tr[k].l].p<tr[tr[k].r].p)
		{
			rotate_right(k);//由于插入与删除是相对的操作,删除点后会面临着高度过小导致的失衡,所以要有旋转 
			delete_x(k,x);
		}
		else
		{
			rotate_left(k);
			delete_x(k,x);
		}
		return ;
	}
	tr[k].size--;
	if(tr[k].w<x)//若当前点小于目标值,则到右子树删除 
		delete_x(tr[k].r,x);
	else
		delete_x(tr[k].l,x);
}
int get_rank(int x)//查询排名 
{
	int k=root;//当前节点编号 
	int rank=0;//这里的rank主要记录该点在右子树中的排名,初始为0 
	while(k)//一直循环到当前节点不存在 
	{
		if(x==tr[k].w)
			return rank+tr[tr[k].l].size+1;//说明已经查找到该节点,此时直接返回当前节点的排名
			//当前节点的排名为:左子树的大小(即左子树的节点个数)+当前节点(也就是最后的+1) 
		if(x<tr[k].w)
			k=tr[k].l;//如果要查询的节点比当前节点的值小,说明要查询的点肯定在当前点的左子树上 
		if(x>tr[k].w)
		{
			rank+=tr[tr[k].l].size+tr[k].c;
			//若要该节点的值大于当前节点,说明该点肯定在当前节点的右子树上 
			k=tr[k].r;//同上
		}
	}
	return rank;
}
int get_kthmath(int x)//查询排名为x的数(查询第k大)
{
	int k=root;
	while(k)
	{
		if(x>tr[k].l && x<=tr[tr[k].l].size+tr[k].c)
		//tr[tr[k].l].size+tr[k].c表示右子树除右子树外节点的个数,也就是左子树和当前节点相加的排名数 
		//如果该排名比当前节点大,说明在当前节点或右子树上;
		//如果该排名小于tr[tr[k].l].size+tr[k].c,则说明在左子树或当前节点上 
		//所以若同时满足两条件,则当前点即为所求 
			return tr[k].w;
		if(x<=tr[tr[k].l].size)
			k=tr[k].l;//递归到左子树 
		else
		{
			x-=tr[tr[k].l].size+tr[k].c;
			//如果当前节点x>tr[tr[k].l].size+tr[k].c,则应该去右子树找第k-tr[tr[k].l].size-tr[k].c大的数 
			k=tr[k].r;//递归到右子树 
		}
	}
}
int get_pre(int x)//查询前驱
{
	int k=root;//当前节点编号 
	int pre=-INF;//前驱一定比当前节点小,所以将前驱初始值设为无穷小,方便更新
	while(k)//一直查询到当前节点不存在时再结束 
	{
		if(x>tr[k].w)//若要查询的点比当前的点要大,则说明它的前驱可能只可能是当前点的右儿子 
		{
			pre=tr[k].w;//更新前驱 
			k=tr[k].r;//递归到右儿子 
		}
		else
			k=tr[k].l;//递归到左儿子 
	}
	return pre;
}
int get_nxt(int x)//查询后继
{ 
	int k=root;
	int nxt=INF;
	while(k)
	{
		if(x<tr[k].w)
		{
			nxt=tr[k].w;
			k=tr[k].l;
		}
		else
			k=tr[k].r;
	}
	return nxt;
}
int main()
{
	memset(tr,0,sizeof(tr));
	n=read();
	root=tt=0;
	while(n--)
	{
		int op,x;
		op=read(),x=read();
		if(op==1)
			insert_x(root,x);
		else if(op==2)
			delete_x(root,x);
		else if(op==3)
			printf("%d\n",get_rank(x));
		else if(op==4)
			printf("%d\n",get_kthmath(x));
		else if(op==5)
			printf("%d\n",get_pre(x));
		else if(op==6)
			printf("%d\n",get_nxt(x));
	}
	return 0;
}
2023/7/2 17:47
加载中...