求助 替罪羊树 MLE 只有8pts
查看原帖
求助 替罪羊树 MLE 只有8pts
396994
Winston12321_楼主2023/3/25 17:10

代码如下(不好读的地方请留言):

#include <iostream>
#include <vector>
using namespace std;
const int MAXN=100010;
const double alpha=0.7;
int n;
int opt;
int read()
{
	int s=0,f=1;
	char ch=getchar();
	while(!isdigit(ch))
	{
		if(ch=='-') f=-1;
		ch=getchar();
	}
	while(isdigit(ch))
	{
		s=(s<<3)+(s<<1)+(ch^48);
		ch=getchar();
	}
	return f*s;
}
void write(int x)
{
	if(x<0) putchar('-'),x=-x;
	if(x>9) write(x/10);
	putchar(x%10+48);
}
int cnt=1;
struct BST
{
	int lson,rson,val,num,size;
}t[MAXN];
vector<int>FP,FN,FV;
int flatten(int pos)
{
	if(t[pos].lson) flatten(t[pos].lson);
	int id=FP.size();
	if(t[pos].num)
	{
		FP.push_back(pos);
		FN.push_back(t[pos].num);
		FV.push_back(t[pos].val);
	}
	if(t[pos].rson) flatten(t[pos].rson);
	return id;
}
void rebuild(int pos,int l,int r)
{
	int mid=l+r>>1;
	t[pos].size=FN[mid];
	t[pos].val=FV[mid];
	t[pos].num=FN[mid];
	if(l<mid)
	{
		t[pos].lson=FP[l+mid-1>>1];
		rebuild(t[pos].lson,l,mid-1);
		t[pos].size+=t[t[pos].lson].size;
	}
	else t[pos].lson=0;
	if(r>mid)
	{
		t[pos].rson=FP[mid+1+r>>1];
		rebuild(t[pos].rson,mid+1,r);
		t[pos].size+=t[t[pos].rson].size;
	}
	else t[pos].rson=0;
}
void TryRe(int pos)
{
	double k=max(t[t[pos].lson].size,t[t[pos].rson].size)/double(t[pos].size);
	if(k>alpha)
	{
		FP.clear();
		FV.clear();
		FN.clear();
		int id=flatten(pos);
		swap(FP[id],FP[FP.size()-1>>1]);
		rebuild(pos,0,FP.size()-1);
	}
}
void insert(int pos,int val)
{
	++t[pos].size;
	if(!t[pos].num && !t[pos].lson && !t[pos].rson)
		return t[pos].val=val,t[pos].num=1,void();
	if(val<t[pos].val)
	{
		if(!t[pos].lson) t[pos].lson=++cnt;
		insert(t[pos].lson,val);
		return;
	}
	if(val>t[pos].val)
	{
		if(!t[pos].rson) t[pos].rson=++cnt;
		insert(t[pos].rson,val);
		return;
	}
	++t[pos].num;
	TryRe(pos);
}
void remove(int pos,int val)
{
	--t[pos].size;
	if(val<t[pos].val) remove(t[pos].lson,val);
	else if(val>t[pos].val) remove(t[pos].rson,val);
	else --t[pos].num;
	TryRe(pos);
}
int countl(int pos,int val)
{
	if(val<t[pos].val) return t[pos].lson?countl(t[pos].lson,val):0;
	if(val>t[pos].val) return t[pos].num+t[t[pos].lson].size+(t[pos].rson?countl(t[pos].rson,val):0);
	return t[t[pos].lson].size;
}
int countg(int pos,int val)
{
	if(val>t[pos].val) return t[pos].rson?countg(t[pos].rson,val):0;
	if(val<t[pos].val) return t[pos].num+t[t[pos].rson].size+(t[pos].lson?countg(t[pos].lson,val):0);
	return t[t[pos].rson].size;
}
int kth(int pos,int k)
{
	if(t[t[pos].lson].size>=k) return kth(t[pos].lson,k);
	if(t[pos].num+t[t[pos].lson].size<k) return kth(t[pos].rson,k);
	return t[pos].val;
}
int main()
{
	n=read();
	while(n--)
	{
		opt=read();
		if(opt==1) insert(1,read());
		else if(opt==2) remove(1,read());
		else if(opt==3) write(countl(1,read())+1),putchar('\n');
		else if(opt==4) write(kth(1,read())),putchar('\n');
		else if(opt==5) write(kth(1,countl(1,read()))),putchar('\n');
		else write(kth(1,t[1].size-countg(1,read())+1)),putchar('\n');
	}
	return 0;
}
2023/3/25 17:10
加载中...