求助,splay16pts
查看原帖
求助,splay16pts
536743
arrow_king楼主2023/1/8 21:48

rt,萌新刚学splay,除前两个点以外全T,怀疑是rotate或者是splay操作写挂了或者是树写的不平衡求调

#include<iostream>
#include<cstdio>
using namespace std;
#define ll long long
#define il inline
#define N 200005
il ll read() {
	ll x=0,f=1;char c=getchar();
	while(c<'0'||c>'9') {if(c=='-') {f=-1;} c=getchar();}
	while(c>='0'&&c<='9') {x=(x<<3)+(x<<1)+(c^48);c=getchar();}
	return x*f;
}
int ch[N][2],sum[N],fa[N],cnt[N],a[N],root,tot;
il void push_up(int now) {
	sum[now]=sum[ch[now][0]]+sum[ch[now][1]]+cnt[now];
}
il int getwh(int x) {
	return ch[fa[x]][0]==x?0:1;
}
il void rotate(int x) {
	int y=fa[x],z=fa[y],k=getwh(x);
	fa[x]=z;
	ch[z][getwh(y)]=x;
	fa[ch[x][k^1]]=y;
	ch[y][k]=ch[x][k^1];
	fa[y]=x;
	ch[x][k^1]=y;
	push_up(y);
	push_up(x);
}
il void splay(int x,int tar) {
	while(fa[x]!=tar) {
		int y=fa[x],z=fa[y];
		if(z!=tar) {
			if(getwh(x)==getwh(y)) rotate(y);
			else rotate(x);
		}
		rotate(x);
	}
	if(!tar) root=x;
}
il void findv(int v) {
	int x=root;
	if(!x) return;
	while(ch[x][v>a[x]]&&v!=a[x]) x=ch[x][v>a[x]];
	splay(x,0);
}
il void insert(int v) {
	int x=root,y=0;
	while(x&&a[x]!=v) {
		y=x;
		x=ch[x][v>a[x]];
	}
	if(x) cnt[x]++;
	else {
		x=++tot;
		if(y) ch[y][v>a[y]]=x;
		ch[x][0]=ch[x][1]=0;
		fa[x]=y,a[x]=v,cnt[x]=sum[x]=1;
	}
	splay(x,0);
}
il int nextt(int v,int f) {
	findv(v);
	int x=root;
	if(a[x]>v&&f) return x;
	if(a[x]<v&&!f) return x;
	x=ch[x][f];
	while(ch[x][f^1]) x=ch[x][f^1];
	splay(x,0);
	return x;
}
il void del(int v) {
	int pre=nextt(v,0),nxt=nextt(v,1);
	splay(pre,0);
	splay(nxt,pre);
	int tmp=ch[nxt][0];
	if(cnt[tmp]>1) {
		cnt[tmp]--;
		splay(tmp,0);
	}
	else {
		ch[nxt][0]=0;
		fa[tmp]=0;
	}
}
il int kth(int k) {
	int x=root;
	if(sum[x]<k) return 0;
	while(1) {
		int lc=ch[x][0];
		if(k>sum[lc]+cnt[x]) {
			k-=sum[lc]+cnt[x];
			x=ch[x][1];
		}
		else {
			if(sum[lc]>=k) x=lc;
			else return a[x];
		}
	}
}
il int getrank(int v) {
	int x=root,rank=0;
	while(x) {
		if(a[x]<v) {
			rank+=sum[ch[x][0]]+cnt[x];
			x=ch[x][1];
		}
		else x=ch[x][0];
	}
	return rank+1;
}
int n,opt,x;
int main() {
	n=read();
	for(int i=1;i<=n;i++) {
		opt=read(),x=read();
		switch(opt) {
			case 1:{
				insert(x);
				break;
			}
			case 2:{
				del(x);
				break;
			}
			case 3:{
				printf("%d\n",getrank(x));
				break;
			}
			case 4:{
				printf("%d\n",kth(x));
				break;
			}
			default:{
				printf("%d\n",a[nextt(x,opt-5)]);
				break;
			}
		}
	}
	return 0;
}
2023/1/8 21:48
加载中...