splay48分WA求调
查看原帖
splay48分WA求调
311764
__mrfs楼主2023/8/4 15:14
#include<bits/stdc++.h>
using namespace std;
const int N =1e5+10,INF=1e9;
struct Node{
    int s[2],v,p;
    int siz,cnt;
    void init(int _v,int _p){
        v=_v,p=_p;
        siz=1,cnt=1;
    }
}tr[N];
int root,idx;
void pushup(int u){
    tr[u].siz=tr[u].cnt+tr[tr[u].s[0]].siz+tr[tr[u].s[1]].siz;
}
void rotate(int x){
    int y=tr[x].p,z=tr[y].p;
    int k=tr[y].s[1]==x;
    tr[y].s[k]=tr[x].s[k^1],tr[tr[x].s[k^1]].p=y;
    tr[x].s[k^1]=y,tr[y].p=x;
    tr[z].s[tr[z].s[1]==y]=x,tr[x].p=z;
    pushup(y),pushup(x);
}
void splay(int x,int k){
    while(tr[x].p!=k){
        int y=tr[x].p;
        int z=tr[y].p;
        if(z!=k){
            if((tr[z].s[1]==y)^(tr[y].s[1]==x)) rotate(x);
            else rotate(y);
        }
        rotate(x);
    }
    if(!k) root=x;
}
void insert(int x){//插入
    int u=root,p=0;
    while(u){
        if(tr[u].v==x){
            tr[u].cnt++;
            splay(u,0);
            return;
        }
        p=u;
        u=tr[u].s[tr[u].v<x];
    }
    u=++idx;
    if(p) tr[p].s[tr[p].v<x]=u;
    tr[u].init(x,p);
    splay(u,0);
}
int find_pre(int v){//找前驱
    int u=root,res;
    while(u){
        if(tr[u].v<v) res=u,u=tr[u].s[1];
        else u=tr[u].s[0];
    }
    return res;
}
int find_suf(int v){//找后继
    int u=root,res;
    while(u){
        if(tr[u].v>v) res=u,u=tr[u].s[0];
        else u=tr[u].s[1];
    }
    return res;
}
void delet(int v){//删除
    int L=find_pre(v);
    int R=find_suf(v);
    splay(L,0);
    splay(R,L);
    if(tr[tr[R].s[0]].cnt==1) tr[R].s[0]=0;
    else tr[tr[R].s[0]].cnt--;
    pushup(R);
    pushup(L);
}
int find_key(int v){//查询排名
    int u=root,res=0;
    while(u){
        if(tr[u].v>v) u=tr[u].s[0];
        else if(tr[u].v<v) res+=tr[u].cnt+tr[tr[u].s[0]].siz,u=tr[u].s[1];
        else {
            res+=tr[tr[u].s[0]].siz;
            res+=tr[u].cnt;
            return res;
        }
    }
}
int find_value(int k){//查询数值
    int u=root;
    while(u){
        if(tr[tr[u].s[0]].siz>=k) u=tr[u].s[0];
        else if(tr[tr[u].s[0]].siz<k&&tr[tr[u].s[0]].siz+tr[u].cnt>=k) return tr[u].v;
        else k-=tr[tr[u].s[0]].siz+tr[u].cnt,u=tr[u].s[1];
    }
}
int main(){
    insert(-INF),insert(INF);
    int n;
    int tot=2;
    scanf("%d",&n);
    while(n--){
        int op,x;
        scanf("%d%d",&op,&x);
        if(op==1) {
            insert(x);
            tot++;
        }
        else if(op==2){
            delet(x);
            tot--;
        } 
        else if(op==3) printf("%d\n",find_key(x)-1);
        else if(op==4) printf("%d\n",find_value(x+1));
        else if(op==5) printf("%d\n",tr[find_pre(x)].v);
        else printf("%d\n",tr[find_suf(x)].v);
    }
    return 0;
}
2023/8/4 15:14
加载中...