萌新刚学splay,一直都调不对,大佬们救救萌新吧
code(样例能过 8pts):
#include<bits/stdc++.h>
#define re register
#define il inline
using namespace std;
typedef long long LL;
inline int read(){
int x=0,f=1;char ch=getchar();
while(ch<'0'||ch>'9'){if(ch=='-')f=-1;ch=getchar();}
while(ch>='0'&&ch<='9')x=(x<<1)+(x<<3)+(ch&15),ch=getchar();
return x*f;
}
const int N=2e5+100;
int root,tot;
struct splay_tree{
int fu,son[2],val,cnt,siz;
}tr[N];
inline void update(int x){
tr[x].siz=tr[tr[x].son[0]].siz+tr[tr[x].son[1]].siz+tr[x].cnt;
}
inline void rotate(int x){
int y=tr[x].fu,z=tr[y].fu;
int k=(x==tr[y].son[1]);
tr[z].son[y==tr[z].son[1]]=x;
tr[x].fu=z;
tr[y].son[k]=tr[x].son[k^1];
tr[tr[x].son[k^1]].fu=y;
tr[x].son[k^1]=y;
tr[y].fu=x;
update(y),update(x);
}
inline void splay(int x,int goal){
while(tr[x].fu!=goal){
int y=tr[x].fu;int z=tr[y].fu;
if(z!=goal) {
((tr[z].son[0]==y)^(tr[y].son[0]==x))?rotate(x):rotate(y);
}
rotate(x);
}
if(!goal)root=x;
}
inline void find(int x){
int u=root;
if(!u)return;
while(tr[u].son[x>tr[u].val]&&x!=tr[u].val)
u=tr[u].son[x>tr[u].val];
splay(u,0);
}
inline int get_rank(int x){
find(x);
return tr[tr[root].son[0]].siz+1+(tr[root].val<x)?1:0;
}
inline int get_next(int x,int ch){
find(x);
if(tr[root].val<x&&!ch)return root;
if(tr[root].val>x&&ch) return root;
int u=tr[root].son[ch];
while(tr[u].son[ch^1])u=tr[u].son[ch^1];
return u;
}
inline void Insert(int x){
int u=root,ff=0;
while(u&&tr[u].val!=x){
ff=u;
u=tr[u].son[x>tr[u].val];
}
if(u) tr[u].cnt++;
else {
u=++tot;
if(ff) tr[ff].son[x>tr[ff].val]=u;
tr[u].siz=tr[u].cnt=1;
tr[u].val=x;
tr[u].fu=ff;
}
splay(u,0);
}
inline void Delete(int x){
int lst=get_next(x,0),nxt=get_next(x,1);
splay(lst,0),splay(nxt,lst);
int del=tr[nxt].son[0];
if(tr[del].cnt>1){
tr[del].cnt--;
splay(del,0);
}
else tr[nxt].son[0]=0;
}
inline int get_kth(int k){
int u=root;
if(tr[u].siz<k)return 0;
while(1){
if(k>tr[tr[u].son[0]].siz+tr[u].cnt){
k-=tr[tr[u].son[0]].siz+tr[u].cnt;
u=tr[u].son[1];
}
else{
if(tr[tr[u].son[0]].siz>=k){
k-=tr[u].cnt;
u=tr[u].son[0];
}
else return tr[u].val;
}
}
}
int main(){
int T=read();
while(T--){
int opt=read(),x=read();
switch(opt){
case 1:{
Insert(x);
break;
}
case 2:{
Delete(x);
break;
}
case 3:{
printf("%d\n",get_rank(x));
break;
}
case 4:{
printf("%d\n",get_kth(x));
break;
}
case 5:{
printf("%d\n",tr[get_next(x,0)].val);
break;
}
case 6:{
printf("%d\n",tr[get_next(x,1)].val);
break;
}
}
}
return 0;
}