Splay 52pts求调
查看原帖
Splay 52pts求调
482007
TanX_1e18楼主2022/11/11 10:50
#include<bits/stdc++.h>
#define ls(x) tr[x].zi[0]
#define rs(x) tr[x].zi[1]
#define fat(x) tr[x].fa
#define root tr[0].zi[1]
using namespace std;
int n,top;
struct dian
{
    int fa;
    int zi[2];
    int dat;
    int num;
    int size;
}tr[100009];
int where(int x)
{
    return tr[fat(x)].zi[0]==x?0:1;
}
void qf(int x,int f,int how)
{
    tr[f].zi[how]=x;
    tr[x].fa=f;
}
void pushup(int x)
{
    tr[x].size=tr[ls(x)].size+tr[rs(x)].size+tr[x].num;
}
void xuan(int x)
{
    int fu=fat(x),ye=fat(fu);
    int fus=where(x),yes=where(fu);
    qf(tr[x].zi[fus^1],fu,fus);
    qf(fu,x,fus^1);
    qf(x,ye,yes);
    pushup(fu);
    pushup(x);
}
void splay(int x,int to)
{
    to=fat(to);
    while(fat(x)!=to)
    {
        int y=fat(x);
        if(fat(y)==to)
        xuan(x);
        else if(where(x)==where(y))
        {
            xuan(y);
            xuan(x);
        }
        else
        {
            xuan(x);
            xuan(x);
        }
    }
}
int newdian(int v,int f)
{
    top++;
    tr[top].fa=f;
    tr[top].num=1;
    tr[top].size=1;
    tr[top].dat=v;
    return top;
}
void insert(int x)
{
    int now=root;
    if(root==0)
    {
        root=newdian(x,0);
        return;
    }
    while(1)
    {
        tr[now].size++;
        if(tr[now].dat==x)
        {
            tr[now].num++;
            splay(now,root);
            return;
        }
        int zuoyou=x<tr[now].dat?0:1;
        if(tr[now].zi[zuoyou]==0)
        {
            int p=newdian(x,now);
            tr[now].zi[zuoyou]=p;
            splay(p,root);
            return;
        }
        now=tr[now].zi[zuoyou];
    }
}
int find(int x)
{
    int now=root;
    while(1)
    {
        if(now==0)
        {
            return 0;
        }
        if(tr[now].dat==x)
        {
            splay(now,root);
            return now;
        }
        int zuoyou= x<=tr[now].dat ? 0 : 1;
        now=tr[now].zi[zuoyou];
    }
    return 0;
}
void delet(int x)
{
    int pos=find(x);
    if(pos==0)
    return;
    if(tr[pos].num>1)
    {
        tr[pos].num--;
        tr[pos].size--;
        return;
    }
    if(tr[pos].zi[0]==0&&tr[pos].zi[1]==0)
    {
        root=0;
        return;
    }
    else
    {
        if(tr[pos].zi[0]==0)
        {
            root=tr[pos].zi[1];
            tr[root].fa=0;
            return;
        }
        else
        {
            int l=ls(pos);
            while(rs(l)!=0)
            {
                l=rs(l);
            }
            splay(l,ls(pos));
            qf(rs(pos),l,1);
            qf(l,0,1);
            pushup(l);
            return;
        }
    }
}
int pai(int x)
{
    return tr[ls(find(x))].size+1;
}
int wei(int x)
{
    int now=root;
    while(1)
    {
        int tm=tr[now].size-tr[rs(now)].size;
        if(tr[ls(now)].size<x&&x<=tm)
        {
            splay(now,root);
            return tr[now].dat;
        }
        if(x<tm)
        {
            now=ls(now);
        }
        else
        {
            now=rs(now);
            x-=tm;
        }
    }
}
int lower(int x)
{
    int now=root,ans=-114514514;
    while(now!=0)
    {
        if(tr[now].dat<x)
        ans=max(ans,tr[now].dat);
        int nt= x<=tr[now].dat?0:1;
        now=tr[now].zi[nt];
    }
    return ans;
}
int upper(int x)
{
    int now=root,ans=114514514;
    while(now!=0)
    {
        if(tr[now].dat>x)
        ans=min(ans,tr[now].dat);
        int nt= x<=tr[now].dat?0:1;
        now=tr[now].zi[nt];
    }
    return ans;
}
int main()
{
    cin>>n;
    for(int i=1;i<=n;i++)
    {
        int opt,x;
        cin>>opt>>x;
        switch(opt)
        {
            case 1:insert(x);break;
            case 2:delet(x);break;
            case 3:cout<<pai(x)<<endl;break;
            case 4:cout<<wei(x)<<endl;break;
            case 5:cout<<lower(x)<<endl;break;
            case 6:cout<<upper(x)<<endl;break;
        }
    }
    return 0;
}

2022/11/11 10:50
加载中...