treap 93分求助 WA on #1
查看原帖
treap 93分求助 WA on #1
444236
Lesiris楼主2023/8/30 20:47

RT

#include<bits/stdc++.h>
using namespace std;
const int N=10e5,INF=1e9;
int sz[N],key[N],cnt[N],sn[N][2],rd[N],k,x,n,opt,tot;
inline void push_up(int k)
{
	sz[k]=sz[sn[k][0]]+sz[sn[k][1]]+cnt[k];
}
inline void rotate(int &k,int d)
{
	int k1=sn[k][d^1];
	sn[k][d^1]=sn[k1][d];
	sn[k1][d]=k;
	push_up(k); 
	push_up(k1);
	k=k1;
}
void insert(int &k,int x)
{
	if(!k)
	{
		k=++tot;
		sz[k]=cnt[k]=1;
		key[k]=x;
		rd[k]=rand();
		return;
	}
	if(key[k]==x)
	{
		sz[k]++;
		cnt[k]++;
		return;
	}
	int d=x>key[k];
	insert(sn[k][d],x);
	if(rd[k]<rd[sn[k][d]]) rotate(k,d^1);
	push_up(k);
}
void del(int &k,int x)
{
	if(!k) return;
	if(x!=key[k]) del(sn[k][x>key[k]],x);
	else
	{
		if(cnt[k]>1)
		{
			cnt[k]--; 
			sz[k]--; 
			return;
		}
		else if(!sz[sn[k][0]]&&!sz[sn[k][1]])
		{
			cnt[k]--;
			sz[k]--;
			k=0;
			return;
		}
		else
		{
			int d;
			if(sz[sn[k][0]]*sz[sn[k][1]])
			{
				d=rd[sn[k][0]]>rd[sn[k][1]];
			}
			else d=sz[sn[k][0]]>=1;
			rotate(k,d);
			del(sn[k][d],x);
		}
	}
	push_up(k);
}
int get_rnk(int k, int x)
{
	if(!k) return 0;
	if(key[k]==x) return sz[sn[k][0]]+1;
	if(key[k]>x) return get_rnk(sn[k][0],x);
	return sz[sn[k][0]]+cnt[k]+get_rnk(sn[k][1],x);
}
int get_val(int k, int x)
{
	if(!k) return 0;
	if(sz[sn[k][0]]>=x) return get_val(sn[k][0],x);
	else if(sz[sn[k][0]]+cnt[k]>=x) return key[k];
	return get_val(sn[k][1],x-sz[sn[k][0]]-cnt[k]);
}
int get_pre(int k, int x)
{
	if(!k) return -INF;
	if(key[k]>=x) return get_pre(sn[k][0],x);
	return max(key[k],get_pre(sn[k][1],x));
}
int get_suf(int k, int x)
{
	if(!k) return (1<<30);
	if(key[k]<=x) return get_suf(sn[k][1],x);
	return min(key[k],get_suf(sn[k][0],x));
}
inline long long read()
{
    long long 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 * 10 + ch - '0';
        ch = getchar();
    }
    return x * f;
}
int main()
{
	srand(time(0));
	n=read();
	while(n--)
	{
		opt=read();
		x=read();
		if(opt==1)
		{
			insert(k,x);
		}
		if(opt==2)
		{
			del(k,x);
		}
		if(opt==3)
		{
			printf("%d\n",get_rnk(k,x));
		}
		if(opt==4)
		{
			printf("%d\n",get_val(k,x));
		}
		if(opt==5)
		{
			printf("%d\n",get_pre(k,x));
		}
		if(opt==6)
		{
			printf("%d\n",get_suf(k,x));
		}
	}
	return 0;
}
2023/8/30 20:47
加载中...